初始化

初始化项目
This commit is contained in:
truewhile
2026-08-23 22:12:32 +08:00
parent 0bcb1fec87
commit 71bf60c69c
631 changed files with 2121 additions and 76050 deletions
-265
View File
@@ -1,265 +0,0 @@
// Package service — AI integration (OpenAI-compatible chat completions).
//
// AIService is a thin wrapper around any OpenAI-compatible REST endpoint
// (OpenAI, DeepSeek, Qwen, Ollama, …). Today we expose two operations:
//
// - SmartSearch: interpret a free-form Chinese / English query and
// return a normalised JSON intent the React UI can
// translate into filter params.
// - Recommend: given a list of recently-watched titles, generate
// a short list of "you might like…" recommendations.
//
// The service is disabled (every method returns nil) when ai.enabled is
// false or ai.api_key is empty.
package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
)
// AIService talks to an OpenAI-compatible chat-completions endpoint.
type AIService struct {
cfg *config.Config
log *zap.Logger
client *http.Client
apiConfig *APIConfigService
}
// NewAIService is the constructor.
func NewAIService(cfg *config.Config, log *zap.Logger, apiConfig *APIConfigService) *AIService {
timeout := time.Duration(cfg.AI.Timeout) * time.Second
if timeout <= 0 {
timeout = 30 * time.Second
}
return &AIService{
cfg: cfg,
log: log,
apiConfig: apiConfig,
client: NewExternalHTTPClient(timeout),
}
}
// Enabled reports whether the AI integration is configured.
func (a *AIService) Enabled() bool {
return a.cfg.AI.Enabled && strings.TrimSpace(a.cfg.AI.APIKey) != ""
}
// EnabledFor reports whether the AI integration is configured for a request.
func (a *AIService) EnabledFor(ctx context.Context) bool {
return a.resolveRuntimeConfig(ctx).Enabled
}
// AIStatus is returned to the UI for connection-state display.
type AIStatus struct {
Enabled bool `json:"enabled"`
Provider string `json:"provider"`
Model string `json:"model"`
}
// Status resolves live database-backed AI config for the UI.
func (a *AIService) Status(ctx context.Context) AIStatus {
cfg := a.resolveRuntimeConfig(ctx)
return AIStatus{Enabled: cfg.Enabled, Provider: cfg.Provider, Model: cfg.Model}
}
// SearchIntent is the structured output the smart search endpoint returns.
type SearchIntent struct {
Query string `json:"query"`
Year int `json:"year,omitempty"`
Genre string `json:"genre,omitempty"`
Type string `json:"type,omitempty"` // movie / tv / anime / music
Sort string `json:"sort,omitempty"` // recent / rating / random
Language string `json:"language,omitempty"`
}
// SmartSearch turns a natural-language query into a structured intent.
// Returns a best-effort intent on parse failure (raw query passes through).
func (a *AIService) SmartSearch(ctx context.Context, raw string) (*SearchIntent, error) {
runtime := a.resolveRuntimeConfig(ctx)
if !runtime.Enabled {
return &SearchIntent{Query: raw}, nil
}
const sys = "You are a media-library search assistant. Read the user's query and " +
"output a JSON object with the keys: query (string), year (int, optional), " +
"genre (string, optional), type (movie|tv|anime|music, optional), sort " +
"(recent|rating|random, optional), language (zh|en, optional). Respond with " +
"JSON only, no commentary."
out, err := a.complete(ctx, runtime, sys, raw)
if err != nil {
return &SearchIntent{Query: raw}, err
}
var intent SearchIntent
if err := json.Unmarshal([]byte(out), &intent); err != nil {
// Fallback: tolerate non-JSON output by treating the raw text as
// the cleaned query.
intent.Query = strings.TrimSpace(out)
}
if intent.Query == "" {
intent.Query = raw
}
return &intent, nil
}
// Recommend builds a short comma-separated list of titles given the user's
// history. The first call is intentionally best-effort: a future iteration
// may chain media DB lookups onto each suggestion.
func (a *AIService) Recommend(ctx context.Context, history []string, max int) ([]string, error) {
runtime := a.resolveRuntimeConfig(ctx)
if !runtime.Enabled || len(history) == 0 {
return nil, nil
}
if max <= 0 || max > 20 {
max = 8
}
sys := fmt.Sprintf("You are a film / TV recommendation assistant. Reply with %d "+
"comma-separated titles only, no commentary, in the same language as the input.", max)
usr := "I recently watched: " + strings.Join(history, "; ")
out, err := a.complete(ctx, runtime, sys, usr)
if err != nil {
return nil, err
}
parts := strings.Split(out, ",")
titles := make([]string, 0, len(parts))
for _, p := range parts {
p = strings.TrimSpace(p)
p = strings.Trim(p, "\"'`")
if p != "" {
titles = append(titles, p)
}
}
return titles, nil
}
// complete is the shared helper — POST /v1/chat/completions.
func (a *AIService) complete(ctx context.Context, runtime aiRuntimeConfig, system, user string) (string, error) {
payload := map[string]any{
"model": runtime.Model,
"temperature": 0.2,
"messages": []map[string]string{
{"role": "system", "content": system},
{"role": "user", "content": user},
},
}
body, _ := json.Marshal(payload)
endpoint := strings.TrimRight(runtime.APIBase, "/") + "/chat/completions"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+runtime.APIKey)
resp, err := a.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
raw, _ := io.ReadAll(resp.Body)
return "", fmt.Errorf("ai %d: %s", resp.StatusCode, strings.TrimSpace(string(raw)))
}
type choice struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
}
var out struct {
Choices []choice `json:"choices"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
return "", err
}
if len(out.Choices) == 0 {
return "", errors.New("ai: empty completion")
}
return strings.TrimSpace(out.Choices[0].Message.Content), nil
}
// ChatTurn is one message in a multi-turn assistant transcript.
type ChatTurn struct {
Role string `json:"role"`
Content string `json:"content"`
}
// Chat sends an entire transcript to the LLM. When the AI is disabled
// we return a deterministic offline reply so the assistant UI still
// has something to render.
func (a *AIService) Chat(ctx context.Context, history []ChatTurn) (string, error) {
runtime := a.resolveRuntimeConfig(ctx)
if !runtime.Enabled || len(history) == 0 {
return offlineReply(history), nil
}
// Build a chat/completions payload preserving the history order.
msgs := make([]map[string]string, 0, len(history)+1)
msgs = append(msgs, map[string]string{
"role": "system",
"content": "You are MediaStationGo's helpful media-library assistant. " +
"Respond concisely in the user's language. " +
"Never invent file paths or media that don't exist.",
})
for _, t := range history {
msgs = append(msgs, map[string]string{"role": t.Role, "content": t.Content})
}
payload := map[string]any{
"model": runtime.Model,
"temperature": 0.4,
"messages": msgs,
}
body, _ := json.Marshal(payload)
endpoint := strings.TrimRight(runtime.APIBase, "/") + "/chat/completions"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+runtime.APIKey)
resp, err := a.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
raw, _ := io.ReadAll(resp.Body)
return "", fmt.Errorf("ai %d: %s", resp.StatusCode, strings.TrimSpace(string(raw)))
}
type choice struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
}
var out struct {
Choices []choice `json:"choices"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
return "", err
}
if len(out.Choices) == 0 {
return "", errors.New("ai: empty completion")
}
return strings.TrimSpace(out.Choices[0].Message.Content), nil
}
// offlineReply returns a deterministic stand-in response so the UI's
// chat view stays functional when the AI provider is not configured.
func offlineReply(history []ChatTurn) string {
if len(history) == 0 {
return "Hi — AI provider is not configured. Set up OpenAI/DeepSeek in API Configs to chat with me."
}
last := history[len(history)-1].Content
if len(last) > 80 {
last = last[:80] + "…"
}
return "(offline) Heard: " + last + "\n请在 API 配置中接入 LLM 后重试。"
}
-60
View File
@@ -1,60 +0,0 @@
package service
import (
"context"
"strings"
"go.uber.org/zap"
)
type aiRuntimeConfig struct {
Enabled bool
Provider string
APIBase string
APIKey string
Model string
}
func (a *AIService) resolveRuntimeConfig(ctx context.Context) aiRuntimeConfig {
out := aiRuntimeConfig{
Enabled: a.cfg.AI.Enabled && strings.TrimSpace(a.cfg.AI.APIKey) != "",
Provider: strings.TrimSpace(a.cfg.AI.Provider),
APIBase: strings.TrimSpace(a.cfg.AI.APIBase),
APIKey: strings.TrimSpace(a.cfg.AI.APIKey),
Model: strings.TrimSpace(a.cfg.AI.Model),
}
if out.Provider == "" {
out.Provider = "openai"
}
if out.APIBase == "" {
out.APIBase = "https://api.openai.com/v1"
}
if out.Model == "" {
out.Model = "gpt-4o-mini"
}
if a.apiConfig != nil {
resolved, err := a.apiConfig.Resolve(ctx, "openai")
if err != nil {
if a.log != nil {
a.log.Warn("ai: failed to resolve openai api config", zap.Error(err))
}
return out
}
if resolved.BaseURL != "" {
out.APIBase = strings.TrimSpace(resolved.BaseURL)
}
if resolved.APIKey != "" {
out.APIKey = strings.TrimSpace(resolved.APIKey)
}
if resolved.Enabled && out.APIKey != "" {
out.Enabled = true
out.Provider = "openai"
return out
}
if !resolved.Enabled && (resolved.APIKey != "" || resolved.BaseURL != "" || resolved.Extra != "") {
out.Enabled = false
}
}
return out
}
-70
View File
@@ -1,70 +0,0 @@
package service
import (
"context"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
crypto := NewCryptoService("test-secret", zap.NewNop())
apiConfig := NewAPIConfigService(zap.NewNop(), repo, crypto)
key := "sk-test"
baseURL := "https://example.test/v1"
enabled := true
if _, err := apiConfig.Update(context.Background(), "openai", APIConfigPatch{
APIKey: &key,
BaseURL: &baseURL,
Enabled: &enabled,
}); err != nil {
t.Fatal(err)
}
ai := NewAIService(&config.Config{
AI: config.AIConfig{
Enabled: false,
Model: "gpt-4o-mini",
},
}, zap.NewNop(), apiConfig)
status := ai.Status(context.Background())
if !status.Enabled {
t.Fatalf("AI status disabled, want enabled from database config")
}
if status.Provider != "openai" {
t.Fatalf("provider = %q, want openai", status.Provider)
}
}
func TestAIStatusHonorsDisabledDatabaseOpenAIConfig(t *testing.T) {
db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
apiConfig := NewAPIConfigService(zap.NewNop(), repo, NewCryptoService("test-secret", zap.NewNop()))
key := "sk-test"
enabled := false
if _, err := apiConfig.Update(context.Background(), "openai", APIConfigPatch{
APIKey: &key,
Enabled: &enabled,
}); err != nil {
t.Fatal(err)
}
ai := NewAIService(&config.Config{
AI: config.AIConfig{
Enabled: true,
APIKey: "sk-file",
Model: "gpt-4o-mini",
},
}, zap.NewNop(), apiConfig)
if ai.Status(context.Background()).Enabled {
t.Fatalf("AI status enabled, want disabled when database config is explicitly disabled")
}
}
-203
View File
@@ -1,203 +0,0 @@
// Package service — Aria2 下载适配器。
//
// Aria2Adapter 实现了 DownloadAdapter 接口,通过 Aria2 JSON-RPC API
// 管理下载任务。
package service
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"sync"
"time"
)
var errAria2ListUnavailable = errors.New("aria2 task list unavailable")
// Aria2Adapter 是 Aria2 的 DownloadAdapter 实现。
type Aria2Adapter struct {
mu sync.Mutex
cfg DownloadClientConfig
client *http.Client
idSeq int
}
// NewAria2Adapter 创建新的 Aria2 适配器。
func NewAria2Adapter() *Aria2Adapter {
return &Aria2Adapter{
client: NewInternalHTTPClient(20 * time.Second),
}
}
// AddTorrent 通过 URL 添加种子或磁力链接。
func (a *Aria2Adapter) AddTorrent(ctx context.Context, torrentURL, savePath string) (string, error) {
a.mu.Lock()
defer a.mu.Unlock()
// Aria2 addUri 的参数: [secret, [uris], options]
uris := []string{torrentURL}
options := map[string]string{}
if savePath != "" {
options["dir"] = savePath
}
result, err := a.rpcLocked(ctx, "aria2.addUri", []interface{}{uris, options})
if err != nil {
return "", err
}
var gid string
if err := json.Unmarshal(result, &gid); err != nil {
return "", err
}
return gid, nil
}
// AddMagnet 通过磁力链接添加下载。
func (a *Aria2Adapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) {
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()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "aria2.pause", []interface{}{hash})
return err
}
// Resume 恢复下载任务(通过 GID)。
func (a *Aria2Adapter) Resume(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "aria2.unpause", []interface{}{hash})
return err
}
// Remove 移除下载任务。
func (a *Aria2Adapter) Remove(ctx context.Context, hash string, deleteFiles bool) error {
a.mu.Lock()
defer a.mu.Unlock()
_, 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)
}
_ = deleteFiles // aria2 RPC removes the task/result but has no delete-local-data flag.
return nil
}
// List 列出所有活动/等待/已停止的任务。
func (a *Aria2Adapter) List(ctx context.Context, filter string) ([]TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
var allResults []TorrentInfo
var listErrs []error
var successfulCalls int
// 获取活动任务
active, err := a.rpcLocked(ctx, "aria2.tellActive", []interface{}{
[]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", "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", "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 != "" {
filtered := make([]TorrentInfo, 0, len(allResults))
for _, item := range allResults {
if strings.EqualFold(item.State, filter) {
filtered = append(filtered, item)
}
}
return filtered, errors.Join(listErrs...)
}
return allResults, errors.Join(listErrs...)
}
// GetInfo 获取单个任务信息。
func (a *Aria2Adapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
result, err := a.rpcLocked(ctx, "aria2.tellStatus", []interface{}{
hash,
[]string{"gid", "bittorrent", "files", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections"},
})
if err != nil {
return nil, err
}
var item map[string]interface{}
if err := json.Unmarshal(result, &item); err != nil {
return nil, err
}
info := a.parseSingleItem(item)
if info == nil {
return nil, fmt.Errorf("task %s not found", hash)
}
return info, nil
}
-116
View File
@@ -1,116 +0,0 @@
package service
import (
"encoding/json"
"time"
)
// parseAria2Items 解析 Aria2 返回的任务列表。
func (a *Aria2Adapter) parseAria2Items(raw json.RawMessage) []TorrentInfo {
var items []map[string]interface{}
if err := json.Unmarshal(raw, &items); err != nil {
return nil
}
result := make([]TorrentInfo, 0, len(items))
for _, item := range items {
info := a.parseSingleItem(item)
if info != nil {
result = append(result, *info)
}
}
return result
}
// parseSingleItem 解析单个 Aria2 任务项。
func (a *Aria2Adapter) parseSingleItem(item map[string]interface{}) *TorrentInfo {
gid := strVal(item["gid"])
totalLength := toInt64(item["totalLength"])
completedLength := toInt64(item["completedLength"])
dlSpeed := toInt64(item["downloadSpeed"])
upSpeed := toInt64(item["uploadSpeed"])
status := strVal(item["status"])
dir := strVal(item["dir"])
numSeeders := int(toInt64(item["numSeeders"]))
connections := int(toInt64(item["connections"]))
var name 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"])
}
contentPath = downloaderPayloadPath(dir, name)
}
if name == "" {
// 尝试从 files 获取文件名
if files, ok := item["files"].([]interface{}); ok && len(files) > 0 {
if f, ok := files[0].(map[string]interface{}); ok {
filePath := strVal(f["path"])
if filePath != "" {
contentPath = filePath
name = downloaderPathBase(filePath)
}
if name == "" {
name = strVal(f["uris"])
}
}
}
}
if name == "" {
name = gid
}
var progress float64
if totalLength > 0 {
progress = float64(completedLength) / float64(totalLength)
}
// Aria2 状态映射
state := canonicalTorrentState(aria2StatusStr(status), progress)
return &TorrentInfo{
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,
}
}
// aria2StatusStr 将 Aria2 状态转为可读字符串。
func aria2StatusStr(status string) string {
switch status {
case "active":
return "downloading"
case "waiting":
return "queued"
case "paused":
return "paused"
case "error":
return "error"
case "complete":
return "seeding"
case "removed":
return "removed"
default:
return status
}
}
func aria2MaxInt(a, b int) int {
if a > b {
return a
}
return b
}
-166
View File
@@ -1,166 +0,0 @@
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
)
// aria2Request 是 Aria2 JSON-RPC 请求结构。
type aria2Request struct {
JSONRPC string `json:"jsonrpc"`
Method string `json:"method"`
ID string `json:"id"`
Params []interface{} `json:"params"`
}
// aria2Response 是 Aria2 JSON-RPC 响应结构。
type aria2Response struct {
JSONRPC string `json:"jsonrpc"`
ID string `json:"id"`
Result json.RawMessage `json:"result"`
Error *aria2Error `json:"error"`
}
// aria2Error 是 Aria2 JSON-RPC 错误结构。
type aria2Error struct {
Code int `json:"code"`
Message string `json:"message"`
}
// Initialize 配置并初始化 Aria2 RPC 连接。
func (a *Aria2Adapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
a.mu.Lock()
defer a.mu.Unlock()
endpoint, err := normalizeDownloadClientEndpoint("aria2", cfg.Host)
if err != nil {
return err
}
cfg.Host = endpoint
a.cfg = cfg
a.idSeq = 0
return a.getVersionLocked(ctx)
}
// Ping 测试连接。
func (a *Aria2Adapter) Ping(ctx context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
return a.getVersionLocked(ctx)
}
// getVersionLocked 内部版本检查(调用者必须持有锁)。
func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host)
if err != nil {
return err
}
req := &aria2Request{
JSONRPC: "2.0",
Method: "aria2.getVersion",
ID: a.nextID(),
Params: []interface{}{"token:" + a.cfg.Password},
}
body, err := json.Marshal(req)
if err != nil {
return err
}
httpReq, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
if err != nil {
return err
}
httpReq.Header.Set("Content-Type", "application/json")
if a.cfg.Username != "" {
httpReq.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(httpReq)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return fmt.Errorf("aria2 rpc: %d", resp.StatusCode)
}
return nil
}
// rpcLocked 发送 JSON-RPC 请求(调用者必须持有锁)。
func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []interface{}) (json.RawMessage, error) {
rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host)
if err != nil {
return nil, err
}
if params == nil {
params = []interface{}{}
}
// 如果 secret 不在 params 中,添加到第一位
if len(params) > 0 {
if secret, ok := params[0].(string); ok && strings.HasPrefix(secret, "token:") {
// 已经有 secret
} else {
newParams := make([]interface{}, 0, len(params)+1)
newParams = append(newParams, "token:"+a.cfg.Password)
newParams = append(newParams, params...)
params = newParams
}
} else {
params = []interface{}{"token:" + a.cfg.Password}
}
req := &aria2Request{
JSONRPC: "2.0",
Method: method,
ID: a.nextID(),
Params: params,
}
body, err := json.Marshal(req)
if err != nil {
return nil, err
}
httpReq, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
httpReq.Header.Set("Content-Type", "application/json")
if a.cfg.Username != "" {
httpReq.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(httpReq)
if err != nil {
return nil, err
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, err
}
var rpcResp aria2Response
if err := json.Unmarshal(respBody, &rpcResp); err != nil {
return nil, err
}
if rpcResp.Error != nil {
return nil, fmt.Errorf("aria2 rpc error [%d]: %s", rpcResp.Error.Code, rpcResp.Error.Message)
}
return rpcResp.Result, nil
}
// nextID 生成递增的请求 ID。
func (a *Aria2Adapter) nextID() string {
a.idSeq++
return fmt.Sprintf("msg-%d", a.idSeq)
}
-209
View File
@@ -1,209 +0,0 @@
// Package service — multi-turn AI assistant chat.
//
// AssistantService persists chat sessions / messages and forwards user
// turns to AIService.Chat() for the actual LLM call. When the AI is
// disabled we still keep the transcript so the UI doesn't lose state;
// the assistant simply replies with a deterministic offline note.
//
// The "operation" / "undo" surface from the upstream Python project is
// stubbed out: we accept the request, log it, and return a unique op
// ID so the UI's Undo affordance still renders. Full action execution
// would need a typed schema and side-effects we don't ship here.
package service
import (
"context"
"errors"
"strings"
"time"
"github.com/google/uuid"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// AssistantService coordinates AssistantSession + AssistantMessage rows
// against the underlying AIService.
type AssistantService struct {
log *zap.Logger
repo *repository.Container
ai *AIService
}
// NewAssistantService is the constructor.
func NewAssistantService(log *zap.Logger, repo *repository.Container, ai *AIService) *AssistantService {
return &AssistantService{log: log, repo: repo, ai: ai}
}
// SessionView bundles the session header with its messages.
type SessionView struct {
Session model.AssistantSession `json:"session"`
Messages []model.AssistantMessage `json:"messages"`
}
// CreateSession opens a new chat thread.
func (s *AssistantService) CreateSession(ctx context.Context, userID, title string) (*model.AssistantSession, error) {
if title == "" {
title = "New chat"
}
sess := &model.AssistantSession{UserID: userID, Title: title}
if err := s.repo.Assistant.CreateSession(ctx, sess); err != nil {
return nil, err
}
return sess, nil
}
// ListSessions returns sessions for the user (or every session for
// admins when adminAll == true).
func (s *AssistantService) ListSessions(ctx context.Context, userID string, adminAll bool) ([]model.AssistantSession, error) {
if adminAll {
return s.repo.Assistant.ListSessions(ctx, "")
}
return s.repo.Assistant.ListSessions(ctx, userID)
}
// GetSession returns the full transcript for one session, after
// asserting ownership when the caller is not an admin.
func (s *AssistantService) GetSession(ctx context.Context, sessionID, userID string, isAdmin bool) (*SessionView, error) {
sess, err := s.repo.Assistant.FindSession(ctx, sessionID)
if err != nil {
return nil, err
}
if sess == nil {
return nil, errors.New("session not found")
}
if !isAdmin && sess.UserID != userID {
return nil, errors.New("forbidden")
}
msgs, err := s.repo.Assistant.ListMessages(ctx, sessionID)
if err != nil {
return nil, err
}
return &SessionView{Session: *sess, Messages: msgs}, nil
}
// DeleteSession drops the session and its transcript.
func (s *AssistantService) DeleteSession(ctx context.Context, sessionID, userID string, isAdmin bool) error {
sess, err := s.repo.Assistant.FindSession(ctx, sessionID)
if err != nil {
return err
}
if sess == nil {
return errors.New("session not found")
}
if !isAdmin && sess.UserID != userID {
return errors.New("forbidden")
}
return s.repo.Assistant.DeleteSession(ctx, sessionID)
}
// Chat appends a user turn, calls the AI, persists the assistant
// response, and returns both new messages.
func (s *AssistantService) Chat(ctx context.Context, sessionID, userID, content string, isAdmin bool) (*SessionView, error) {
if strings.TrimSpace(content) == "" {
return nil, errors.New("content required")
}
sess, err := s.repo.Assistant.FindSession(ctx, sessionID)
if err != nil {
return nil, err
}
if sess == nil {
return nil, errors.New("session not found")
}
if !isAdmin && sess.UserID != userID {
return nil, errors.New("forbidden")
}
// Append the user turn.
userMsg := &model.AssistantMessage{
SessionID: sessionID,
Role: "user",
Content: strings.TrimSpace(content),
}
if err := s.repo.Assistant.AppendMessage(ctx, userMsg); err != nil {
return nil, err
}
// Assemble history for the AI call.
prior, _ := s.repo.Assistant.ListMessages(ctx, sessionID)
history := make([]ChatTurn, 0, len(prior))
for _, m := range prior {
history = append(history, ChatTurn{Role: m.Role, Content: m.Content})
}
// Call the LLM (or fall back to a deterministic offline reply).
reply, err := s.ai.Chat(ctx, history)
if err != nil {
s.log.Warn("assistant chat failed", zap.Error(err))
reply = "(AI 暂未配置或调用失败,请稍后再试。)"
}
asstMsg := &model.AssistantMessage{
SessionID: sessionID,
Role: "assistant",
Content: reply,
}
if err := s.repo.Assistant.AppendMessage(ctx, asstMsg); err != nil {
return nil, err
}
return s.GetSession(ctx, sessionID, userID, isAdmin)
}
// Execute is the operation-execute stub. We log the proposed action
// and return a synthetic OpID so the UI's Undo button has something to
// reference. Real execution would need a typed action schema we don't
// ship here.
func (s *AssistantService) Execute(ctx context.Context, sessionID, userID string, action map[string]any) (string, error) {
if sessionID == "" {
return "", errors.New("session_id required")
}
opID := uuid.NewString()
s.log.Info("assistant.execute (stub)",
zap.String("session_id", sessionID),
zap.String("user_id", userID),
zap.String("op_id", opID),
zap.Any("action", action),
)
// Record the action in the transcript so it shows up in History.
_ = s.repo.Assistant.AppendMessage(ctx, &model.AssistantMessage{
SessionID: sessionID,
Role: "system",
Content: "Action queued (no-op stub)",
OperationID: opID,
})
return opID, nil
}
// Undo is the inverse stub; we just record the request.
func (s *AssistantService) Undo(ctx context.Context, opID string) error {
s.log.Info("assistant.undo (stub)", zap.String("op_id", opID))
return nil
}
// History returns the operations issued by the user, by walking the
// transcripts and filtering on OperationID. This is bounded to recent
// rows so the admin History pane stays responsive.
func (s *AssistantService) History(ctx context.Context, userID string, isAdmin bool) ([]map[string]any, error) {
sessions, err := s.ListSessions(ctx, userID, isAdmin)
if err != nil {
return nil, err
}
out := make([]map[string]any, 0)
cutoff := time.Now().AddDate(0, 0, -30)
for _, sess := range sessions {
msgs, _ := s.repo.Assistant.ListMessages(ctx, sess.ID)
for _, m := range msgs {
if m.OperationID == "" || m.CreatedAt.Before(cutoff) {
continue
}
out = append(out, map[string]any{
"op_id": m.OperationID,
"session": sess.ID,
"created_at": m.CreatedAt,
"content": m.Content,
})
}
}
return out, nil
}
+8 -3
View File
@@ -40,10 +40,15 @@ var (
ErrUserExpired = errors.New("user account has expired")
)
// MaxUsers is kept for compatibility with tests and callers; dynamic runtime
// checks use LicensedMaxUsers so official licensed builds can raise the quota.
// MaxUsers 是单实例允许的最大用户数(开源版本固定上限)。
const MaxUsers = OpenSourceUserLimit
// OpenSourceUserLimit 是开源版本的用户数上限。
const OpenSourceUserLimit = 20
// UserLimit 是注册/邀请码发放时的固定用户数上限(授权管理已移除,固定为开源上限)。
const UserLimit = OpenSourceUserLimit
// SeedAdmin makes sure at least one admin user exists. It mirrors the
// legacy default behaviour: if no admin row is found we create
// `admin / admin123` (overridable through ADMIN_INITIAL_PASSWORD) and warn.
@@ -100,7 +105,7 @@ func (s *AuthService) Register(ctx context.Context, username, password string) (
}
if n, err := s.repo.User.Count(ctx); err != nil {
return nil, nil, err
} else if n >= LicensedMaxUsers(ctx, s.repo) {
} else if n >= UserLimit {
return nil, nil, ErrUserLimitReached
}
hash, err := hashPassword(password)
+1 -55
View File
@@ -2,7 +2,6 @@ package service
import (
"context"
"encoding/json"
"errors"
"fmt"
"path/filepath"
@@ -21,7 +20,7 @@ import (
func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *ProfileService, *PermissionService) {
t.Helper()
db := newServiceTestDB(t, &model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{})
db := newServiceTestDB(t, &model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.Setting{})
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
@@ -58,44 +57,6 @@ func TestRegisterRejectsMoreThanTwentyUsers(t *testing.T) {
}
}
func TestRegisterUsesLicensedUserLimit(t *testing.T) {
ctx := context.Background()
repos, auth, _, _ := newAuthTestServices(t)
maxUsers := 25
state := LicenseActivationState{Valid: true, LicenseType: "plus", MaxUsers: &maxUsers}
raw, _ := json.Marshal(state)
if err := repos.Setting.Set(ctx, LicenseSettingActivation, string(raw)); err != nil {
t.Fatal(err)
}
for i := 0; i < OpenSourceUserLimit; i++ {
if err := repos.User.Create(ctx, &model.User{
Username: fmt.Sprintf("licensed-%02d", i),
PasswordHash: "hash",
Role: "user",
Tier: "free",
}); err != nil {
t.Fatal(err)
}
}
if _, _, err := auth.Register(ctx, "extra", "password"); err != nil {
t.Fatalf("licensed user limit should allow user 21: %v", err)
}
}
func TestLicensedMaxUsersCanBeUnlimited(t *testing.T) {
ctx := context.Background()
repos, _, _, _ := newAuthTestServices(t)
state := LicenseActivationState{Valid: true, LicenseType: "enterprise", UnlimitedUsers: true}
raw, _ := json.Marshal(state)
if err := repos.Setting.Set(ctx, LicenseSettingActivation, string(raw)); err != nil {
t.Fatal(err)
}
if got := LicensedMaxUsers(ctx, repos); got <= 1_000_000 {
t.Fatalf("unlimited license should return a very high limit, got %d", got)
}
}
func TestRegisterDefaultsAdultLibrariesHidden(t *testing.T) {
_, auth, _, _ := newAuthTestServices(t)
user, _, err := auth.Register(context.Background(), "viewer", "password")
@@ -114,14 +75,6 @@ func TestDeletedUserCanBeRecreatedWithSameUsername(t *testing.T) {
if err != nil {
t.Fatalf("register old user: %v", err)
}
if err := repos.DB.Create(&model.TelegramBinding{
TelegramUserID: 10001,
TelegramName: "@viewer",
ChatID: 10001,
UserID: user.ID,
}).Error; err != nil {
t.Fatalf("create telegram binding: %v", err)
}
if err := repos.User.Delete(ctx, user.ID); err != nil {
t.Fatalf("delete user: %v", err)
}
@@ -136,13 +89,6 @@ func TestDeletedUserCanBeRecreatedWithSameUsername(t *testing.T) {
if _, err := auth.Login(ctx, "viewer", "new-password"); err != nil {
t.Fatalf("login recreated user: %v", err)
}
var bindings int64
if err := repos.DB.Model(&model.TelegramBinding{}).Where("telegram_user_id = ?", 10001).Count(&bindings).Error; err != nil {
t.Fatalf("count bindings: %v", err)
}
if bindings != 0 {
t.Fatalf("deleted user telegram bindings should be removed, got %d", bindings)
}
}
func TestRegisterReleasesLegacySoftDeletedUsername(t *testing.T) {
-77
View File
@@ -1,77 +0,0 @@
// Package service — Bangumi discovery calendar.
package service
import (
"context"
"strconv"
"strings"
)
// Calendar returns Bangumi's public on-air anime calendar as a recommendation
// rail. It needs no token, but NewBangumiProvider still attaches one when set.
func (b *BangumiProvider) Calendar(ctx context.Context) ([]ExternalMediaResult, error) {
type subject struct {
ID int `json:"id"`
Name string `json:"name"`
NameCN string `json:"name_cn"`
Summary string `json:"summary"`
AirDate string `json:"air_date"`
Images struct {
Large string `json:"large"`
Common string `json:"common"`
} `json:"images"`
Rating struct {
Score float32 `json:"score"`
} `json:"rating"`
}
type day struct {
Items []subject `json:"items"`
}
var days []day
if err := b.getJSON(ctx, b.base+"/calendar", &days); err != nil {
return nil, err
}
out := make([]ExternalMediaResult, 0, 24)
seen := map[int]struct{}{}
for _, day := range days {
for _, item := range day.Items {
if _, ok := seen[item.ID]; ok {
continue
}
seen[item.ID] = struct{}{}
title := strings.TrimSpace(item.NameCN)
if title == "" {
title = strings.TrimSpace(item.Name)
}
if title == "" {
continue
}
poster := item.Images.Large
if poster == "" {
poster = item.Images.Common
}
poster = normalizeBangumiImageURL(poster)
year := 0
if len(item.AirDate) >= 4 {
year, _ = strconv.Atoi(item.AirDate[:4])
}
out = append(out, ExternalMediaResult{
Source: "bangumi",
MediaType: "anime",
Title: title,
OriginalName: item.Name,
Overview: item.Summary,
PosterURL: poster,
Year: year,
Rating: item.Rating.Score,
BangumiID: item.ID,
SubscribeKeyword: buildSubscribeKeyword(title, year),
SubscribeAliases: buildSubscribeAliases(title, item.Name, year),
})
if len(out) >= 24 {
return out, nil
}
}
}
return out, nil
}
-74
View File
@@ -1,74 +0,0 @@
package service
import (
"context"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// BootCloudLibraries optionally scans cloud libraries after startup. It is
// disabled by default for huge cloud mounts; normal automatic refresh is handled
// by the nightly cloud_sync scheduler window, and operators can still scan
// manually at any time.
func (c *Container) BootCloudLibraries(ctx context.Context) {
if c == nil || c.Repo == nil || c.Scan == nil {
return
}
if !bootCloudLibraryScanEnabled(ctx, c.Repo) {
c.Log.Info("boot: cloud library scans disabled; use manual scan or nightly cloud sync")
return
}
libs, err := c.Repo.Library.List(ctx)
if err != nil {
c.Log.Warn("boot cloud libraries: list failed", zap.Error(err))
return
}
libs = FilterScannableCloudLibraries(ctx, c.Repo, libs)
cloudLibs := make([]model.Library, 0)
for _, lib := range libs {
if !lib.Enabled {
continue
}
if _, ok := ParseCloudLibraryMount(lib.Path); ok {
cloudLibs = append(cloudLibs, lib)
}
}
if len(cloudLibs) == 0 {
return
}
c.Log.Info("boot: scheduling cloud library scans", zap.Int("count", len(cloudLibs)))
// 延迟3秒后启动,避免和系统初始化任务冲突
time.AfterFunc(3*time.Second, func() {
go c.runBootCloudLibraryScanQueue(cloudLibs)
})
}
func (c *Container) runBootCloudLibraryScanQueue(cloudLibs []model.Library) {
for _, lib := range cloudLibs {
libID := lib.ID
libName := lib.Name
scanCtx, cancel := cloudScanContext(context.Background(), cloudScanTimeout(context.Background(), c.Repo, 24*time.Hour))
c.Log.Info("boot: scanning cloud library", zap.String("id", libID), zap.String("name", libName))
if _, err := c.Scan.ScanLibraryWithoutAutoScrape(scanCtx, libID); err != nil {
c.Log.Warn("boot: cloud library scan failed", zap.String("id", libID), zap.String("name", libName), zap.Error(err))
} else {
c.Log.Info("boot: cloud library scan completed", zap.String("id", libID), zap.String("name", libName))
}
cancel()
}
}
func bootCloudLibraryScanEnabled(ctx context.Context, repo *repository.Container) bool {
if repo == nil || repo.Setting == nil {
return false
}
value, err := repo.Setting.Get(ctx, "cloud.boot_scan_enabled")
if err != nil {
return false
}
return parseBoolSetting(value, false)
}
-102
View File
@@ -1,102 +0,0 @@
package service
import (
"context"
"strings"
"time"
"go.uber.org/zap"
)
const cloudStorageMissingConfigWarnPrefix = "cloud.storage.missing_config_warned."
// BootCloudStorageHealthCheck验证所有已配置的云盘存储在启动时是否可用
func (c *Container) BootCloudStorageHealthCheck(ctx context.Context) {
if c == nil || c.StorageCfg == nil {
return
}
configs, err := c.StorageCfg.List(ctx)
if err != nil {
c.Log.Warn("boot: cloud storage health check failed to list configs", zap.Error(err))
return
}
cloudConfigs := make([]StorageView, 0)
for _, cfg := range configs {
if cfg.Enabled && IsAdminCloudConfigurable(cfg.Type) {
cloudConfigs = append(cloudConfigs, cfg)
}
}
if len(cloudConfigs) == 0 {
c.Log.Info("boot: no enabled cloud storage configured")
return
}
c.Log.Info("boot: checking cloud storage health", zap.Int("count", len(cloudConfigs)))
for _, cfg := range cloudConfigs {
go func(typ string) {
checkCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
provider, err := c.StorageCfg.CloudProvider(checkCtx, typ)
if err != nil {
if c.warnMissingCloudStorageConfigOnce(checkCtx, typ, err) {
return
}
c.Log.Warn("boot: cloud storage unavailable", zap.String("type", typ), zap.Error(err))
return
}
if err := provider.Ping(checkCtx); err != nil {
if c.warnMissingCloudStorageConfigOnce(checkCtx, typ, err) {
return
}
c.Log.Warn("boot: cloud storage ping failed", zap.String("type", typ), zap.Error(err))
} else {
c.Log.Info("boot: cloud storage healthy", zap.String("type", typ))
}
}(cfg.Type)
}
}
func (c *Container) warnMissingCloudStorageConfigOnce(ctx context.Context, typ string, err error) bool {
reason := cloudStorageMissingConfigReason(err)
if reason == "" {
return false
}
if c == nil || c.Repo == nil || c.Repo.Setting == nil {
if c != nil && c.Log != nil {
c.Log.Warn("boot: cloud storage config incomplete; skipping health check", zap.String("type", typ), zap.String("reason", reason), zap.Error(err))
}
return true
}
key := cloudStorageMissingConfigWarnPrefix + strings.TrimSpace(typ) + "." + reason
if value, getErr := c.Repo.Setting.Get(ctx, key); getErr == nil && strings.EqualFold(strings.TrimSpace(value), "true") {
return true
}
if c.Log != nil {
c.Log.Warn("boot: cloud storage config incomplete; skipping health check", zap.String("type", typ), zap.String("reason", reason), zap.Error(err))
}
if setErr := c.Repo.Setting.Set(ctx, key, "true"); setErr != nil && c.Log != nil {
c.Log.Debug("remember cloud storage config warning failed", zap.String("type", typ), zap.Error(setErr))
}
return true
}
func cloudStorageMissingConfigReason(err error) string {
if err == nil {
return ""
}
msg := strings.ToLower(strings.TrimSpace(err.Error()))
switch {
case strings.Contains(msg, "missing cookie") || (strings.Contains(msg, "missing") && strings.Contains(msg, "cookie")):
return "missing_cookie"
case strings.Contains(msg, "missing webdav url"):
return "missing_webdav_url"
default:
return ""
}
}
@@ -1,53 +0,0 @@
package service
import (
"context"
"errors"
"testing"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestCloudStorageMissingConfigReason(t *testing.T) {
cases := []struct {
err error
want string
}{
{errors.New("115: missing cookie"), "missing_cookie"},
{errors.New("openlist: missing cookie"), "missing_cookie"},
{errors.New("clouddrive2: missing WebDAV URL"), "missing_webdav_url"},
{errors.New("openlist: token expired"), ""},
}
for _, tc := range cases {
if got := cloudStorageMissingConfigReason(tc.err); got != tc.want {
t.Fatalf("reason(%q) = %q, want %q", tc.err, got, tc.want)
}
}
}
func TestWarnMissingCloudStorageConfigOncePersistsMarker(t *testing.T) {
db := newServiceTestDB(t, &model.Setting{})
core, observed := observer.New(zap.WarnLevel)
c := &Container{
Log: zap.New(core),
Repo: repository.New(db),
}
err := errors.New("115: missing cookie")
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
t.Fatal("missing config should be handled")
}
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
t.Fatal("missing config should still be classified on second call")
}
if observed.FilterMessage("boot: cloud storage config incomplete; skipping health check").Len() != 1 {
t.Fatalf("warn count = %d, want 1", observed.Len())
}
if c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", errors.New("network timeout")) {
t.Fatal("non-missing config error should not be swallowed")
}
}
-246
View File
@@ -1,246 +0,0 @@
package service
import (
"context"
"strings"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotCleanupRulesDefaultToEmpty(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("default cleanup rules should be empty, got %+v", cfg.AccountCleanupRules)
}
}
func TestBotCleanupRulesCanBeDeletedUntilEmpty(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
if _, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule add watch_hours watch_3_5d_6h 观看3到5天满6小时 3 5 6"); err != nil {
t.Fatal(err)
}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule del watch_3_5d_6h")
if err != nil {
t.Fatal(err)
}
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("cleanup rules should stay empty after deleting the last rule; reply=%q rules=%+v", reply.Text, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule list")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "暂无规则") {
t.Fatalf("expected empty rule list, got %q", reply.Text)
}
}
func TestBotCleanupRunPreviewsBeforeConfirm(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
newUser := &model.User{Username: "newbie", PasswordHash: "x", Role: "user", IsActive: true}
newUser.CreatedAt = now
for _, user := range []*model.User{stale, recent, newUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "any"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") || !strings.Contains(reply.Text, "stale") || !strings.Contains(reply.Text, "/cleanup run confirm") {
t.Fatalf("cleanup run should preview candidates and confirmation command, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("cleanup preview must not delete the stale user")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/deleted")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") {
t.Fatalf("/deleted alias should preview only, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("/deleted preview alias must not delete users")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已清理 <b>1</b>") {
t.Fatalf("cleanup confirm should delete exactly one stale user, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatal("stale user should be deleted after explicit confirmation")
}
for _, user := range []*model.User{recent, newUser, admin} {
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatalf("%s should be kept by保号 rules/protection", user.Username)
}
}
}
func TestBotCleanupLegacyCountModeStillKeepsSingleMatchedRule(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
for _, user := range []*model.User{recent, stale} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "count"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRequiredCount, "2"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if strings.Contains(reply.Text, "recent") {
t.Fatalf("user matching one keep rule must not be a cleanup candidate, got %q", reply.Text)
}
if !strings.Contains(reply.Text, "stale") {
t.Fatalf("user matching no keep rules should be a candidate, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if got, _ := repos.User.FindByID(ctx, recent.ID); got == nil {
t.Fatal("legacy count mode must not delete a user matching one keep rule")
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatalf("stale user should be deleted after confirm, reply=%q", reply.Text)
}
}
func TestBotCleanupConfirmRequiresEnabledRules(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
user := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
user.CreatedAt = time.Now().Add(-30 * 24 * time.Hour)
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "没有启用的保号规则") {
t.Fatalf("cleanup confirm without rules should be blocked, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatal("cleanup confirm without enabled rules must not delete users")
}
}
func TestBotCleanupRuleListInfersDaysAndHidesDuplicateNames(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"login_7d","type":"recent_login","enabled":true,"window_days_min":1,"window_days_max":5,"min_count":1},
{"id":"new_7d","name":"new_7d","type":"account_age_grace","enabled":true,"window_days_min":1,"window_days_max":1,"min_count":1}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule")
if err != nil {
t.Fatal(err)
}
for _, bad := range []string{"login_7d</code> · login_7d", "new_7d</code> · new_7d", "5 天内登录", "新号宽限 1 天", "add watch_hours", "Mgo 保号规则命令"} {
if strings.Contains(reply.Text, bad) {
t.Fatalf("rule list still contains bad fragment %q: %s", bad, reply.Text)
}
}
for _, want := range []string{"login_7d", "7 天内登录", "new_7d", "新号宽限 7 天"} {
if !strings.Contains(reply.Text, want) {
t.Fatalf("rule list missing %q: %s", want, reply.Text)
}
}
}
-175
View File
@@ -1,175 +0,0 @@
package service
import (
"context"
"strings"
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotRegistrationCommandUsesOpenRegQuota(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingOpenRegEnabled, "false"); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.TelegramBinding{
TelegramUserID: 9051,
TelegramName: "@root",
ChatID: 9051,
UserID: admin.ID,
}).Error; err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9051"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9051, Username: "root"}, Chat: TelegramChat{ID: 9051, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/registration on 2")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "2 个名额") {
t.Fatalf("expected quota feedback, got %q", reply.Text)
}
capacity := bot.loadCapacity(ctx)
if !capacity.OpenRegOn || capacity.OpenRegLimit != 2 || capacity.OpenRegUsed != 0 {
t.Fatalf("registration command should open quota-aware registration, got %+v", capacity)
}
}
func TestBotUserCommandsAndAdminGate(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
user := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.TelegramBinding{
TelegramUserID: 9101,
TelegramName: "@viewer",
ChatID: 9101,
UserID: user.ID,
}).Error; err != nil {
t.Fatal(err)
}
now := time.Now()
if err := repos.UserDevice.Create(ctx, &model.UserDevice{
UserID: user.ID, DeviceID: "dev-1", DeviceName: "iPhone", Client: "Infuse", FirstSeenAt: now, LastSeenAt: now,
}); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9101, Username: "viewer"}, Chat: TelegramChat{ID: 9101, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/antishare on")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "仅管理员") {
t.Fatalf("regular user should not manage policy, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/devices")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "我的登录设备") {
t.Fatalf("expected device list, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/kick 1")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已踢下线") {
t.Fatalf("expected kick feedback, got %q", reply.Text)
}
if kicked := bot.device; kicked != nil {
t.Fatal("test should not require wired device service")
}
if ok := NewDeviceService(zap.NewNop(), repos).IsDeviceKicked(ctx, user.ID, "dev-1"); !ok {
t.Fatal("device should be marked kicked")
}
}
func TestBotAdminCodeAndUserCommands(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
user := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.TelegramBinding{
TelegramUserID: 9301,
TelegramName: "@root",
ChatID: 9301,
UserID: admin.ID,
}).Error; err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "root"}, Chat: TelegramChat{ID: 9301, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/gencode renew 90 7")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已生成续期码") {
t.Fatalf("expected generated renew code, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/renew_user viewer 30")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "viewer") {
t.Fatalf("renew command should return user actions, got %q", reply.Text)
}
updated, _ := repos.User.FindByID(ctx, user.ID)
if updated.ExpiredAt == nil || updated.ExpiredAt.Before(time.Now()) {
t.Fatalf("renew_user should set future expiry, got %v", updated.ExpiredAt)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/delete_user viewer")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "需要确认") {
t.Fatalf("delete without confirm should be rejected, got %q", reply.Text)
}
}
func TestBotGroupMenuShowsAdminActionsOnlyForAdmins(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301","group_chat_id":"-1001"}`}
adminMsg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "admin"}, Chat: TelegramChat{ID: -1001, Type: "group"}}
reply, err := bot.executeCommand(ctx, channel, adminMsg, "/menu")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "管理员入口") || len(reply.Buttons) == 0 {
t.Fatalf("admin group menu should expose management actions, got %#v", reply)
}
userMsg := &TelegramMessage{From: TelegramUser{ID: 9302, Username: "user"}, Chat: TelegramChat{ID: -1001, Type: "group"}}
reply, err = bot.executeCommand(ctx, channel, userMsg, "/menu")
if err != nil {
t.Fatal(err)
}
if strings.Contains(reply.Text, "管理员入口") {
t.Fatalf("non-admin group menu must not expose management actions, got %#v", reply)
}
}
-253
View File
@@ -1,253 +0,0 @@
package service
import (
"context"
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestDeviceKickAndConcurrency(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
dev := NewDeviceService(zap.NewNop(), repos)
u := &model.User{Username: "carol", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
dev.RecordLogin(ctx, u.ID, "dev-1", "iPhone", "Infuse", "1.2.3.4")
dev.RecordPlayback(ctx, u.ID, "dev-1", "iPhone", "Infuse")
devices, _ := dev.ListDevices(ctx, u.ID)
if len(devices) != 1 {
t.Fatalf("expected 1 device, got %d", len(devices))
}
// 踢下线后命中 kicked
if err := dev.KickDevice(ctx, u.ID, "dev-1"); err != nil {
t.Fatal(err)
}
if !dev.IsDeviceKicked(ctx, u.ID, "dev-1") {
t.Fatal("device should be kicked")
}
// 重新登录清除 kicked
dev.RecordLogin(ctx, u.ID, "dev-1", "iPhone", "Infuse", "1.2.3.4")
if dev.IsDeviceKicked(ctx, u.ID, "dev-1") {
t.Fatal("re-login should clear kicked flag")
}
// 并发播放计数
now := time.Now()
for i, id := range []string{"d1", "d2", "d3", "d4"} {
_ = repos.UserDevice.Create(ctx, &model.UserDevice{
UserID: u.ID, DeviceID: id, FirstSeenAt: now, LastSeenAt: now, LastPlayAt: &now,
})
_ = i
}
n, err := repos.UserDevice.CountConcurrentPlaying(ctx, u.ID, now.Add(-time.Minute))
if err != nil {
t.Fatal(err)
}
if n < 4 {
t.Fatalf("expected >=4 concurrent playing, got %d", n)
}
}
func TestTerminalDeviceLimitDeduplicatesAppsOnSameDevice(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
dev := NewDeviceService(zap.NewNop(), repos)
now := time.Date(2026, 6, 25, 21, 30, 0, 0, time.UTC)
tracker := NewSessionTrackerService(zap.NewNop())
tracker.now = func() time.Time { return now }
dev.SetSessionTracker(tracker)
u := &model.User{Username: "device-user", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAntiShareEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingMaxLoggedClients, "3"); err != nil {
t.Fatal(err)
}
for _, login := range []struct {
id string
name string
client string
}{
{id: "phone-infuse", name: "iPhone", client: "Infuse"},
{id: "phone-emby", name: " iPhone ", client: "Emby"},
{id: "phone-jellyfin", name: "IPHONE", client: "Jellyfin"},
} {
dev.RecordLogin(ctx, u.ID, login.id, login.name, login.client, "1.2.3.4")
now = now.Add(time.Second)
}
count, err := repos.UserDevice.CountActiveClients(ctx, u.ID, now.Add(-24*time.Hour))
if err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("same terminal through multiple apps should count as 1, got %d", count)
}
devices, err := dev.ListDevices(ctx, u.ID)
if err != nil {
t.Fatal(err)
}
if len(devices) != 1 {
t.Fatalf("same terminal should show as one device row, got %#v", devices)
}
rawRows, err := repos.UserDevice.ListByUser(ctx, u.ID)
if err != nil {
t.Fatal(err)
}
if len(rawRows) != 1 {
t.Fatalf("same terminal should be persisted as one canonical row, got %#v", rawRows)
}
if devices[0].DeviceID != "phone-jellyfin" || devices[0].Client != "Jellyfin" {
t.Fatalf("merged device row should keep latest login channel, got %#v", devices[0])
}
got, _ := repos.User.FindByID(ctx, u.ID)
if !got.IsActive {
t.Fatal("same terminal through multiple apps must not disable the account")
}
dev.RecordLogin(ctx, u.ID, "tablet", "iPad", "Infuse", "1.2.3.4")
dev.RecordLogin(ctx, u.ID, "pc", "Windows PC", "Browser", "1.2.3.4")
count, err = repos.UserDevice.CountActiveClients(ctx, u.ID, now.Add(-24*time.Hour))
if err != nil {
t.Fatal(err)
}
if count != 3 {
t.Fatalf("three distinct terminal devices should count as 3, got %d", count)
}
got, _ = repos.User.FindByID(ctx, u.ID)
if !got.IsActive {
t.Fatal("device limit is inclusive; 3 of 3 terminals should stay active")
}
dev.RecordLogin(ctx, u.ID, "tv", "Apple TV", "Emby", "1.2.3.4")
got, _ = repos.User.FindByID(ctx, u.ID)
if got.IsActive {
t.Fatal("fourth distinct terminal should disable the account")
}
}
func TestDeviceKickAppliesToMergedTerminal(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
dev := NewDeviceService(zap.NewNop(), repos)
u := &model.User{Username: "kick-merged", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
dev.RecordLogin(ctx, u.ID, "phone-infuse", "iPhone", "Infuse", "1.2.3.4")
dev.RecordLogin(ctx, u.ID, "phone-emby", " iPhone ", "Emby", "1.2.3.4")
if err := dev.KickDevice(ctx, u.ID, "phone-emby"); err != nil {
t.Fatal(err)
}
if !dev.IsTerminalKicked(ctx, u.ID, "phone-jellyfin", "IPHONE", "Jellyfin") {
t.Fatal("same terminal with a new app/device id should still be kicked")
}
dev.RecordLogin(ctx, u.ID, "phone-jellyfin", "IPHONE", "Jellyfin", "1.2.3.4")
if dev.IsTerminalKicked(ctx, u.ID, "phone-jellyfin", "IPHONE", "Jellyfin") {
t.Fatal("re-login should clear kicked state for the merged terminal")
}
rawRows, err := repos.UserDevice.ListByUser(ctx, u.ID)
if err != nil {
t.Fatal(err)
}
if len(rawRows) != 1 || rawRows[0].DeviceID != "phone-jellyfin" {
t.Fatalf("merged terminal should keep one latest row, got %#v", rawRows)
}
}
func TestRecordPlaybackMergesChangingDeviceIDOnSameTerminal(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
dev := NewDeviceService(zap.NewNop(), repos)
u := &model.User{Username: "play-merged", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
dev.RecordPlayback(ctx, u.ID, "tv-emby", "Living Room TV", "Emby")
dev.RecordPlayback(ctx, u.ID, "tv-infuse", " living room tv ", "Infuse")
rawRows, err := repos.UserDevice.ListByUser(ctx, u.ID)
if err != nil {
t.Fatal(err)
}
if len(rawRows) != 1 {
t.Fatalf("same playback terminal should persist one row, got %#v", rawRows)
}
if rawRows[0].DeviceID != "tv-infuse" || rawRows[0].Client != "Infuse" || rawRows[0].LastPlayAt == nil {
t.Fatalf("merged playback row should keep latest playback channel, got %#v", rawRows[0])
}
count, err := repos.UserDevice.CountConcurrentPlaying(ctx, u.ID, time.Now().Add(-time.Minute))
if err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("same terminal playback should count once, got %d", count)
}
}
func TestConcurrentPlaybackDeduplicatesAppsOnSameDevice(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
u := &model.User{Username: "play-user", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
now := time.Now()
fp := fingerprint("Infuse", "Living Room TV")
for _, row := range []model.UserDevice{
{UserID: u.ID, DeviceID: "tv-emby", DeviceName: "Living Room TV", Client: "Emby", Fingerprint: fp, FirstSeenAt: now, LastSeenAt: now, LastPlayAt: &now},
{UserID: u.ID, DeviceID: "tv-jellyfin", DeviceName: "living room tv", Client: "Jellyfin", Fingerprint: fp, FirstSeenAt: now, LastSeenAt: now, LastPlayAt: &now},
{UserID: u.ID, DeviceID: "phone", DeviceName: "iPhone", Client: "Infuse", Fingerprint: fingerprint("Infuse", "iPhone"), FirstSeenAt: now, LastSeenAt: now, LastPlayAt: &now},
} {
if err := repos.UserDevice.Create(ctx, &row); err != nil {
t.Fatal(err)
}
}
count, err := repos.UserDevice.CountConcurrentPlaying(ctx, u.ID, now.Add(-time.Minute))
if err != nil {
t.Fatal(err)
}
if count != 2 {
t.Fatalf("same terminal playback through multiple apps should count as 1 terminal, got %d", count)
}
}
func TestProtectedAdminNeverViolated(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
dev := NewDeviceService(zap.NewNop(), repos)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
_ = repos.Setting.Set(ctx, SettingAntiShareEnabled, "true")
cfg := loadBotConfig(ctx, repos)
// 多次违规也不应删除/警告/禁用管理员
for i := 0; i < 5; i++ {
dev.registerFingerprintWarning(ctx, admin.ID, "test", cfg)
}
got, _ := repos.User.FindByID(ctx, admin.ID)
if got == nil {
t.Fatal("admin must never be auto-deleted")
}
if !got.IsActive {
t.Fatal("admin must never be auto-disabled")
}
if got.ShareWarnings != 0 {
t.Fatalf("admin should accrue no warnings, got %d", got.ShareWarnings)
}
}
-285
View File
@@ -1,285 +0,0 @@
package service
import (
"context"
"fmt"
"strconv"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// ── 容量 / 开注名额 ──────────────────────────────────────────────────────────
// capacityInfo 描述当前用户容量(随凭证授权实时变化)与开注名额状态。
type capacityInfo struct {
UsedUsers int64
MaxUsers int64 // 来自 LicensedMaxUsers,随授权实时变化
OpenRegOn bool
OpenRegLimit int // 0 = 不限(仅受 MaxUsers 约束)
OpenRegUsed int
}
// Remaining 返回还能注册多少个账号(同时受授权上限与开注名额约束)。
func (c capacityInfo) Remaining() int64 {
byLicense := c.MaxUsers - c.UsedUsers
if byLicense < 0 {
byLicense = 0
}
if c.OpenRegLimit > 0 {
byQuota := int64(c.OpenRegLimit - c.OpenRegUsed)
if byQuota < 0 {
byQuota = 0
}
if byQuota < byLicense {
return byQuota
}
}
return byLicense
}
// loadCapacity reads live capacity + open-reg quota state.
func (s *TelegramBotService) loadCapacity(ctx context.Context) capacityInfo {
used, _ := s.repo.User.Count(ctx)
info := capacityInfo{
UsedUsers: used,
MaxUsers: LicensedMaxUsers(ctx, s.repo),
OpenRegOn: s.openRegEnabled(ctx),
OpenRegLimit: s.intSetting(ctx, SettingOpenRegLimit, 0),
OpenRegUsed: s.intSetting(ctx, SettingOpenRegUsed, 0),
}
return info
}
func (s *TelegramBotService) intSetting(ctx context.Context, key string, fallback int) int {
v, err := s.repo.Setting.Get(ctx, key)
if err != nil {
return fallback
}
return parseIntSettingDefault(v, fallback)
}
// openRegEnabled reports whether bot registration is currently open. It honours
// both the new open-reg switch and the legacy registration switch.
func (s *TelegramBotService) openRegEnabled(ctx context.Context) bool {
if v, _ := s.repo.Setting.Get(ctx, SettingOpenRegEnabled); v != "" {
return parseBoolSetting(v, false)
}
return s.registrationEnabled(ctx)
}
// openRegistration opens registration for `limit` new accounts (0 = unlimited,
// bounded only by the license). Resets the used counter.
func (s *TelegramBotService) openRegistration(ctx context.Context, limit int) error {
if limit < 0 {
limit = 0
}
if err := s.repo.Setting.Set(ctx, SettingOpenRegEnabled, "true"); err != nil {
return err
}
if err := s.repo.Setting.Set(ctx, SettingOpenRegLimit, strconv.Itoa(limit)); err != nil {
return err
}
if err := s.repo.Setting.Set(ctx, SettingOpenRegUsed, "0"); err != nil {
return err
}
// 与旧开关同步,兼容系统设置页。
return s.setRegistrationEnabled(ctx, true)
}
// closeRegistration disables bot registration.
func (s *TelegramBotService) closeRegistration(ctx context.Context) error {
if err := s.repo.Setting.Set(ctx, SettingOpenRegEnabled, "false"); err != nil {
return err
}
return s.setRegistrationEnabled(ctx, false)
}
// consumeOpenRegSlot increments the used counter and auto-closes registration
// once the quota is exhausted. Call after a successful bot registration.
func (s *TelegramBotService) consumeOpenRegSlot(ctx context.Context) {
limit := s.intSetting(ctx, SettingOpenRegLimit, 0)
used := s.intSetting(ctx, SettingOpenRegUsed, 0) + 1
_ = s.repo.Setting.Set(ctx, SettingOpenRegUsed, strconv.Itoa(used))
if limit > 0 && used >= limit {
_ = s.closeRegistration(ctx)
}
}
// ── 兑换码 ──────────────────────────────────────────────────────────────────
// generateCode creates a random redemption code of the given kind. durationDays
// sets the account validity granted on redeem (0 = permanent). validDays sets
// how long the code itself stays redeemable (0 = never expires).
func (s *TelegramBotService) generateCode(ctx context.Context, kind string, durationDays, validDays int, createdBy string) (*model.RegistrationCode, error) {
return s.generateCodeWithUses(ctx, kind, durationDays, validDays, 1, createdBy)
}
func (s *TelegramBotService) generateCodeWithUses(ctx context.Context, kind string, durationDays, validDays, maxUses int, createdBy string) (*model.RegistrationCode, error) {
if kind != model.RegistrationCodeRegister && kind != model.RegistrationCodeRenew {
kind = model.RegistrationCodeRegister
}
if maxUses <= 0 {
maxUses = 1
}
code := &model.RegistrationCode{
Code: randomCode(12),
Kind: kind,
DurationDays: durationDays,
MaxUses: maxUses,
CreatedByID: createdBy,
}
if validDays > 0 {
exp := time.Now().Add(time.Duration(validDays) * 24 * time.Hour)
code.ExpiresAt = &exp
}
if err := s.repo.RegCode.Create(ctx, code); err != nil {
return nil, err
}
return code, nil
}
// lookupRedeemableCode validates a code without consuming it. Callers mark it
// used only after the dependent action (account create / renew) succeeds, so a
// failed action never burns a code.
func (s *TelegramBotService) lookupRedeemableCode(ctx context.Context, raw, wantKind string) (*model.RegistrationCode, string) {
code := normalizeRedemptionCode(raw)
if code == "" {
return nil, "请提供兑换码。"
}
rc, err := s.repo.RegCode.FindByCode(ctx, code)
if err != nil || rc == nil {
return nil, "兑换码无效。"
}
if rc.IsUsed() {
return nil, "兑换码已被使用。"
}
if rc.IsExpired() {
return nil, "兑换码已过期。"
}
if wantKind != "" && rc.Kind != wantKind {
switch rc.Kind {
case model.RegistrationCodeRenew:
return nil, "这是续期兑换码,请在「我的账号」里使用它续期。"
default:
return nil, "这是注册兑换码,请用于注册新账号。"
}
}
return rc, ""
}
func normalizeRedemptionCode(raw string) string {
code := strings.ToUpper(strings.TrimSpace(raw))
code = strings.NewReplacer(" ", "", "-", "", "_", "").Replace(code)
return code
}
func looksLikeRedemptionCode(raw string) bool {
code := normalizeRedemptionCode(raw)
if len(code) < 8 || len(code) > 32 {
return false
}
for _, ch := range code {
if !strings.ContainsRune(codeAlphabet, ch) {
return false
}
}
return true
}
// ── 续期 ────────────────────────────────────────────────────────────────────
// renewUser extends a user's expiry by durationDays. A nil/zero current expiry
// starts from now; a future expiry is extended from that point. durationDays<=0
// sets the account to never expire (permanent).
func renewExpiry(current *time.Time, durationDays int) *time.Time {
if durationDays <= 0 {
return nil // permanent
}
base := time.Now()
if current != nil && current.After(base) {
base = *current
}
exp := base.Add(time.Duration(durationDays) * 24 * time.Hour)
return &exp
}
// applyRenewal renews a user account and clears any expiry-related suspension.
func (s *TelegramBotService) applyRenewal(ctx context.Context, userID string, durationDays int) error {
u, err := s.repo.User.FindByID(ctx, userID)
if err != nil || u == nil {
return fmt.Errorf("user not found")
}
exp := renewExpiry(u.ExpiredAt, durationDays)
updates := map[string]any{"expired_at": exp, "is_active": true}
return s.repo.User.UpdateFields(ctx, userID, updates)
}
// ── 签到 ────────────────────────────────────────────────────────────────────
// signInResult 描述一次签到的结果。
type signInResult struct {
AlreadySigned bool
Streak int
Total int
}
// signIn records a daily sign-in for the user, tracking consecutive-day streaks
// only (no points). A second sign-in on the same calendar day is a no-op.
func (s *TelegramBotService) signIn(ctx context.Context, userID string) (signInResult, error) {
now := time.Now()
today := now.Truncate(24 * time.Hour)
rec, err := s.repo.SignIn.Get(ctx, userID)
if err != nil {
return signInResult{}, err
}
if rec == nil {
rec = &model.SignIn{UserID: userID, LastSignIn: now, StreakDays: 1, TotalDays: 1}
if err := s.repo.SignIn.Save(ctx, rec); err != nil {
return signInResult{}, err
}
return signInResult{Streak: 1, Total: 1}, nil
}
last := rec.LastSignIn.Truncate(24 * time.Hour)
switch {
case last.Equal(today):
return signInResult{AlreadySigned: true, Streak: rec.StreakDays, Total: rec.TotalDays}, nil
case last.Equal(today.Add(-24 * time.Hour)):
rec.StreakDays++
default:
rec.StreakDays = 1 // streak broken
}
rec.TotalDays++
rec.LastSignIn = now
if err := s.repo.SignIn.Save(ctx, rec); err != nil {
return signInResult{}, err
}
return signInResult{Streak: rec.StreakDays, Total: rec.TotalDays}, nil
}
// ── helpers ─────────────────────────────────────────────────────────────────
const codeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" // no ambiguous 0/O/1/I
func randomCode(n int) string {
b := make([]byte, n)
secureRandomBytes(b)
out := make([]byte, n)
for i := range b {
out[i] = codeAlphabet[int(b[i])%len(codeAlphabet)]
}
return string(out)
}
// formatExpiry renders a user's expiry status for display.
func formatExpiry(t *time.Time) string {
if t == nil {
return "永久有效"
}
if time.Now().After(*t) {
return "已过期(" + t.Format("2006-01-02") + ")"
}
days := int(time.Until(*t).Hours() / 24)
return fmt.Sprintf("%s(剩 %d 天)", t.Format("2006-01-02"), days)
}
-135
View File
@@ -1,135 +0,0 @@
package service
import (
"context"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"go.uber.org/zap"
)
func newBotTestService(t *testing.T) (*repository.Container, *TelegramBotService) {
t.Helper()
db := newServiceTestDB(t, model.AllModels()...)
repos := repository.New(db)
cfg := &config.Config{}
cfg.Secrets.JWTSecret = "test-secret"
log := zap.NewNop()
perms := NewPermissionService(log, repos)
tokenSvc := NewTokenService(cfg, log, repos)
auth := NewAuthService(cfg, log, repos, tokenSvc, perms)
crypto := NewCryptoService("test-secret", log)
bot := NewTelegramBotService(log, repos, crypto, auth)
return repos, bot
}
// ── pure logic ──────────────────────────────────────────────────────────────
func TestRenewExpiry(t *testing.T) {
// 永久(0 天)→ nil
if got := renewExpiry(nil, 0); got != nil {
t.Fatalf("expected nil for permanent, got %v", got)
}
// 从现在起 +30 天(当前为空)
got := renewExpiry(nil, 30)
if got == nil || got.Before(time.Now().Add(29*24*time.Hour)) {
t.Fatalf("expected ~30d expiry, got %v", got)
}
// 已有未来到期 → 在原到期基础上叠加
future := time.Now().Add(10 * 24 * time.Hour)
got = renewExpiry(&future, 30)
if got == nil || got.Before(future.Add(29*24*time.Hour)) {
t.Fatalf("expected stacking on future expiry, got %v", got)
}
// 已过期 → 从现在起算
past := time.Now().Add(-10 * 24 * time.Hour)
got = renewExpiry(&past, 5)
if got == nil || got.Before(time.Now().Add(4*24*time.Hour)) {
t.Fatalf("expected fresh window from now, got %v", got)
}
}
func TestCapacityRemaining(t *testing.T) {
cases := []struct {
name string
c capacityInfo
want int64
}{
{"license only", capacityInfo{UsedUsers: 5, MaxUsers: 20}, 15},
{"quota tighter", capacityInfo{UsedUsers: 5, MaxUsers: 100, OpenRegLimit: 10, OpenRegUsed: 3}, 7},
{"license tighter", capacityInfo{UsedUsers: 95, MaxUsers: 100, OpenRegLimit: 50, OpenRegUsed: 0}, 5},
{"full", capacityInfo{UsedUsers: 20, MaxUsers: 20}, 0},
{"quota exhausted", capacityInfo{UsedUsers: 1, MaxUsers: 100, OpenRegLimit: 5, OpenRegUsed: 5}, 0},
}
for _, tc := range cases {
if got := tc.c.Remaining(); got != tc.want {
t.Errorf("%s: Remaining()=%d want %d", tc.name, got, tc.want)
}
}
}
func TestRandomWindowDays(t *testing.T) {
for i := 0; i < 200; i++ {
d := randomWindowDays(3, 5)
if d < 3 || d > 5 {
t.Fatalf("randomWindowDays(3,5)=%d out of range", d)
}
}
if d := randomWindowDays(4, 4); d != 4 {
t.Fatalf("randomWindowDays(4,4)=%d want 4", d)
}
}
func TestFingerprintStability(t *testing.T) {
a := fingerprint("Infuse", "iPhone")
b := fingerprint("infuse", " iPhone ")
if a != b {
t.Fatalf("fingerprint should be case/space-insensitive: %s != %s", a, b)
}
if a != fingerprint("Emby", "iPhone") {
t.Fatal("different apps on the same terminal must share one fingerprint")
}
if a == fingerprint("Infuse", "iPad") {
t.Fatal("different device names must yield different fingerprints")
}
}
// ── DB-backed ─────────────────────────────────────────────────────────────
func TestSignInStreak(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
u := &model.User{Username: "alice", PasswordHash: "x", Role: "user"}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
res, err := bot.signIn(ctx, u.ID)
if err != nil || res.Streak != 1 || res.Total != 1 {
t.Fatalf("first sign-in: %+v err=%v", res, err)
}
// 同日重复签到 → 不增长
res, _ = bot.signIn(ctx, u.ID)
if !res.AlreadySigned || res.Streak != 1 {
t.Fatalf("same-day re-signin should be no-op: %+v", res)
}
// 模拟昨天签到 → 连续 +1
rec, _ := repos.SignIn.Get(ctx, u.ID)
rec.LastSignIn = time.Now().Add(-24 * time.Hour)
_ = repos.SignIn.Save(ctx, rec)
res, _ = bot.signIn(ctx, u.ID)
if res.Streak != 2 || res.Total != 2 {
t.Fatalf("consecutive day should bump streak: %+v", res)
}
// 中断(前天)→ 重置为 1
rec, _ = repos.SignIn.Get(ctx, u.ID)
rec.LastSignIn = time.Now().Add(-72 * time.Hour)
_ = repos.SignIn.Save(ctx, rec)
res, _ = bot.signIn(ctx, u.ID)
if res.Streak != 1 {
t.Fatalf("broken streak should reset to 1: %+v", res)
}
}
@@ -1,124 +0,0 @@
package service
import (
"context"
"strings"
"testing"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotAdminCommandsManageDevicePolicy(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.TelegramBinding{
TelegramUserID: 9001,
TelegramName: "@root",
ChatID: 9001,
UserID: admin.ID,
}).Error; err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/antishare on play=4 login=5 warn=3")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "防共享:<b>已开启</b>") {
t.Fatalf("expected antishare enabled reply, got %q", reply.Text)
}
cfg := loadBotConfig(ctx, repos)
if !cfg.AntiShareEnabled || cfg.MaxConcurrentPlay != 4 || cfg.MaxLoggedClients != 5 || cfg.WarnThreshold != 3 {
t.Fatalf("unexpected device policy: %+v", cfg)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_mode count 2")
if err != nil {
t.Fatal(err)
}
cfg = loadBotConfig(ctx, repos)
if cfg.AccountCleanupKeepMode != "any" || cfg.AccountCleanupRequiredCount != 1 {
t.Fatalf("unexpected cleanup mode: %+v; reply=%q", cfg, reply.Text)
}
if !strings.Contains(reply.Text, "满足任意一条") {
t.Fatalf("cleanup mode should explain fixed any-rule policy, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule add recent_login login_7d 七天内登录 7")
if err != nil {
t.Fatal(err)
}
cfg = loadBotConfig(ctx, repos)
found := false
for _, rule := range cfg.AccountCleanupRules {
if rule.ID == "login_7d" && rule.Type == "recent_login" && rule.WindowDaysMax == 7 {
found = true
}
}
if !found {
t.Fatalf("cleanup rule not added; reply=%q rules=%+v", reply.Text, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule edit recent_login login_7d 十四天内登录 14")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已更新规则") {
t.Fatalf("expected cleanup rule update reply, got %q", reply.Text)
}
cfg = loadBotConfig(ctx, repos)
matches := 0
for _, rule := range cfg.AccountCleanupRules {
if rule.ID == "login_7d" {
matches++
if rule.Type != "recent_login" || rule.WindowDaysMax != 14 || rule.Name != "十四天内登录" {
t.Fatalf("cleanup rule should be updated in place, got %+v", rule)
}
}
}
if matches != 1 {
t.Fatalf("cleanup rule update should not create duplicates, got %d rules=%+v", matches, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule 修改 recent_login login_7d 二十一天内登录 21")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已更新规则") {
t.Fatalf("expected Chinese cleanup rule update reply, got %q", reply.Text)
}
cfg = loadBotConfig(ctx, repos)
matches = 0
for _, rule := range cfg.AccountCleanupRules {
if rule.ID == "login_7d" {
matches++
if rule.WindowDaysMax != 21 || rule.Name != "二十一天内登录" {
t.Fatalf("Chinese cleanup rule update should update values, got %+v", rule)
}
}
}
if matches != 1 {
t.Fatalf("Chinese cleanup rule update should not create duplicates, got %d rules=%+v", matches, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule add account_age_grace new_7d 7")
if err != nil {
t.Fatal(err)
}
cfg = loadBotConfig(ctx, repos)
found = false
for _, rule := range cfg.AccountCleanupRules {
if rule.ID == "new_7d" && rule.Type == "account_age_grace" && rule.MinCount == 7 {
found = true
}
}
if !found {
t.Fatalf("cleanup shorthand rule not added; reply=%q rules=%+v", reply.Text, cfg.AccountCleanupRules)
}
}
@@ -1,141 +0,0 @@
package service
import (
"context"
"encoding/json"
"strings"
"testing"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotRedeemRegisterRequiresAllowedTelegramUser(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "outsider"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "不在管理员配置") {
t.Fatalf("outsider should not redeem register code, got %q", reply.Text)
}
channel.Config = `{"admin_user_ids":"9201"}`
reply, err = bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("allowed user should redeem register code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9201); binding == nil {
t.Fatal("redeemed account should be bound to telegram user")
}
}
func TestBotRedeemRegisterCodeCreatesOnlyOneAccount(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9201,9202"}`}
first := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "first"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, first, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("first redeem should succeed, got %q", reply.Text)
}
second := &TelegramMessage{From: TelegramUser{ID: 9202, Username: "second"}, Chat: TelegramChat{ID: 9202, Type: "private"}}
reply, err = bot.executeCommand(ctx, channel, second, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换码已被使用") && !strings.Contains(reply.Text, "兑换码刚刚被使用") {
t.Fatalf("second redeem should be rejected as used, got %q", reply.Text)
}
var users int64
if err := repos.DB.Model(&model.User{}).Count(&users).Error; err != nil {
t.Fatal(err)
}
if users != 1 {
t.Fatalf("one register code must create exactly one user, got %d", users)
}
if binding := bot.telegramBinding(ctx, 9202); binding != nil {
t.Fatal("second telegram user must not be bound by an already-used register code")
}
}
func TestBotRegisterCommandAcceptsRegistrationCode(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "codeuser"}, Chat: TelegramChat{ID: 9301, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/register "+strings.ToLower(code.Code[:4])+"-"+strings.ToLower(code.Code[4:]))
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("/register CODE should redeem registration code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9301); binding == nil {
t.Fatal("register code should bind the newly created account")
}
}
func TestBotPlainRegistrationCodeMessageRedeems(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.NotifyChannel{
Name: "Telegram",
Type: "telegram",
Enabled: true,
Config: `{"admin_user_ids":"9302"}`,
}).Error; err != nil {
t.Fatal(err)
}
update, _ := json.Marshal(TelegramUpdate{
UpdateID: 1,
Message: &TelegramMessage{
MessageID: 12,
Text: strings.ToLower(code.Code),
From: TelegramUser{ID: 9302, Username: "plaincode"},
Chat: TelegramChat{ID: 9302, Type: "private"},
},
})
if err := bot.HandleWebhook(ctx, update); err != nil {
t.Fatal(err)
}
if binding := bot.telegramBinding(ctx, 9302); binding == nil {
t.Fatal("plain code private message should redeem and bind account")
}
var used model.RegistrationCode
if err := repos.DB.Where("code = ?", code.Code).First(&used).Error; err != nil {
t.Fatal(err)
}
if used.UsedAt == nil || used.UsedByUserID == "" {
t.Fatal("plain code message should mark registration code as used")
}
}
-94
View File
@@ -1,94 +0,0 @@
package service
import (
"context"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestRegistrationCodeRedeemOnce(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRenew, 30, 0, "")
if err != nil {
t.Fatal(err)
}
// 首次校验通过
rc, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew)
if rc == nil {
t.Fatalf("expected valid code, got msg=%q", msg)
}
// 标记使用后不可再用
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-1"); err != nil {
t.Fatal(err)
}
if _, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew); msg == "" {
t.Fatal("used code must not validate again")
}
// 第二次 MarkUsed 应失败(防止双花)
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-2"); err == nil {
t.Fatal("double-spend should be rejected")
}
// 类型不匹配应被拒
reg, _ := bot.generateCode(ctx, model.RegistrationCodeRegister, 0, 0, "")
if _, msg := bot.lookupRedeemableCode(ctx, reg.Code, model.RegistrationCodeRenew); msg == "" {
t.Fatal("register code should not validate as renew")
}
}
func TestRegistrationCodeCanBeGeneratedForMultipleUses(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCodeWithUses(ctx, model.RegistrationCodeRenew, 30, 0, 2, "")
if err != nil {
t.Fatal(err)
}
rc, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew)
if rc == nil {
t.Fatalf("expected valid code, got msg=%q", msg)
}
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-1"); err != nil {
t.Fatal(err)
}
rc, msg = bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew)
if rc == nil {
t.Fatalf("code should remain redeemable after first use, got msg=%q", msg)
}
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-2"); err != nil {
t.Fatal(err)
}
if _, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew); msg == "" {
t.Fatal("code should be exhausted after max uses")
}
var used model.RegistrationCode
if err := repos.DB.Where("id = ?", code.ID).First(&used).Error; err != nil {
t.Fatal(err)
}
if used.UsedCount != 2 || used.UsedAt == nil {
t.Fatalf("expected exhausted code with used_count=2, got %+v", used)
}
}
func TestRenewalClearsExpiry(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
past := time.Now().Add(-time.Hour)
u := &model.User{Username: "bob", PasswordHash: "x", Role: "user", IsActive: false, ExpiredAt: &past}
if err := repos.User.Create(ctx, u); err != nil {
t.Fatal(err)
}
if err := bot.applyRenewal(ctx, u.ID, 30); err != nil {
t.Fatal(err)
}
got, _ := repos.User.FindByID(ctx, u.ID)
if !got.IsActive {
t.Fatal("renewal should re-activate account")
}
if got.ExpiredAt == nil || got.ExpiredAt.Before(time.Now()) {
t.Fatalf("renewal should set future expiry, got %v", got.ExpiredAt)
}
}
-141
View File
@@ -1,141 +0,0 @@
package service
import (
"context"
"strings"
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotAdminUnbindMultipleUsers(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
viewer := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
guest := &model.User{Username: "guest", PasswordHash: "x", Role: "user", IsActive: true}
for _, user := range []*model.User{admin, viewer, guest} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
bindings := []model.TelegramBinding{
{TelegramUserID: 9401, TelegramName: "@root", ChatID: 9401, UserID: admin.ID},
{TelegramUserID: 9402, TelegramName: "@viewer", ChatID: 9402, UserID: viewer.ID},
{TelegramUserID: 9403, TelegramName: "@guest", ChatID: 9403, UserID: guest.ID},
}
for i := range bindings {
if err := repos.DB.Create(&bindings[i]).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9401"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9401, Username: "root"}, Chat: TelegramChat{ID: 9401, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind viewer,guest missing root")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>2</b>") || !strings.Contains(reply.Text, "root(管理员)") || !strings.Contains(reply.Text, "missing") {
t.Fatalf("unexpected unbind reply: %q", reply.Text)
}
for _, user := range []*model.User{viewer, guest} {
var count int64
if err := repos.DB.Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("%s binding count = %d, want 0", user.Username, count)
}
}
if binding := bot.telegramBinding(ctx, 9401); binding == nil {
t.Fatal("admin binding should be protected from /unbind by username")
}
}
func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
oldTime := time.Now().Add(-45 * 24 * time.Hour)
recentTime := time.Now().Add(-2 * 24 * time.Hour)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true, LastLoginAt: &oldTime}
oldUser := &model.User{Username: "old", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime}
realtimeUser := &model.User{Username: "realtime", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime}
recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime}
for _, user := range []*model.User{admin, oldUser, realtimeUser, recentUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
for _, binding := range []model.TelegramBinding{
{TelegramUserID: 9501, TelegramName: "@root", ChatID: 9501, UserID: admin.ID},
{TelegramUserID: 9502, TelegramName: "@old", ChatID: 9502, UserID: oldUser.ID},
{TelegramUserID: 9503, TelegramName: "@recent", ChatID: 9503, UserID: recentUser.ID},
{TelegramUserID: 9505, TelegramName: "@realtime", ChatID: 9505, UserID: realtimeUser.ID},
{TelegramUserID: 9504, TelegramName: "@ghost", ChatID: 9504, UserID: "missing-user"},
} {
row := binding
if err := repos.DB.Create(&row).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9501"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9501, Username: "root"}, Chat: TelegramChat{ID: 9501, Type: "private"}}
tracker := NewSessionTrackerService(zap.NewNop())
tracker.RecordActivity(ctx, realtimeUser.ID, realtimeUser.Username, "phone-1", "iPhone", "Infuse", "192.0.2.10")
device := NewDeviceService(zap.NewNop(), repos)
device.SetSessionTracker(tracker)
bot.SetDeviceService(device)
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "old") {
t.Fatalf("unexpected inactive unbind reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9502); binding != nil {
t.Fatal("old user binding should be removed")
}
if binding := bot.telegramBinding(ctx, 9501); binding == nil {
t.Fatal("admin binding should be skipped by inactive cleanup")
}
if binding := bot.telegramBinding(ctx, 9503); binding == nil {
t.Fatal("recent user binding should remain")
}
if binding := bot.telegramBinding(ctx, 9505); binding == nil {
t.Fatal("realtime active user binding should remain")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "tg:9504") {
t.Fatalf("unexpected duplicate cleanup reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9504); binding != nil {
t.Fatal("invalid binding should be removed")
}
}
func TestTelegramMembershipChatIDsIncludesCommandChatID(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 1 || got[0] != "-100123" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want command_chat_id", got)
}
}
func TestTelegramMembershipChatIDsDedupesGroupChannelAndCommandIDs(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"group_chat_id":"-100123","channel_chat_id":"-100124","command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 2 || got[0] != "-100123" || got[1] != "-100124" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want deduped ids", got)
}
}
-137
View File
@@ -1,137 +0,0 @@
// Package cloud implements pluggable cloud-disk (网盘) providers used by the
// external-storage subsystem to expose remote files as playable media via
// HTTP 302 redirects.
//
// The design offloads playback to the cloud provider: instead of the
// host downloading and re-streaming bytes, a provider resolves a file to a
// short-lived direct download URL and the player is 302-redirected straight to
// the cloud CDN. The host only performs a tiny redirect, freeing its CPU and
// bandwidth.
//
// Each provider authenticates with a cookie (obtained via the web UI, an API
// cookie, or a QR-code login flow). Providers are intentionally side-effect
// free and take an *http.Client so they can be exercised against httptest
// mock servers in unit tests.
package cloud
import (
"context"
"errors"
"net/http"
"strings"
"time"
)
// timeNow is a seam so tests can pin timestamps.
var timeNow = time.Now
// Provider types recognised by the registry.
const (
Type115 = "cloud115" // 115 网盘
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
)
// ErrUnsupported is returned for an unknown provider type.
var ErrUnsupported = errors.New("unsupported cloud provider")
// FileEntry is one item in a cloud directory listing.
type FileEntry struct {
ID string `json:"id"` // provider-native file id
Name string `json:"name"`
IsDir bool `json:"is_dir"`
Size int64 `json:"size"`
// PickCode is 115-specific; other providers use ID directly.
PickCode string `json:"pick_code,omitempty"`
}
// DirectLink is a resolved playback target.
type DirectLink struct {
URL string `json:"url"`
// Headers that must accompany a request to URL (e.g. User-Agent, Cookie).
Headers map[string]string `json:"-"`
// Proxy reports whether URL requires the host to reverse-proxy the bytes
// (because the headers cannot be carried by a plain browser 302). When
// false the play handler issues a pure 302 redirect (true offload).
Proxy bool `json:"-"`
}
// Provider is the common cloud-disk interface.
type Provider interface {
// Type returns the provider key.
Type() string
// Ping validates the stored credentials (cookie). Cheap, used by the
// storage-config Test() probe.
Ping(ctx context.Context) error
// List returns the entries under dirID. An empty dirID means the root.
List(ctx context.Context, dirID string) ([]FileEntry, error)
// Resolve turns a provider-native file reference (id or pickcode) into a
// short-lived direct download link suitable for 302 playback.
Resolve(ctx context.Context, fileRef string) (*DirectLink, error)
}
// MutableProvider is implemented by cloud bridges that support safe folder
// management through their official API or standard WebDAV methods.
type MutableProvider interface {
Provider
Mkdir(ctx context.Context, parentDir, name string) (*FileEntry, error)
Rename(ctx context.Context, ref, name string) (*FileEntry, error)
}
// MovableProvider is implemented by writable cloud bridges that can move an
// entry across directories, optionally renaming it in the same operation.
type MovableProvider interface {
MutableProvider
Move(ctx context.Context, ref, targetDir, name string) (*FileEntry, error)
}
// New constructs a provider of the given type from a free-form config map
// (as persisted by StorageConfigService). The client is shared so callers can
// inject timeouts / test transports.
func New(typ string, cfg map[string]any, client *http.Client) (Provider, error) {
if client == nil {
client = http.DefaultClient
}
switch typ {
case Type115:
return new115(cfg, client), nil
case TypeCloudDrive2:
return newCloudDrive2(cfg, client), nil
case TypeOpenList:
return newOpenList(cfg, client), nil
default:
return nil, ErrUnsupported
}
}
// IsCloudType reports whether typ is a cloud-disk provider.
func IsCloudType(typ string) bool {
return typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
}
// str coerces a config value to a trimmed string.
func str(v any) string {
if v == nil {
return ""
}
if s, ok := v.(string); ok {
return strings.TrimSpace(s)
}
return ""
}
// boolish coerces a config value to bool ("true"/"1"/true → true).
func boolish(v any) bool {
switch t := v.(type) {
case bool:
return t
case string:
s := strings.ToLower(strings.TrimSpace(t))
return s == "1" || s == "true" || s == "yes" || s == "on"
default:
return false
}
}
// defaultUA is a desktop browser UA accepted by upstream cloud providers.
const defaultUA = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0 Safari/537.36"
-212
View File
@@ -1,212 +0,0 @@
package cloud
import (
"context"
"encoding/base64"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
)
func Test115ListAndResolve(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/files":
if r.URL.Query().Get("cid") != "0" {
t.Errorf("bad cid %q", r.URL.Query().Get("cid"))
}
w.Write([]byte(`{"state":true,"data":[
{"cid":"100","n":"Movies","s":0},
{"fid":"200","n":"Inception.mkv","s":456,"pc":"pick200"}]}`))
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
// The downurl endpoint is m115-encrypted end-to-end (the server side
// requires 115's private key), so stub the decrypted payload via the seam
// and assert the pickcode->URL extraction. The live crypto/transport path is
// exercised by integration testing against the real 115 API.
p115, ok := p.(*pan115Provider)
if !ok {
t.Fatalf("expected *pan115Provider, got %T", p)
}
p115.downURLPayload = func(ctx context.Context, pickcode string) ([]byte, error) {
if pickcode != "pick200" {
t.Errorf("bad pickcode %q", pickcode)
}
return []byte(`{"200":{"file_name":"Inception.mkv","file_size":"456","url":{"url":"https://cdn.115/x.mkv?t=1"}}}`), nil
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 2 {
t.Fatalf("want 2 entries: %#v", entries)
}
if !entries[0].IsDir || entries[0].ID != "100" {
t.Fatalf("dir entry wrong: %#v", entries[0])
}
if entries[1].IsDir || entries[1].PickCode != "pick200" || entries[1].Size != 456 {
t.Fatalf("file entry wrong: %#v", entries[1])
}
link, err := p.Resolve(context.Background(), "pick200")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if link.URL != "https://cdn.115/x.mkv?t=1" {
t.Fatalf("bad url: %s", link.URL)
}
if link.Proxy {
t.Fatalf("115 should default to 302 (no proxy)")
}
}
func Test115ListPaginates(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/files" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
count := 100
if offset > 0 {
count = 1
}
items := make([]string, 0, count)
for i := 0; i < count; i++ {
n := offset + i
items = append(items, fmt.Sprintf(`{"fid":"%d","n":"Movie.%03d.mkv","s":%d,"pc":"pick%d"}`, n, n, n, n))
}
w.Write([]byte(`{"state":true,"data":[` + strings.Join(items, ",") + `]}`))
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "0")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 101 {
t.Fatalf("entries = %d, want 101", len(entries))
}
if entries[100].ID != "100" || entries[100].PickCode != "pick100" {
t.Fatalf("last entry wrong: %#v", entries[100])
}
}
// Test115DownURLEndpointAndError exercises the live fetchDownURLPayload path:
// it must POST an m115-encrypted `data` body to /app/chrome/downurl?t=... and
// surface 115's error when state=false (no decryption needed for that branch).
func Test115DownURLEndpointAndError(t *testing.T) {
var gotData, gotT string
pro := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/chrome/downurl" {
t.Errorf("unexpected path %s", r.URL.Path)
}
gotT = r.URL.Query().Get("t")
_ = r.ParseForm()
gotData = r.PostFormValue("data")
w.Write([]byte(`{"state":false,"error":"not exist"}`))
}))
defer pro.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1", "pro_base": pro.URL}, pro.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "pickX")
if err == nil || !strings.Contains(err.Error(), "not exist") {
t.Fatalf("want upstream error surfaced, got %v", err)
}
if gotT == "" {
t.Errorf("missing t query param")
}
if gotData == "" {
t.Errorf("missing encrypted data body")
}
if _, derr := base64.StdEncoding.DecodeString(gotData); derr != nil {
t.Errorf("data body is not base64: %v", derr)
}
}
func Test115QRFlow(t *testing.T) {
// status sequence: waiting -> scanned -> confirmed
calls := 0
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/1.0/web/1.0/token/":
w.Write([]byte(`{"state":1,"data":{"uid":"U1","time":1700,"sign":"S1"}}`))
case "/get/status/":
if r.URL.Query().Get("uid") != "U1" {
t.Errorf("bad uid %q", r.URL.Query().Get("uid"))
}
calls++
switch calls {
case 1:
w.Write([]byte(`{"state":1,"data":{"status":0}}`))
case 2:
w.Write([]byte(`{"state":1,"data":{"status":1}}`))
default:
w.Write([]byte(`{"state":1,"data":{"status":2}}`))
}
default:
t.Errorf("unexpected api path %s", r.URL.Path)
}
}))
defer api.Close()
passport := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/1.0/web/1.0/login/qrcode/" {
t.Errorf("unexpected passport path %s", r.URL.Path)
}
w.Write([]byte(`{"state":1,"data":{"cookie":{"UID":"u","CID":"c","SEID":"s"}}}`))
}))
defer passport.Close()
oldA, oldP := qr115APIBase, qr115PassportBase
qr115APIBase, qr115PassportBase = api.URL, passport.URL
defer func() { qr115APIBase, qr115PassportBase = oldA, oldP }()
ctx := context.Background()
sess, err := QRStart(ctx, api.Client())
if err != nil {
t.Fatalf("qr start: %v", err)
}
if sess.UID != "U1" || sess.QRImageURL == "" {
t.Fatalf("bad session: %#v", sess)
}
want := []string{"waiting", "scanned", "confirmed"}
for i, exp := range want {
st, err := QRPoll(ctx, api.Client(), sess)
if err != nil {
t.Fatalf("poll %d: %v", i, err)
}
if st.State != exp {
t.Fatalf("poll %d: want %s got %s", i, exp, st.State)
}
if exp == "confirmed" {
if st.Cookie == "" || !containsAll(st.Cookie, "UID=u", "SEID=s") {
t.Fatalf("confirmed must yield cookie: %q", st.Cookie)
}
}
}
}
func containsAll(s string, subs ...string) bool {
for _, sub := range subs {
if !strings.Contains(s, sub) {
return false
}
}
return true
}
-29
View File
@@ -1,29 +0,0 @@
package cloud
import (
"net/http"
"testing"
)
func TestDeprecatedProviderPlaybackOverrideKeysAreIgnored(t *testing.T) {
pan115 := new115(map[string]any{"cookie": "UID=1; CID=2", "force_proxy": "true"}, http.DefaultClient)
if pan115.proxy {
t.Fatalf("115 should keep safe direct mode; force_proxy is deprecated")
}
cd2 := newCloudDrive2(map[string]any{"url": "http://example.test/dav", "force_302": "true"}, http.DefaultClient)
if !cd2.proxy {
t.Fatalf("clouddrive2 should keep safe proxy mode; force_302 is deprecated")
}
}
func TestUnsupportedProvider(t *testing.T) {
if _, err := New("dropbox", nil, nil); err != ErrUnsupported {
t.Fatalf("want ErrUnsupported, got %v", err)
}
if _, err := New("quark", nil, nil); err != ErrUnsupported {
t.Fatalf("quark should be unsupported, got %v", err)
}
if IsCloudType("quark") {
t.Fatal("quark should not be an active cloud provider")
}
}
-237
View File
@@ -1,237 +0,0 @@
package cloud
import (
"context"
"encoding/base64"
"fmt"
"net/http"
"net/url"
"path"
"strings"
)
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
//
// CloudDrive2 integrates many cloud disks (115 / 123 / Aliyun and more).
// Treating it as a WebDAV-backed cloud provider lets MediaStationGo
// browse, mount and upload to those disks without carrying every provider's
// private chunk-upload protocol in this project.
type cloudDrive2Provider struct {
typ string
name string
base *url.URL
username string
password string
token string
ua string
apiBase *url.URL
client *http.Client
proxy bool
}
func newCloudDrive2(cfg map[string]any, client *http.Client) *cloudDrive2Provider {
return newCloudDAVProvider(TypeCloudDrive2, "clouddrive2", cfg, client, "/dav")
}
func newOpenList(cfg map[string]any, client *http.Client) *cloudDrive2Provider {
return newCloudDAVProvider(TypeOpenList, "openlist", cfg, client, "/dav")
}
func newCloudDAVProvider(typ, name string, cfg map[string]any, client *http.Client, defaultDAVPath string) *cloudDrive2Provider {
rawURL := webDAVURLFromConfig(cfg, defaultDAVPath)
u, _ := url.Parse(strings.TrimRight(rawURL, "/"))
var apiBase *url.URL
if typ == TypeOpenList {
apiBase = openListAPIBaseFromConfig(cfg, rawURL, defaultDAVPath)
}
ua := str(cfg["ua"])
if ua == "" {
ua = defaultUA
}
proxy := true
return &cloudDrive2Provider{
typ: typ,
name: name,
base: u,
username: str(cfg["username"]),
password: str(cfg["password"]),
token: str(cfg["token"]),
ua: ua,
apiBase: apiBase,
client: client,
proxy: proxy,
}
}
func (p *cloudDrive2Provider) Type() string { return p.typ }
func (p *cloudDrive2Provider) Ping(ctx context.Context) error {
_, err := p.List(ctx, "")
return err
}
func (p *cloudDrive2Provider) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
if err := p.validate(); err != nil {
return nil, err
}
ref := normalizeCloudDAVPath(fileRef)
if ref == "/" {
return nil, fmt.Errorf("%s: file reference required", p.name)
}
if p.typ == TypeOpenList && isCloudVideoPlaybackCandidate(ref) {
if p.apiBase == nil {
return nil, fmt.Errorf("%s: pure 302 playback requires an OpenList API server address; configure server/api_url so /api/fs/get can return raw_url", p.name)
}
link, err := p.resolveOpenListAPIDirect(ctx, ref)
if err != nil {
return nil, fmt.Errorf("%s: pure 302 playback requires OpenList raw_url for %s: %w", p.name, ref, err)
}
return link, nil
}
if p.typ == TypeCloudDrive2 && isCloudVideoPlaybackCandidate(ref) {
link, err := p.resolveCloudDAVRedirectDirect(ctx, ref)
if err != nil {
return nil, fmt.Errorf("%s: pure 302 playback requires CloudDrive2/WebDAV to return a CDN Location for %s: %w", p.name, ref, err)
}
return link, nil
}
headers := map[string]string{
"User-Agent": p.ua,
}
if p.token != "" {
headers["Authorization"] = p.token
} else if p.username != "" {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(p.username+":"+p.password))
}
return &DirectLink{URL: p.urlFor(ref), Headers: headers, Proxy: p.proxy}, nil
}
func (p *cloudDrive2Provider) validate() error {
if p.base == nil || p.base.Scheme == "" || p.base.Host == "" {
return fmt.Errorf("%s: missing WebDAV URL", p.name)
}
return nil
}
func webDAVURLFromConfig(cfg map[string]any, defaultDAVPath string) string {
rawURL := str(cfg["url"])
if rawURL == "" {
rawURL = str(cfg["webdav_url"])
}
if rawURL != "" {
return ensureDefaultDAVPath(rawURL, defaultDAVPath)
}
return defaultWebDAVURL(str(cfg["server"]), defaultDAVPath)
}
func defaultWebDAVURL(server, defaultDAVPath string) string {
server = strings.TrimRight(strings.TrimSpace(server), "/")
if server == "" {
return ""
}
davPath := strings.TrimSpace(defaultDAVPath)
if davPath == "" {
return server
}
if !strings.HasPrefix(davPath, "/") {
davPath = "/" + davPath
}
return server + davPath
}
func openListAPIBaseFromConfig(cfg map[string]any, webDAVURL, defaultDAVPath string) *url.URL {
raw := str(cfg["server"])
if raw == "" {
raw = firstNonEmpty(str(cfg["api_url"]), webDAVURL)
}
raw = strings.TrimRight(strings.TrimSpace(raw), "/")
if raw == "" {
return nil
}
u, err := url.Parse(raw)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil
}
davPath := strings.Trim(strings.TrimSpace(defaultDAVPath), "/")
if davPath != "" {
pathParts := strings.Split(strings.TrimRight(u.Path, "/"), "/")
if len(pathParts) > 0 && strings.EqualFold(pathParts[len(pathParts)-1], davPath) {
u.Path = strings.Join(pathParts[:len(pathParts)-1], "/")
if u.Path == "" {
u.Path = "/"
}
}
}
u.RawPath = ""
u.RawQuery = ""
u.Fragment = ""
return u
}
func (p *cloudDrive2Provider) openListAPIURL(apiPath string) string {
if p.apiBase == nil {
return ""
}
u := *p.apiBase
u.RawPath = ""
basePath := strings.TrimRight(u.Path, "/")
apiPath = "/" + strings.TrimLeft(apiPath, "/")
if basePath == "" || basePath == "/" {
u.Path = apiPath
} else {
u.Path = basePath + apiPath
}
return u.String()
}
func ensureDefaultDAVPath(rawURL, defaultDAVPath string) string {
rawURL = strings.TrimRight(strings.TrimSpace(rawURL), "/")
if rawURL == "" {
return ""
}
u, err := url.Parse(rawURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return rawURL
}
if strings.TrimSpace(defaultDAVPath) == "" {
return rawURL
}
if u.Path == "" || u.Path == "/" {
davPath := strings.TrimSpace(defaultDAVPath)
if !strings.HasPrefix(davPath, "/") {
davPath = "/" + davPath
}
u.Path = davPath
u.RawPath = ""
return strings.TrimRight(u.String(), "/")
}
return rawURL
}
func normalizeCloudDAVPath(p string) string {
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
if p == "" || p == "." {
return "/"
}
if !strings.HasPrefix(p, "/") {
p = "/" + p
}
cleaned := path.Clean(p)
if cleaned == "." {
return "/"
}
return cleaned
}
func sameCloudDAVPath(a, b string) bool {
return strings.TrimRight(normalizeCloudDAVPath(a), "/") == strings.TrimRight(normalizeCloudDAVPath(b), "/")
}
func firstNonEmpty(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
return strings.TrimSpace(v)
}
}
return ""
}
-172
View File
@@ -1,172 +0,0 @@
package cloud
import (
"context"
"encoding/base64"
"encoding/xml"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strings"
)
func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry, error) {
if err := p.validate(); err != nil {
return nil, err
}
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
return p.listOpenListAPI(ctx, dir)
}
target := normalizeCloudDAVPath(dir)
req, err := http.NewRequestWithContext(ctx, "PROPFIND", p.urlFor(target), strings.NewReader(cloudDAVPropfindBody))
if err != nil {
return nil, err
}
p.auth(req)
req.Header.Set("Depth", "1")
req.Header.Set("Content-Type", "application/xml; charset=utf-8")
req.Header.Set("Accept", "application/xml,text/xml,*/*")
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.urlFor(target), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.decorateDAVStatusError(resp, target)
}
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
var multi cloudDAVMultiStatus
if err := xml.Unmarshal(body, &multi); err != nil {
return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err)
}
basePath := strings.TrimRight(p.base.EscapedPath(), "/")
currentID := normalizeCloudDAVPath(target)
out := make([]FileEntry, 0, len(multi.Responses))
for _, item := range multi.Responses {
entryPath, err := p.entryIDFromHref(item.Href, basePath)
if err != nil || entryPath == "" || sameCloudDAVPath(entryPath, currentID) {
continue
}
name := firstNonEmpty(item.PropStat.Prop.DisplayName, path.Base(strings.TrimRight(entryPath, "/")))
if decoded, err := url.PathUnescape(name); err == nil {
name = decoded
}
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: entryPath,
Name: name,
IsDir: item.PropStat.Prop.ResourceType.Collection != nil || strings.HasSuffix(item.Href, "/"),
Size: parseDAVSize(item.PropStat.Prop.ContentLength),
})
}
return out, nil
}
func (p *cloudDrive2Provider) resolveCloudDAVRedirectDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
target := p.urlFor(fileRef)
headers := map[string]string{
"User-Agent": p.ua,
}
if p.token != "" {
headers["Authorization"] = p.token
} else if p.username != "" {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(p.username+":"+p.password))
}
location, status, err := p.firstHTTPRedirectLocation(ctx, target, headers)
if err != nil {
return nil, decorateDAVTransportError(p.name, target, err)
}
if location == "" {
return nil, fmt.Errorf("%s: WebDAV %s returned http %d without CDN Location; refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, status)
}
return &DirectLink{URL: location, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) firstHTTPRedirectLocation(ctx context.Context, target string, headers map[string]string) (string, int, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
if err != nil {
return "", 0, err
}
req.Header.Set("Accept", "*/*")
req.Header.Set("Accept-Encoding", "identity")
req.Header.Set("Range", "bytes=0-0")
if strings.TrimSpace(p.ua) != "" {
req.Header.Set("User-Agent", p.ua)
}
for key, value := range headers {
key = strings.TrimSpace(key)
if key != "" && strings.TrimSpace(value) != "" {
req.Header.Set(key, value)
}
}
client := p.client
if client == nil {
client = http.DefaultClient
}
noFollow := *client
noFollow.CheckRedirect = func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
}
resp, err := noFollow.Do(req)
if err != nil {
return "", 0, err
}
defer resp.Body.Close()
status := resp.StatusCode
if status >= 300 && status < 400 {
rawLocation := strings.TrimSpace(resp.Header.Get("Location"))
if rawLocation == "" {
return "", status, fmt.Errorf("%s: upstream returned redirect http %d without Location", p.name, status)
}
location, err := resolveHTTPRedirectLocation(target, rawLocation)
if err != nil {
return "", status, err
}
return location, status, nil
}
return "", status, nil
}
func resolveHTTPRedirectLocation(baseURL, rawLocation string) (string, error) {
rawLocation = strings.TrimSpace(rawLocation)
if rawLocation == "" {
return "", fmt.Errorf("empty redirect Location")
}
if strings.HasPrefix(rawLocation, "//") {
base, err := url.Parse(baseURL)
if err != nil || base.Scheme == "" {
return "", fmt.Errorf("protocol-relative redirect Location without base scheme")
}
rawLocation = base.Scheme + ":" + rawLocation
}
location, err := url.Parse(rawLocation)
if err != nil {
return "", fmt.Errorf("invalid redirect Location: %w", err)
}
if location.IsAbs() {
if location.Scheme != "http" && location.Scheme != "https" {
return "", fmt.Errorf("unsupported redirect Location scheme %q", location.Scheme)
}
return location.String(), nil
}
base, err := url.Parse(baseURL)
if err != nil {
return "", fmt.Errorf("invalid redirect base URL: %w", err)
}
return base.ResolveReference(location).String(), nil
}
func (p *cloudDrive2Provider) auth(req *http.Request) {
req.Header.Set("User-Agent", p.ua)
if p.token != "" {
req.Header.Set("Authorization", p.token)
return
}
if p.username != "" {
req.SetBasicAuth(p.username, p.password)
}
}
@@ -1,55 +0,0 @@
package cloud
import (
"fmt"
"io"
"net/http"
"strings"
)
func (p *cloudDrive2Provider) decorateDAVStatusError(resp *http.Response, target string) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
detail := compactDAVErrorBody(string(body))
if detail == "" {
if resp.StatusCode == http.StatusMethodNotAllowed {
return fmt.Errorf("%s: list %s returned http %d;请确认填写的是 WebDAV 地址(通常以 /dav 结尾),并且桥接网盘已在 OpenList/CloudDrive2 内完成登录或 Cookie 保存", p.name, target, resp.StatusCode)
}
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
return fmt.Errorf("%s: list %s returned http %d;请填写 OpenList/CloudDrive2 的 Token 或用户名密码,或确认 WebDAV 凭据可用", p.name, target, resp.StatusCode)
}
return fmt.Errorf("%s: list %s returned http %d", p.name, target, resp.StatusCode)
}
if resp.StatusCode == http.StatusMethodNotAllowed {
return fmt.Errorf("%s: list %s returned http %d:%s;请确认填写的是 WebDAV 地址(通常以 /dav 结尾),并且桥接网盘已在 OpenList/CloudDrive2 内完成登录或 Cookie 保存", p.name, target, resp.StatusCode, detail)
}
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
return fmt.Errorf("%s: list %s returned http %d:%s;请检查 WebDAV 用户名/密码、Authorization Token,或先在 OpenList/CloudDrive2 中保存对应网盘 Cookie", p.name, target, resp.StatusCode, detail)
}
return fmt.Errorf("%s: list %s returned http %d:%s", p.name, target, resp.StatusCode, detail)
}
func compactDAVErrorBody(raw string) string {
raw = strings.TrimSpace(strings.ReplaceAll(raw, "\x00", ""))
if raw == "" {
return ""
}
raw = strings.Join(strings.Fields(raw), " ")
if len([]rune(raw)) > 180 {
return string([]rune(raw)[:180]) + "…"
}
return raw
}
func decorateDAVTransportError(name, target string, err error) error {
if err == nil {
return nil
}
message := err.Error()
if strings.Contains(message, "server gave HTTP response to HTTPS client") {
return fmt.Errorf("%s: %w;当前地址使用 https://,但服务端返回 HTTP。请改用 http:// 地址,例如 OpenList 默认 WebDAV 通常是 http://host:5244/dav/;如果必须使用 https,请在 OpenList 前配置反向代理和证书", name, err)
}
if strings.Contains(message, "first record does not look like a TLS handshake") {
return fmt.Errorf("%s: %w;疑似把 HTTP 服务配置成了 https://,请检查 %s 的协议头", name, err, target)
}
return err
}
@@ -1,85 +0,0 @@
package cloud
import (
"net/url"
"strconv"
"strings"
)
func (p *cloudDrive2Provider) urlFor(remotePath string) string {
u := *p.base
u.RawPath = ""
basePath := strings.TrimRight(u.Path, "/")
remote := strings.Trim(normalizeCloudDAVPath(remotePath), "/")
switch {
case basePath == "" || basePath == "/":
if remote == "" {
u.Path = "/"
} else {
u.Path = "/" + remote
}
case remote == "":
u.Path = basePath
default:
u.Path = basePath + "/" + remote
}
return u.String()
}
func (p *cloudDrive2Provider) entryIDFromHref(href, basePath string) (string, error) {
if href == "" {
return "", nil
}
parsed, err := url.Parse(href)
if err != nil {
return "", err
}
hrefPath := parsed.EscapedPath()
if hrefPath == "" {
hrefPath = href
}
if basePath != "" && basePath != "/" {
hrefPath = strings.TrimPrefix(hrefPath, basePath)
}
if decoded, err := url.PathUnescape(hrefPath); err == nil {
hrefPath = decoded
}
return normalizeCloudDAVPath(hrefPath), nil
}
const cloudDAVPropfindBody = `<?xml version="1.0" encoding="utf-8"?>
<d:propfind xmlns:d="DAV:">
<d:prop>
<d:displayname/>
<d:getcontentlength/>
<d:resourcetype/>
</d:prop>
</d:propfind>`
type cloudDAVMultiStatus struct {
Responses []cloudDAVResponse `xml:"response"`
}
type cloudDAVResponse struct {
Href string `xml:"href"`
PropStat cloudDAVPropStat `xml:"propstat"`
}
type cloudDAVPropStat struct {
Prop cloudDAVProp `xml:"prop"`
}
type cloudDAVProp struct {
DisplayName string `xml:"displayname"`
ContentLength string `xml:"getcontentlength"`
ResourceType cloudDAVResourceType `xml:"resourcetype"`
}
type cloudDAVResourceType struct {
Collection *struct{} `xml:"collection"`
}
func parseDAVSize(raw string) int64 {
n, _ := strconv.ParseInt(strings.TrimSpace(raw), 10, 64)
return n
}
@@ -1,233 +0,0 @@
package cloud
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strings"
)
func (p *cloudDrive2Provider) Mkdir(ctx context.Context, parentDir, name string) (*FileEntry, error) {
cleanName, err := cleanCloudEntryName(name)
if err != nil {
return nil, err
}
parent := normalizeCloudDAVPath(parentDir)
target := joinOpenListAPIPath(parent, cleanName)
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIMkdir(ctx, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
if err := p.webDAVMkdir(ctx, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
func (p *cloudDrive2Provider) Rename(ctx context.Context, ref, name string) (*FileEntry, error) {
cleanName, err := cleanCloudEntryName(name)
if err != nil {
return nil, err
}
source := normalizeCloudDAVPath(ref)
if source == "/" {
return nil, fmt.Errorf("%s: cannot rename root directory", p.name)
}
target := joinOpenListAPIPath(path.Dir(source), cleanName)
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIRename(ctx, source, cleanName); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
if err := p.webDAVRename(ctx, source, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
func (p *cloudDrive2Provider) Move(ctx context.Context, ref, targetDir, name string) (*FileEntry, error) {
source := normalizeCloudDAVPath(ref)
if source == "/" {
return nil, fmt.Errorf("%s: cannot move root directory", p.name)
}
cleanName := strings.TrimSpace(name)
if cleanName == "" {
cleanName = path.Base(source)
}
var err error
cleanName, err = cleanCloudEntryName(cleanName)
if err != nil {
return nil, err
}
targetDir = normalizeCloudDAVPath(targetDir)
target := joinOpenListAPIPath(targetDir, cleanName)
if sameCloudDAVPath(source, target) {
return &FileEntry{ID: target, Name: cleanName}, nil
}
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIMove(ctx, source, targetDir, cleanName); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName}, nil
}
if err := p.webDAVRename(ctx, source, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName}, nil
}
func cleanCloudEntryName(name string) (string, error) {
name = strings.TrimSpace(name)
if name == "" || name == "." || name == ".." {
return "", fmt.Errorf("entry name is required")
}
if strings.ContainsAny(name, `/\`) {
return "", fmt.Errorf("entry name cannot contain path separators")
}
return name, nil
}
func (p *cloudDrive2Provider) openListAPIMkdir(ctx context.Context, target string) error {
return p.openListAPIPost(ctx, "/api/fs/mkdir", map[string]string{"path": normalizeCloudDAVPath(target)}, "mkdir")
}
func (p *cloudDrive2Provider) openListAPIRename(ctx context.Context, source, name string) error {
return p.openListAPIPost(ctx, "/api/fs/rename", map[string]string{
"path": normalizeCloudDAVPath(source),
"name": name,
}, "rename")
}
func (p *cloudDrive2Provider) openListAPIMove(ctx context.Context, source, targetDir, targetName string) error {
targetDir = normalizeCloudDAVPath(targetDir)
sourceName := path.Base(normalizeCloudDAVPath(source))
if sameCloudDAVPath(path.Dir(source), targetDir) {
if sourceName == targetName {
return nil
}
return p.openListAPIRename(ctx, source, targetName)
}
if err := p.openListAPIPost(ctx, "/api/fs/move", map[string]any{
"src_dir": normalizeCloudDAVPath(path.Dir(source)),
"dst_dir": targetDir,
"names": []string{sourceName},
}, "move"); err != nil {
return err
}
if sourceName != targetName {
moved := joinOpenListAPIPath(targetDir, sourceName)
return p.openListAPIRename(ctx, moved, targetName)
}
return nil
}
func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath string, payload any, action string) error {
token, err := p.openListAPIToken(ctx)
if err != nil {
return err
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("%s: api %s returned http %d", p.name, action, resp.StatusCode)
}
var decoded struct {
Code int `json:"code"`
Message string `json:"message"`
}
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return fmt.Errorf("%s: decode api %s: %w", p.name, action, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return fmt.Errorf("%s: api %s failed: %s", p.name, action, msg)
}
return nil
}
func (p *cloudDrive2Provider) webDAVMkdir(ctx context.Context, target string) error {
req, err := http.NewRequestWithContext(ctx, "MKCOL", p.urlFor(target), nil)
if err != nil {
return err
}
p.auth(req)
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.urlFor(target), err)
}
defer resp.Body.Close()
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent:
return nil
case http.StatusMethodNotAllowed:
return fmt.Errorf("%s: mkdir %s returned http %d; the folder may already exist or this WebDAV backend is read-only", p.name, target, resp.StatusCode)
default:
return p.decorateDAVMutationStatusError(resp, "mkdir", target)
}
}
func (p *cloudDrive2Provider) webDAVRename(ctx context.Context, source, target string) error {
req, err := http.NewRequestWithContext(ctx, "MOVE", p.urlFor(source), nil)
if err != nil {
return err
}
p.auth(req)
req.Header.Set("Destination", p.webDAVDestination(target))
req.Header.Set("Overwrite", "F")
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.urlFor(source), err)
}
defer resp.Body.Close()
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent:
return nil
default:
return p.decorateDAVMutationStatusError(resp, "rename", source)
}
}
func (p *cloudDrive2Provider) webDAVDestination(target string) string {
raw := p.urlFor(target)
u, err := url.Parse(raw)
if err != nil {
return raw
}
u.RawQuery = ""
u.Fragment = ""
return u.String()
}
func (p *cloudDrive2Provider) decorateDAVMutationStatusError(resp *http.Response, action, target string) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
detail := compactDAVErrorBody(string(body))
if detail == "" {
return fmt.Errorf("%s: %s %s returned http %d", p.name, action, target, resp.StatusCode)
}
return fmt.Errorf("%s: %s %s returned http %d:%s", p.name, action, target, resp.StatusCode, detail)
}
@@ -1,240 +0,0 @@
package cloud
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
)
func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
const pageSize = 500
target := normalizeCloudDAVPath(dir)
out := make([]FileEntry, 0, pageSize)
for pageNum := 1; ; pageNum++ {
payload := map[string]any{
"path": target,
"password": "",
"page": pageNum,
"per_page": pageSize,
"refresh": false,
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/list"), bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/list"), err)
}
var decoded openListListResponse
decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded)
resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.openListAPIStatusError("list", target, resp.StatusCode)
}
if decodeErr != nil {
return nil, fmt.Errorf("%s: decode api list: %w", p.name, decodeErr)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api list %s failed: %s", p.name, target, msg)
}
for _, item := range decoded.Data.Content {
name := strings.TrimSpace(item.Name)
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: joinOpenListAPIPath(target, name),
Name: name,
IsDir: item.IsDir,
Size: item.Size,
})
}
total := decoded.Data.Total
if total > 0 {
if len(out) >= total || len(decoded.Data.Content) == 0 {
break
}
continue
}
if len(decoded.Data.Content) == 0 || len(decoded.Data.Content) < pageSize {
break
}
}
return out, nil
}
func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
payload, _ := json.Marshal(map[string]string{"path": normalizeCloudDAVPath(fileRef), "password": ""})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/get"), bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.openListAPIStatusError("get", fileRef, resp.StatusCode)
}
var decoded openListGetResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return nil, fmt.Errorf("%s: decode api get: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api get %s failed: %s", p.name, fileRef, msg)
}
raw := firstNonEmpty(decoded.Data.RawURL, decoded.Data.URL)
if raw == "" {
return nil, fmt.Errorf("%s: api get %s returned empty raw_url", p.name, fileRef)
}
resolved, err := p.resolveOpenListPlaybackURL(raw)
if err != nil {
return nil, err
}
headers := normalizeOpenListPlaybackHeaders(decoded.Data.Header)
if len(headers) > 0 {
return nil, fmt.Errorf("%s: api get %s returned raw_url that requires headers (%s); refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, strings.Join(sortedHeaderNames(headers), ","))
}
resolved, err = p.resolveOpenListCDNRedirect(ctx, fileRef, resolved)
if err != nil {
return nil, err
}
return &DirectLink{URL: resolved, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) resolveOpenListCDNRedirect(ctx context.Context, fileRef, rawURL string) (string, error) {
if p.apiBase == nil || !sameURLHost(rawURL, p.apiBase) {
return rawURL, nil
}
location, status, err := p.firstHTTPRedirectLocation(ctx, rawURL, nil)
if err != nil {
return "", fmt.Errorf("%s: probe raw_url %s failed: %w", p.name, fileRef, err)
}
if location != "" {
return location, nil
}
return "", fmt.Errorf("%s: api get %s returned an OpenList-hosted raw_url with http %d and no CDN Location; refusing OpenList/WebDAV proxy fallback for pure 302 playback", p.name, fileRef, status)
}
func (p *cloudDrive2Provider) openListAPIStatusError(action, target string, status int) error {
if status == http.StatusUnauthorized || status == http.StatusForbidden {
return fmt.Errorf("%s: api %s %s returned http %d;请检查 OpenList Token 或用户名密码,并确认填写的是 OpenList 服务地址而不是 /dav 地址", p.name, action, target, status)
}
return fmt.Errorf("%s: api %s %s returned http %d", p.name, action, target, status)
}
func (p *cloudDrive2Provider) hasOpenListAPICredentials() bool {
return strings.TrimSpace(p.token) != "" || (strings.TrimSpace(p.username) != "" && p.password != "")
}
func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) {
if token := strings.TrimSpace(p.token); token != "" {
return token, nil
}
if strings.TrimSpace(p.username) == "" || p.password == "" {
return "", nil
}
payload, _ := json.Marshal(map[string]string{
"username": p.username,
"password": p.password,
})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/auth/login"), bytes.NewReader(payload))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
resp, err := p.client.Do(req)
if err != nil {
return "", decorateDAVTransportError(p.name, p.openListAPIURL("/api/auth/login"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return "", fmt.Errorf("%s: api login returned http %d", p.name, resp.StatusCode)
}
var decoded openListLoginResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return "", fmt.Errorf("%s: decode api login: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return "", fmt.Errorf("%s: api login failed: %s", p.name, msg)
}
token := strings.TrimSpace(decoded.Data.Token)
if token == "" {
return "", fmt.Errorf("%s: api login returned empty token", p.name)
}
p.token = token
return token, nil
}
func (p *cloudDrive2Provider) resolveOpenListPlaybackURL(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", fmt.Errorf("%s: empty playback URL", p.name)
}
if strings.HasPrefix(raw, "//") {
if p.apiBase == nil || p.apiBase.Scheme == "" {
return "", fmt.Errorf("%s: protocol-relative playback URL without API base", p.name)
}
raw = p.apiBase.Scheme + ":" + raw
}
u, err := url.Parse(raw)
if err != nil {
return "", fmt.Errorf("%s: invalid playback URL: %w", p.name, err)
}
if u.IsAbs() {
if u.Scheme != "http" && u.Scheme != "https" {
return "", fmt.Errorf("%s: unsupported playback URL scheme %q", p.name, u.Scheme)
}
return u.String(), nil
}
if p.apiBase == nil {
return "", fmt.Errorf("%s: relative playback URL without API base", p.name)
}
base := *p.apiBase
base.RawPath = ""
base.RawQuery = ""
base.Fragment = ""
return base.ResolveReference(u).String(), nil
}
@@ -1,126 +0,0 @@
package cloud
import (
"encoding/json"
"net/url"
"path"
"sort"
"strings"
)
func sortedHeaderNames(headers map[string]string) []string {
if len(headers) == 0 {
return nil
}
out := make([]string, 0, len(headers))
for key := range headers {
key = strings.TrimSpace(key)
if key != "" {
out = append(out, key)
}
}
sort.Strings(out)
return out
}
func sameURLHost(raw string, base *url.URL) bool {
if base == nil {
return false
}
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil {
return false
}
if !u.IsAbs() {
return true
}
return strings.EqualFold(u.Host, base.Host)
}
func normalizeOpenListPlaybackHeaders(raw json.RawMessage) map[string]string {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
var obj map[string]any
if err := json.Unmarshal(raw, &obj); err != nil {
return nil
}
out := make(map[string]string, len(obj))
for k, v := range obj {
key := strings.TrimSpace(k)
if key == "" {
continue
}
switch value := v.(type) {
case string:
if strings.TrimSpace(value) != "" {
out[key] = strings.TrimSpace(value)
}
case []any:
parts := make([]string, 0, len(value))
for _, item := range value {
if s, ok := item.(string); ok && strings.TrimSpace(s) != "" {
parts = append(parts, strings.TrimSpace(s))
}
}
if len(parts) > 0 {
out[key] = strings.Join(parts, ", ")
}
}
}
if len(out) == 0 {
return nil
}
return out
}
func isCloudVideoPlaybackCandidate(fileRef string) bool {
switch strings.ToLower(path.Ext(strings.TrimSpace(fileRef))) {
case ".mkv", ".mp4", ".m4v", ".avi", ".mov", ".webm", ".ts", ".rmvb", ".rm", ".3gp", ".mpg", ".mpeg":
return true
default:
return false
}
}
type openListListResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Content []openListListItem `json:"content"`
Total int `json:"total"`
} `json:"data"`
}
type openListListItem struct {
Name string `json:"name"`
Size int64 `json:"size"`
IsDir bool `json:"is_dir"`
}
type openListGetResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
RawURL string `json:"raw_url"`
URL string `json:"url"`
Header json.RawMessage `json:"header"`
} `json:"data"`
}
type openListLoginResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Token string `json:"token"`
} `json:"data"`
}
func joinOpenListAPIPath(dir, name string) string {
dir = strings.TrimRight(normalizeCloudDAVPath(dir), "/")
name = strings.Trim(strings.ReplaceAll(name, "\\", "/"), "/")
if dir == "" || dir == "/" {
return normalizeCloudDAVPath(name)
}
return normalizeCloudDAVPath(dir + "/" + name)
}
-171
View File
@@ -1,171 +0,0 @@
package cloud
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestCloudDrive2WebDAVListAndResolve(t *testing.T) {
var gotAuth, gotDepth, gotRange string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method == "PROPFIND" && r.URL.Path == "/dav":
gotAuth = r.Header.Get("Authorization")
gotDepth = r.Header.Get("Depth")
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="utf-8"?>
<d:multistatus xmlns:d="DAV:">
<d:response>
<d:href>/dav/</d:href>
<d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat>
</d:response>
<d:response>
<d:href>/dav/115/</d:href>
<d:propstat><d:prop><d:displayname>115</d:displayname><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat>
</d:response>
<d:response>
<d:href>/dav/123/Movie.mkv</d:href>
<d:propstat><d:prop><d:displayname>Movie.mkv</d:displayname><d:getcontentlength>789</d:getcontentlength><d:resourcetype/></d:prop></d:propstat>
</d:response>
</d:multistatus>`))
case r.Method == http.MethodGet && r.URL.Path == "/dav/123/Movie.mkv":
gotAuth = r.Header.Get("Authorization")
gotRange = r.Header.Get("Range")
http.Redirect(w, r, "https://cdn.example.test/123/Movie.mkv?sign=1", http.StatusFound)
default:
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeCloudDrive2, map[string]any{"url": srv.URL + "/dav", "username": "u", "password": "p"}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if gotDepth != "1" {
t.Fatalf("Depth = %q, want 1", gotDepth)
}
if !strings.HasPrefix(gotAuth, "Basic ") {
t.Fatalf("missing basic auth: %q", gotAuth)
}
if len(entries) != 2 {
t.Fatalf("entries = %#v", entries)
}
if !entries[0].IsDir || entries[0].ID != "/115" {
t.Fatalf("dir entry wrong: %#v", entries[0])
}
if entries[1].IsDir || entries[1].ID != "/123/Movie.mkv" || entries[1].Size != 789 {
t.Fatalf("file entry wrong: %#v", entries[1])
}
link, err := p.Resolve(context.Background(), entries[1].ID)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if link.URL != "https://cdn.example.test/123/Movie.mkv?sign=1" {
t.Fatalf("bad url: %s", link.URL)
}
if link.Proxy || len(link.Headers) != 0 {
t.Fatalf("clouddrive2 video should resolve to pure 302 link: %#v", link)
}
if gotRange != "bytes=0-0" {
t.Fatalf("resolve should probe with a tiny range, got %q", gotRange)
}
}
func TestCloudDrive2ResolveRejectsWebDAVProxyFallbackWithoutRedirect(t *testing.T) {
var getSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method == "PROPFIND" && r.URL.Path == "/dav":
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="utf-8"?><d:multistatus xmlns:d="DAV:"><d:response><d:href>/dav/</d:href><d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat></d:response></d:multistatus>`))
case r.Method == http.MethodGet && r.URL.Path == "/dav/123/Movie.mkv":
getSeen = true
w.Header().Set("Content-Range", "bytes 0-0/10")
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write([]byte("x"))
default:
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeCloudDrive2, map[string]any{"url": srv.URL + "/dav", "username": "u", "password": "p"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "/123/Movie.mkv")
if err == nil || !strings.Contains(err.Error(), "without CDN Location") || !strings.Contains(err.Error(), "refusing WebDAV/proxy fallback") {
t.Fatalf("resolve error = %v, want pure 302 refusal", err)
}
if !getSeen {
t.Fatal("expected CloudDrive2 WebDAV direct-link probe")
}
}
func TestCloudDrive2MutableProviderUsesWebDAV(t *testing.T) {
var mkcolSeen bool
var destinations []string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method == "MKCOL" && r.URL.Path == "/dav/TV":
mkcolSeen = true
w.WriteHeader(http.StatusCreated)
case r.Method == "MOVE" && r.URL.Path == "/dav/TV":
destinations = append(destinations, r.Header.Get("Destination"))
if r.Header.Get("Overwrite") != "F" {
t.Fatalf("Overwrite = %q, want F", r.Header.Get("Overwrite"))
}
w.WriteHeader(http.StatusCreated)
case r.Method == "MOVE" && r.URL.Path == "/dav/Inbox/Movie.mkv":
destinations = append(destinations, r.Header.Get("Destination"))
if r.Header.Get("Overwrite") != "F" {
t.Fatalf("Overwrite = %q, want F", r.Header.Get("Overwrite"))
}
w.WriteHeader(http.StatusCreated)
default:
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeCloudDrive2, map[string]any{"url": srv.URL + "/dav", "username": "u", "password": "p"}, srv.Client())
if err != nil {
t.Fatal(err)
}
mutable, ok := p.(MutableProvider)
if !ok {
t.Fatal("clouddrive2 should support mutable provider")
}
if _, err := mutable.Mkdir(context.Background(), "", "TV"); err != nil {
t.Fatalf("mkdir: %v", err)
}
if _, err := mutable.Rename(context.Background(), "/TV", "电视剧"); err != nil {
t.Fatalf("rename: %v", err)
}
moved, err := mutable.(MovableProvider).Move(context.Background(), "/Inbox/Movie.mkv", "/电影/欧美电影/Movie (2026)", "Movie (2026).mkv")
if err != nil {
t.Fatalf("move: %v", err)
}
if !mkcolSeen || len(destinations) != 2 {
t.Fatalf("mkcol=%v destinations=%#v, want mkdir and two MOVE calls", mkcolSeen, destinations)
}
if destinations[0] != srv.URL+"/dav/%E7%94%B5%E8%A7%86%E5%89%A7" {
t.Fatalf("rename Destination = %q", destinations[0])
}
if destinations[1] != srv.URL+"/dav/%E7%94%B5%E5%BD%B1/%E6%AC%A7%E7%BE%8E%E7%94%B5%E5%BD%B1/Movie%20%282026%29/Movie%20%282026%29.mkv" {
t.Fatalf("move Destination = %q", destinations[1])
}
if moved.ID != "/电影/欧美电影/Movie (2026)/Movie (2026).mkv" {
t.Fatalf("moved entry = %#v", moved)
}
}
@@ -1,280 +0,0 @@
package cloud
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestOpenListWebDAVListAndResolve(t *testing.T) {
var gotPath, gotDepth string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/api/auth/login" {
http.NotFound(w, r)
return
}
if r.URL.Path == "/api/fs/get" {
http.NotFound(w, r)
return
}
if r.Method != "PROPFIND" || r.URL.Path != "/dav" {
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
gotPath = r.URL.Path
gotDepth = r.Header.Get("Depth")
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="utf-8"?>
<d:multistatus xmlns:d="DAV:">
<d:response>
<d:href>/dav/</d:href>
<d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat>
</d:response>
<d:response>
<d:href>/dav/Cloud/Movie.mkv</d:href>
<d:propstat><d:prop><d:displayname>Movie.mkv</d:displayname><d:getcontentlength>1024</d:getcontentlength><d:resourcetype/></d:prop></d:propstat>
</d:response>
</d:multistatus>`))
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"url": srv.URL + "/dav"}, srv.Client())
if err != nil {
t.Fatal(err)
}
if p.Type() != TypeOpenList {
t.Fatalf("type = %q, want %q", p.Type(), TypeOpenList)
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if gotPath != "/dav" {
t.Fatalf("path = %q, want /dav", gotPath)
}
if gotDepth != "1" {
t.Fatalf("Depth = %q, want 1", gotDepth)
}
if len(entries) != 1 || entries[0].ID != "/Cloud/Movie.mkv" || entries[0].Size != 1024 {
t.Fatalf("entries = %#v", entries)
}
_, err = p.Resolve(context.Background(), entries[0].ID)
if err == nil || !strings.Contains(err.Error(), "pure 302 playback requires OpenList raw_url") {
t.Fatalf("openlist video resolve should require raw_url instead of WebDAV proxy fallback, err=%v", err)
}
}
func TestOpenListListUsesAPIUsernamePasswordWithoutWebDAVFallback(t *testing.T) {
var loginSeen, listSeen, davSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/api/auth/login":
loginSeen = true
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode login body: %v", err)
}
if body["username"] != "alice" || body["password"] != "secret" {
t.Fatalf("login body = %#v", body)
}
_, _ = w.Write([]byte(`{"code":200,"data":{"token":"api-token"}}`))
case "/api/fs/list":
listSeen = true
if r.Header.Get("Authorization") != "api-token" {
t.Fatalf("Authorization = %q, want api-token", r.Header.Get("Authorization"))
}
_, _ = w.Write([]byte(`{"code":200,"data":{"content":[{"name":"Movies","is_dir":true,"size":0},{"name":"Movie.mkv","is_dir":false,"size":1024}],"total":2}}`))
case "/dav":
davSeen = true
w.WriteHeader(http.StatusMultiStatus)
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "username": "alice", "password": "secret"}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if !loginSeen || !listSeen {
t.Fatalf("expected api login/list, login=%v list=%v", loginSeen, listSeen)
}
if davSeen {
t.Fatal("openlist API credentials should not fall back to WebDAV")
}
if len(entries) != 2 || entries[0].ID != "/Movies" || !entries[0].IsDir || entries[1].ID != "/Movie.mkv" || entries[1].Size != 1024 {
t.Fatalf("entries = %#v", entries)
}
}
func TestOpenListMutableProviderUsesAPI(t *testing.T) {
var mkdirPath, renamePath, renameName, moveSrcDir, moveDstDir string
var moveNames []string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/api/fs/mkdir":
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode mkdir body: %v", err)
}
mkdirPath = body["path"]
if r.Header.Get("Authorization") != "alist-token" {
t.Fatalf("mkdir Authorization = %q", r.Header.Get("Authorization"))
}
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
case "/api/fs/rename":
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode rename body: %v", err)
}
renamePath = body["path"]
renameName = body["name"]
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
case "/api/fs/move":
var body struct {
SrcDir string `json:"src_dir"`
DstDir string `json:"dst_dir"`
Names []string `json:"names"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode move body: %v", err)
}
moveSrcDir = body.SrcDir
moveDstDir = body.DstDir
moveNames = body.Names
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
mutable, ok := p.(MutableProvider)
if !ok {
t.Fatal("openlist should support mutable provider")
}
created, err := mutable.Mkdir(context.Background(), "/电视剧", "欧美剧")
if err != nil {
t.Fatalf("mkdir: %v", err)
}
if mkdirPath != "/电视剧/欧美剧" || created.ID != "/电视剧/欧美剧" || !created.IsDir {
t.Fatalf("mkdir path=%q entry=%#v", mkdirPath, created)
}
renamed, err := mutable.Rename(context.Background(), "/电视剧/欧美剧", "美剧")
if err != nil {
t.Fatalf("rename: %v", err)
}
if renamePath != "/电视剧/欧美剧" || renameName != "美剧" || renamed.ID != "/电视剧/美剧" {
t.Fatalf("rename path=%q name=%q entry=%#v", renamePath, renameName, renamed)
}
moved, err := mutable.(MovableProvider).Move(context.Background(), "/待整理/Show.S01E01.mkv", "/动漫/国漫/Show/Season 01", "Show - S01E01.mkv")
if err != nil {
t.Fatalf("move: %v", err)
}
if moveSrcDir != "/待整理" || moveDstDir != "/动漫/国漫/Show/Season 01" || len(moveNames) != 1 || moveNames[0] != "Show.S01E01.mkv" {
t.Fatalf("move src=%q dst=%q names=%#v", moveSrcDir, moveDstDir, moveNames)
}
if renamePath != "/动漫/国漫/Show/Season 01/Show.S01E01.mkv" || renameName != "Show - S01E01.mkv" {
t.Fatalf("post-move rename path=%q name=%q", renamePath, renameName)
}
if moved.ID != "/动漫/国漫/Show/Season 01/Show - S01E01.mkv" {
t.Fatalf("moved entry = %#v", moved)
}
}
func TestOpenListListAPIFailureDoesNotFallbackToWebDAV(t *testing.T) {
var davSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/auth/login":
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":500,"message":"bad password"}`))
case "/dav":
davSeen = true
w.WriteHeader(http.StatusMultiStatus)
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "username": "alice", "password": "bad"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.List(context.Background(), "")
if err == nil || !strings.Contains(err.Error(), "api login failed") || !strings.Contains(err.Error(), "bad password") {
t.Fatalf("list error = %v, want api login failure", err)
}
if davSeen {
t.Fatal("openlist API failure fell back to WebDAV")
}
}
func TestOpenListRootURLDefaultsToDAV(t *testing.T) {
var gotPath string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="utf-8"?><d:multistatus xmlns:d="DAV:"><d:response><d:href>/dav/</d:href><d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat></d:response></d:multistatus>`))
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"url": srv.URL + "/"}, srv.Client())
if err != nil {
t.Fatal(err)
}
if _, err := p.List(context.Background(), ""); err != nil {
t.Fatalf("list: %v", err)
}
if gotPath != "/dav" {
t.Fatalf("path = %q, want /dav", gotPath)
}
}
func TestOpenListURLForKeepsNonASCIIPathSingleEncoded(t *testing.T) {
p := newOpenList(map[string]any{"url": "http://example.test:5244/dav/"}, nil)
got := p.urlFor("/动画电影/爱宠大机密2 (2019) {tmdb-412117}")
if strings.Contains(got, "%25E") {
t.Fatalf("url is double-escaped: %s", got)
}
want := "http://example.test:5244/dav/%E5%8A%A8%E7%94%BB%E7%94%B5%E5%BD%B1/%E7%88%B1%E5%AE%A0%E5%A4%A7%E6%9C%BA%E5%AF%862%20%282019%29%20%7Btmdb-412117%7D"
if got != want {
t.Fatalf("url = %s, want %s", got, want)
}
}
func TestOpenListDAVStatusErrorIncludesBodyHint(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusMethodNotAllowed)
_, _ = w.Write([]byte("请先填写有效Cookie并保存"))
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"url": srv.URL + "/dav"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.List(context.Background(), "")
if err == nil {
t.Fatal("want error")
}
if !strings.Contains(err.Error(), "请先填写有效Cookie并保存") || !strings.Contains(err.Error(), "WebDAV 地址") {
t.Fatalf("unexpected error: %v", err)
}
}
@@ -1,204 +0,0 @@
package cloud
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestOpenListResolveUsesAPIRawURLFor302Playback(t *testing.T) {
var gotPath, gotAuth string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
gotAuth = r.Header.Get("Authorization")
if r.Method != http.MethodPost || r.URL.Path != "/api/fs/get" {
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":200,"data":{"raw_url":"https://cdn.example.test/movie.mkv?sign=1"}}`))
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
link, err := p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if gotPath != "/api/fs/get" {
t.Fatalf("api path = %q, want /api/fs/get", gotPath)
}
if gotAuth != "alist-token" {
t.Fatalf("Authorization = %q, want token", gotAuth)
}
if link.URL != "https://cdn.example.test/movie.mkv?sign=1" {
t.Fatalf("url = %q", link.URL)
}
if link.Proxy {
t.Fatalf("openlist raw_url without required headers should be 302 playback")
}
}
func TestOpenListResolveCollapsesHostedRawURLRedirectToCDN(t *testing.T) {
var probeSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/fs/get":
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":200,"data":{"raw_url":"/d/Cloud/Movie.mkv?sign=1"}}`))
case "/d/Cloud/Movie.mkv":
probeSeen = true
if r.Header.Get("Range") != "bytes=0-0" {
t.Fatalf("probe Range = %q", r.Header.Get("Range"))
}
http.Redirect(w, r, "https://cdn.example.test/movie.mkv?sign=cdn", http.StatusFound)
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
link, err := p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if !probeSeen {
t.Fatal("expected OpenList-hosted raw_url probe")
}
if link.URL != "https://cdn.example.test/movie.mkv?sign=cdn" || link.Proxy || len(link.Headers) != 0 {
t.Fatalf("link = %#v, want collapsed CDN 302 playback", link)
}
}
func TestOpenListResolveLogsInWithUsernamePasswordForAPIRawURL(t *testing.T) {
var loginSeen bool
var gotAuth string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/api/auth/login":
loginSeen = true
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode login body: %v", err)
}
if body["username"] != "alice" || body["password"] != "secret" {
t.Fatalf("login body = %#v", body)
}
_, _ = w.Write([]byte(`{"code":200,"data":{"token":"api-token"}}`))
case "/api/fs/get":
gotAuth = r.Header.Get("Authorization")
_, _ = w.Write([]byte(`{"code":200,"data":{"raw_url":"https://cdn.example.test/movie.mkv?sign=1"}}`))
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "username": "alice", "password": "secret"}, srv.Client())
if err != nil {
t.Fatal(err)
}
link, err := p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if !loginSeen {
t.Fatalf("expected api login before fs/get")
}
if gotAuth != "api-token" {
t.Fatalf("Authorization = %q, want api-token", gotAuth)
}
if link.URL != "https://cdn.example.test/movie.mkv?sign=1" || link.Proxy {
t.Fatalf("link = %#v, want raw_url 302 playback", link)
}
}
func TestOpenListResolveRejectsProxyWhenAPIRawURLNeedsHeaders(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/api/fs/get" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":200,"data":{"raw_url":"/dav/Cloud/Movie.mkv","header":{"Cookie":"sid=abc"}}}`))
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err == nil || !strings.Contains(err.Error(), "refusing WebDAV/proxy fallback") || !strings.Contains(err.Error(), "Cookie") {
t.Fatalf("resolve error = %v, want pure 302 refusal with header names", err)
}
}
func TestOpenListResolveRejectsHostedRawURLWithoutCDNRedirect(t *testing.T) {
var probeSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/fs/get":
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":200,"data":{"raw_url":"/d/Cloud/Movie.mkv?sign=1"}}`))
case "/d/Cloud/Movie.mkv":
probeSeen = true
w.Header().Set("Content-Range", "bytes 0-0/10")
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write([]byte("x"))
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err == nil || !strings.Contains(err.Error(), "OpenList-hosted raw_url") || !strings.Contains(err.Error(), "no CDN Location") {
t.Fatalf("resolve error = %v, want hosted raw_url refusal", err)
}
if !probeSeen {
t.Fatal("expected OpenList-hosted raw_url probe")
}
}
func TestOpenListResolveDoesNotFallbackToWebDAVWhenAPIRawURLFails(t *testing.T) {
var davSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/fs/get":
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":500,"message":"driver cannot provide raw_url"}`))
case "/dav/Cloud/Movie.mkv":
davSeen = true
w.WriteHeader(http.StatusOK)
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "/Cloud/Movie.mkv")
if err == nil || !strings.Contains(err.Error(), "pure 302 playback requires OpenList raw_url") {
t.Fatalf("resolve error = %v, want raw_url requirement", err)
}
if davSeen {
t.Fatal("openlist video resolve fell back to WebDAV after raw_url failure")
}
}
-237
View File
@@ -1,237 +0,0 @@
package cloud
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
)
// pan115Provider implements 115 网盘 via cookie auth.
//
// 115 has removed its desktop clients, so cookies must come from the mobile
// app / web (115.com) or a QR-code login (see QR* helpers below). Directory
// listing uses the public web API; download resolves a file's pickcode to a
// CDN URL that, like Alist's default 115 behaviour, is served by 302 redirect.
type pan115Provider struct {
cookie string
ua string
webBase string // https://webapi.115.com (override in tests)
proBase string // https://proapi.115.com (override in tests)
client *http.Client
proxy bool
// downURLPayload fetches and decrypts the app/chrome/downurl response for a
// pickcode, returning the raw JSON payload (map of file id -> info). It is a
// seam so tests can bypass the live 115 crypto/transport.
downURLPayload func(ctx context.Context, pickcode string) ([]byte, error)
}
const (
pan115WebBase = "https://webapi.115.com"
pan115ProBase = "https://proapi.115.com"
)
func new115(cfg map[string]any, client *http.Client) *pan115Provider {
web := str(cfg["base"])
if web == "" {
web = pan115WebBase
}
ua := str(cfg["ua"])
if ua == "" {
ua = defaultUA
}
// 115 CDN download URLs work with a plain 302 (Alist's recommended mode),
// so offload by default. The global cloud playback setting decides whether
// clients receive a STRMURL entry or a /Videos stream entry.
proxy := false
pro := str(cfg["pro_base"])
if pro == "" {
pro = pan115ProBase
}
p := &pan115Provider{
cookie: str(cfg["cookie"]),
ua: ua,
webBase: strings.TrimRight(web, "/"),
proBase: strings.TrimRight(pro, "/"),
client: client,
proxy: proxy,
}
p.downURLPayload = p.fetchDownURLPayload
return p
}
func (p *pan115Provider) Type() string { return Type115 }
func (p *pan115Provider) get(ctx context.Context, u string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
if err != nil {
return nil, err
}
req.Header.Set("Cookie", p.cookie)
req.Header.Set("User-Agent", p.ua)
req.Header.Set("Accept", "application/json, text/plain, */*")
return p.client.Do(req)
}
func (p *pan115Provider) Ping(ctx context.Context) error {
if p.cookie == "" {
return fmt.Errorf("115: missing cookie")
}
_, err := p.List(ctx, "0")
return err
}
func (p *pan115Provider) List(ctx context.Context, dirID string) ([]FileEntry, error) {
if dirID == "" {
dirID = "0"
}
const pageSize = 100
out := make([]FileEntry, 0, pageSize)
for offset := 0; ; offset += pageSize {
q := url.Values{}
q.Set("aid", "1")
q.Set("cid", dirID)
q.Set("o", "user_ptime")
q.Set("asc", "0")
q.Set("offset", strconv.Itoa(offset))
q.Set("show_dir", "1")
q.Set("limit", strconv.Itoa(pageSize))
q.Set("format", "json")
resp, err := p.get(ctx, p.webBase+"/files?"+q.Encode())
if err != nil {
return nil, err
}
var r struct {
State bool `json:"state"`
Error string `json:"error"`
Data []struct {
Fid string `json:"fid"` // file id (files only)
Cid string `json:"cid"` // category id (dirs use this)
N string `json:"n"` // name
S json.Number `json:"s"` // size
Pc string `json:"pc"` // pickcode
} `json:"data"`
}
err = json.NewDecoder(resp.Body).Decode(&r)
_ = resp.Body.Close()
if err != nil {
return nil, fmt.Errorf("115: decode list: %w", err)
}
if !r.State {
return nil, fmt.Errorf("115: list failed: %s", r.Error)
}
for _, it := range r.Data {
isDir := it.Fid == ""
id := it.Fid
if isDir {
id = it.Cid
}
size, _ := it.S.Int64()
out = append(out, FileEntry{
ID: id,
Name: it.N,
IsDir: isDir,
Size: size,
PickCode: it.Pc,
})
}
if len(r.Data) < pageSize {
break
}
}
return out, nil
}
// Resolve accepts a pickcode (preferred) and returns the CDN download URL.
//
// 115 deprecated the plain web /files/download endpoint (it no longer returns
// file_url for ordinary cookies). We use the current app/chrome/downurl
// endpoint, which takes an m115-encrypted body and returns an m115-encrypted
// payload mapping the file id to a short-lived, OSS-signed CDN URL suitable for
// a 302 redirect (the same approach Alist's 115 driver uses).
func (p *pan115Provider) Resolve(ctx context.Context, pickcode string) (*DirectLink, error) {
if pickcode == "" {
return nil, fmt.Errorf("115: empty pickcode")
}
raw, err := p.downURLPayload(ctx, pickcode)
if err != nil {
return nil, err
}
var payload map[string]struct {
FileName string `json:"file_name"`
FileSize json.Number `json:"file_size"`
URL struct {
URL string `json:"url"`
} `json:"url"`
}
if err := json.Unmarshal(raw, &payload); err != nil {
return nil, fmt.Errorf("115: decode downurl: %w", err)
}
for _, info := range payload {
if info.URL.URL == "" {
continue
}
return &DirectLink{
URL: info.URL.URL,
Headers: map[string]string{
"User-Agent": p.ua,
"Cookie": p.cookie,
},
Proxy: p.proxy,
}, nil
}
return nil, fmt.Errorf("115: download failed: no url")
}
// fetchDownURLPayload performs the live encrypted app/chrome/downurl request and
// returns the decrypted JSON payload.
func (p *pan115Provider) fetchDownURLPayload(ctx context.Context, pickcode string) ([]byte, error) {
key := m115GenerateKey()
params, err := json.Marshal(map[string]string{"pickcode": pickcode})
if err != nil {
return nil, err
}
form := url.Values{}
form.Set("data", m115Encode(params, key))
u := fmt.Sprintf("%s/app/chrome/downurl?t=%d", p.proBase, nowUnix())
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, strings.NewReader(form.Encode()))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Cookie", p.cookie)
req.Header.Set("User-Agent", p.ua)
req.Header.Set("Accept", "application/json, text/plain, */*")
resp, err := p.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
var r struct {
State bool `json:"state"`
Error string `json:"error"`
Data string `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&r); err != nil {
return nil, fmt.Errorf("115: decode downurl: %w", err)
}
if !r.State || r.Data == "" {
msg := r.Error
if msg == "" {
msg = "no data"
}
return nil, fmt.Errorf("115: download failed: %s", msg)
}
out, err := m115Decode(r.Data, key)
if err != nil {
return nil, fmt.Errorf("115: decrypt downurl: %w", err)
}
return out, nil
}
// nowUnix is a seam for deterministic tests.
var nowUnix = func() int64 { return timeNow().Unix() }
-184
View File
@@ -1,184 +0,0 @@
package cloud
// 115 网盘的 app/chrome/downurl 下载接口要求请求体使用 115 自有的 "m115" 加密协议
// (RSA + XOR 混淆),返回的直链也以同样方式加密。普通 web cookie 调用旧的
// /files/download 接口已不再返回 file_url,必须改用该加密接口。
//
// 下面的实现移植自 MIT 许可的 github.com/SheltonZhu/115driver
// (pkg/crypto/m115),alist 等项目同样采用此实现。仅做最小改动:函数加 m115
// 前缀以归入本包命名空间。
//
// Copyright (c) 115driver authors. MIT License.
import (
"bytes"
"crypto/rand"
"encoding/base64"
"io"
"math/big"
)
// m115Key is the random 16-byte session key generated per request.
type m115Key [16]byte
func m115GenerateKey() m115Key {
key := m115Key{}
_, _ = io.ReadFull(rand.Reader, key[:])
return key
}
// m115Encode encrypts request input for the downurl endpoint.
func m115Encode(input []byte, key m115Key) string {
buf := make([]byte, 16+len(input))
copy(buf, key[:])
copy(buf[16:], input)
m115XORTransform(buf[16:], m115XORDeriveKey(key[:], 4))
m115ReverseBytes(buf[16:])
m115XORTransform(buf[16:], m115XORClientKey)
return base64.StdEncoding.EncodeToString(m115RSAEncrypt(buf))
}
// m115Decode decrypts the base64 response payload using the request key.
func m115Decode(input string, key m115Key) ([]byte, error) {
data, err := base64.StdEncoding.DecodeString(input)
if err != nil {
return nil, err
}
data = m115RSADecrypt(data)
output := make([]byte, len(data)-16)
copy(output, data[16:])
m115XORTransform(output, m115XORDeriveKey(data[:16], 12))
m115ReverseBytes(output)
m115XORTransform(output, m115XORDeriveKey(key[:], 4))
return output, nil
}
func m115ReverseBytes(data []byte) {
for i, j := 0, len(data)-1; i < j; i, j = i+1, j-1 {
data[i], data[j] = data[j], data[i]
}
}
// --- RSA layer ---
var (
m115N, _ = big.NewInt(0).SetString(
"8686980c0f5a24c4b9d43020cd2c22703ff3f450756529058b1cf88f09b86021"+
"36477198a6e2683149659bd122c33592fdb5ad47944ad1ea4d36c6b172aad633"+
"8c3bb6ac6227502d010993ac967d1aef00f0c8e038de2e4d3bc2ec368af2e9f1"+
"0a6f1eda4f7262f136420c07c331b871bf139f74f3010e3c4fe57df3afb71683", 16)
m115E, _ = big.NewInt(0).SetString("10001", 16)
m115KeyLength = m115N.BitLen() / 8
)
func m115RSAEncrypt(input []byte) []byte {
buf := &bytes.Buffer{}
for remainSize := len(input); remainSize > 0; {
sliceSize := m115KeyLength - 11
if sliceSize > remainSize {
sliceSize = remainSize
}
m115RSAEncryptSlice(input[:sliceSize], buf)
input = input[sliceSize:]
remainSize -= sliceSize
}
return buf.Bytes()
}
func m115RSAEncryptSlice(input []byte, w io.Writer) {
padSize := m115KeyLength - len(input) - 3
padData := make([]byte, padSize)
_, _ = rand.Read(padData)
buf := make([]byte, m115KeyLength)
buf[0], buf[1] = 0, 2
for i, b := range padData {
buf[2+i] = b%0xff + 0x01
}
buf[padSize+2] = 0
copy(buf[padSize+3:], input)
msg := big.NewInt(0).SetBytes(buf)
ret := big.NewInt(0).Exp(msg, m115E, m115N).Bytes()
if fillSize := m115KeyLength - len(ret); fillSize > 0 {
zeros := make([]byte, fillSize)
_, _ = w.Write(zeros)
}
_, _ = w.Write(ret)
}
func m115RSADecrypt(input []byte) []byte {
buf := &bytes.Buffer{}
for remainSize := len(input); remainSize > 0; {
sliceSize := m115KeyLength
if sliceSize > remainSize {
sliceSize = remainSize
}
m115RSADecryptSlice(input[:sliceSize], buf)
input = input[sliceSize:]
remainSize -= sliceSize
}
return buf.Bytes()
}
func m115RSADecryptSlice(input []byte, w io.Writer) {
msg := big.NewInt(0).SetBytes(input)
ret := big.NewInt(0).Exp(msg, m115E, m115N).Bytes()
for i, b := range ret {
if b == 0 && i != 0 {
_, _ = w.Write(ret[i+1:])
break
}
}
}
// --- XOR layer ---
var (
m115XORKeySeed = []byte{
0xf0, 0xe5, 0x69, 0xae, 0xbf, 0xdc, 0xbf, 0x8a,
0x1a, 0x45, 0xe8, 0xbe, 0x7d, 0xa6, 0x73, 0xb8,
0xde, 0x8f, 0xe7, 0xc4, 0x45, 0xda, 0x86, 0xc4,
0x9b, 0x64, 0x8b, 0x14, 0x6a, 0xb4, 0xf1, 0xaa,
0x38, 0x01, 0x35, 0x9e, 0x26, 0x69, 0x2c, 0x86,
0x00, 0x6b, 0x4f, 0xa5, 0x36, 0x34, 0x62, 0xa6,
0x2a, 0x96, 0x68, 0x18, 0xf2, 0x4a, 0xfd, 0xbd,
0x6b, 0x97, 0x8f, 0x4d, 0x8f, 0x89, 0x13, 0xb7,
0x6c, 0x8e, 0x93, 0xed, 0x0e, 0x0d, 0x48, 0x3e,
0xd7, 0x2f, 0x88, 0xd8, 0xfe, 0xfe, 0x7e, 0x86,
0x50, 0x95, 0x4f, 0xd1, 0xeb, 0x83, 0x26, 0x34,
0xdb, 0x66, 0x7b, 0x9c, 0x7e, 0x9d, 0x7a, 0x81,
0x32, 0xea, 0xb6, 0x33, 0xde, 0x3a, 0xa9, 0x59,
0x34, 0x66, 0x3b, 0xaa, 0xba, 0x81, 0x60, 0x48,
0xb9, 0xd5, 0x81, 0x9c, 0xf8, 0x6c, 0x84, 0x77,
0xff, 0x54, 0x78, 0x26, 0x5f, 0xbe, 0xe8, 0x1e,
0x36, 0x9f, 0x34, 0x80, 0x5c, 0x45, 0x2c, 0x9b,
0x76, 0xd5, 0x1b, 0x8f, 0xcc, 0xc3, 0xb8, 0xf5,
}
m115XORClientKey = []byte{
0x78, 0x06, 0xad, 0x4c, 0x33, 0x86, 0x5d, 0x18,
0x4c, 0x01, 0x3f, 0x46,
}
)
func m115XORDeriveKey(seed []byte, size int) []byte {
key := make([]byte, size)
for i := 0; i < size; i++ {
key[i] = (seed[i] + m115XORKeySeed[size*i]) & 0xff
key[i] ^= m115XORKeySeed[size*(size-i-1)]
}
return key
}
func m115XORTransform(data []byte, key []byte) {
dataSize, keySize := len(data), len(key)
mod := dataSize % 4
if mod > 0 {
for i := 0; i < mod; i++ {
data[i] ^= key[i%keySize]
}
}
for i := mod; i < dataSize; i++ {
data[i] ^= key[(i-mod)%keySize]
}
}
-147
View File
@@ -1,147 +0,0 @@
package cloud
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
)
// QRSession is the handle returned by QRStart; the client renders QRImageURL
// and polls QRPoll until it returns a cookie.
type QRSession struct {
UID string `json:"uid"`
Time int64 `json:"time"`
Sign string `json:"sign"`
QRImageURL string `json:"qr_image_url"`
}
// QR login hosts (overridable for tests).
var (
qr115APIBase = "https://qrcodeapi.115.com"
qr115PassportBase = "https://passportapi.115.com"
)
// QRStart obtains a 115 QR-code login token + image URL.
func QRStart(ctx context.Context, client *http.Client) (*QRSession, error) {
if client == nil {
client = http.DefaultClient
}
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, qr115APIBase+"/api/1.0/web/1.0/token/", nil)
req.Header.Set("User-Agent", defaultUA)
resp, err := client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
var r struct {
State int `json:"state"`
Data struct {
UID string `json:"uid"`
Time int64 `json:"time"`
Sign string `json:"sign"`
} `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&r); err != nil {
return nil, fmt.Errorf("115 qr: decode token: %w", err)
}
if r.State != 1 || r.Data.UID == "" {
return nil, fmt.Errorf("115 qr: token request failed")
}
return &QRSession{
UID: r.Data.UID,
Time: r.Data.Time,
Sign: r.Data.Sign,
QRImageURL: qr115APIBase + "/api/1.0/web/1.0/qrcode?uid=" + url.QueryEscape(r.Data.UID),
}, nil
}
// QRStatus is the poll result.
type QRStatus struct {
// State is one of: "waiting" (not scanned), "scanned" (scanned, awaiting
// confirmation), "confirmed" (login approved; Cookie populated),
// "expired" (token expired/cancelled).
State string `json:"state"`
Cookie string `json:"cookie,omitempty"`
}
// QRPoll checks the QR session status; on confirmation it exchanges the token
// for a session cookie via the passport API.
func QRPoll(ctx context.Context, client *http.Client, sess *QRSession) (*QRStatus, error) {
if client == nil {
client = http.DefaultClient
}
if sess == nil || sess.UID == "" {
return nil, fmt.Errorf("115 qr: nil session")
}
q := url.Values{}
q.Set("uid", sess.UID)
q.Set("time", strconv.FormatInt(sess.Time, 10))
q.Set("sign", sess.Sign)
q.Set("_", strconv.FormatInt(timeNow().UnixMilli(), 10))
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, qr115APIBase+"/get/status/?"+q.Encode(), nil)
req.Header.Set("User-Agent", defaultUA)
resp, err := client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
var r struct {
State int `json:"state"`
Data struct {
Status int `json:"status"` // 0 waiting, 1 scanned, 2 confirmed, -1/-2 expired
} `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&r); err != nil {
return nil, fmt.Errorf("115 qr: decode status: %w", err)
}
switch r.Data.Status {
case 1:
return &QRStatus{State: "scanned"}, nil
case 2:
cookie, err := qr115Exchange(ctx, client, sess.UID)
if err != nil {
return nil, err
}
return &QRStatus{State: "confirmed", Cookie: cookie}, nil
case 0:
return &QRStatus{State: "waiting"}, nil
default:
return &QRStatus{State: "expired"}, nil
}
}
// qr115Exchange swaps an approved uid for a session cookie.
func qr115Exchange(ctx context.Context, client *http.Client, uid string) (string, error) {
form := url.Values{}
form.Set("account", uid)
form.Set("app", "web")
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, qr115PassportBase+"/app/1.0/web/1.0/login/qrcode/", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", defaultUA)
resp, err := client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
var r struct {
State int `json:"state"`
Data struct {
Cookie map[string]string `json:"cookie"`
} `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&r); err != nil {
return "", fmt.Errorf("115 qr: decode login: %w", err)
}
if r.State != 1 || len(r.Data.Cookie) == 0 {
return "", fmt.Errorf("115 qr: login exchange failed")
}
parts := make([]string, 0, len(r.Data.Cookie))
for k, v := range r.Data.Cookie {
parts = append(parts, k+"="+v)
}
return strings.Join(parts, "; "), nil
}
-288
View File
@@ -1,288 +0,0 @@
package service
import (
"context"
"net/url"
"strings"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
const cloudAutoCategoryQueryKey = "auto_category"
func BuildCloudAutoCategoryLibraryPath(provider, displayDir string) string {
return BuildCloudAutoCategoryLibraryPathWithScanDir(provider, "", displayDir)
}
func BuildCloudAutoCategoryLibraryPathWithScanDir(provider, scanDir, displayDir string) string {
base := BuildCloudLibraryPath(provider, scanDir, displayDir)
if base == "" || strings.TrimSpace(displayDir) == "" {
return ""
}
sep := "?"
if strings.Contains(base, "?") {
sep = "&"
}
return base + sep + cloudAutoCategoryQueryKey + "=1"
}
func CloudLibraryAutoCategory(lib model.Library) bool {
u, err := url.Parse(strings.TrimSpace(lib.Path))
if err != nil || strings.ToLower(u.Scheme) != "cloud" {
return false
}
switch strings.ToLower(strings.TrimSpace(u.Query().Get(cloudAutoCategoryQueryKey))) {
case "1", "true", "yes", "on":
return true
default:
return false
}
}
func cloudRootMountNeedsAutoCategory(mount CloudMountInfo) bool {
return strings.TrimSpace(mount.DisplayDir) == "" && strings.TrimSpace(mount.ScanDir) == ""
}
func cloudAutoCategoryDisplayDirForMediaPath(path string) string {
displayDir, _ := cloudAutoCategoryDirsForMediaPath(path)
return displayDir
}
func cloudAutoCategoryDirsForMediaPath(path string) (string, string) {
info, ok := ParseCloudLibraryMount(path)
if !ok {
return "", ""
}
parts := strmSlashParts(info.DisplayDir)
if len(parts) <= 1 {
return "", ""
}
parts = parts[:len(parts)-1]
categoryParts, scanParts := cloudAutoCategoryParts(parts)
if len(categoryParts) == 0 {
return "", ""
}
return strings.Join(categoryParts, "/"), strings.Join(scanParts, "/")
}
func cloudAutoCategoryParts(parts []string) ([]string, []string) {
for i, part := range parts {
root := strmCanonicalRoot(part)
if root != "" {
if i+1 >= len(parts) {
return nil, nil
}
category := strings.TrimSpace(parts[i+1])
if cloudAutoCategoryRootMatches(root, category) {
return []string{root, strmCanonicalCategory(category)}, append([]string(nil), parts[:i+2]...)
}
return nil, nil
}
if root := strmCategoryRoot(part); root != "" {
return []string{root, strmCanonicalCategory(part)}, append([]string(nil), parts[:i+1]...)
}
}
return nil, nil
}
func cloudAutoCategoryRootMatches(root, category string) bool {
category = strings.TrimSpace(category)
if category == "" {
return false
}
if strmCategoryRoot(category) == root {
return true
}
if root == "电影" {
return containsAnyText(strings.ToLower(category), "纪录片", "纪录", "documentary")
}
return false
}
type cloudAutoCategoryTarget struct {
Library *model.Library
RootID string
}
func (s *ScannerService) ensureCloudAutoCategoryTarget(ctx context.Context, rootLib *model.Library, provider, displayDir, scanDir string) (cloudAutoCategoryTarget, error) {
displayDir = normalizeCloudMountDir(provider, displayDir)
scanDir = normalizeCloudMountDir(provider, firstNonEmpty(scanDir, displayDir))
if s == nil || s.repo == nil || s.repo.DB == nil || rootLib == nil || provider == "" || displayDir == "" {
return cloudAutoCategoryTarget{Library: rootLib}, nil
}
path := BuildCloudAutoCategoryLibraryPathWithScanDir(provider, scanDir, displayDir)
if path == "" {
return cloudAutoCategoryTarget{Library: rootLib}, nil
}
name := cloudMountDirBase(displayDir)
if name == "" {
name = displayDir
}
kind := InferCloudMountMediaType(displayDir, name)
target, existingAuto := s.findCloudAutoCategoryTarget(ctx, rootLib.ID, provider, displayDir, name, kind)
if target != nil {
root, err := s.ensureCloudLibraryRoot(ctx, target.ID, name, path)
if err != nil {
return cloudAutoCategoryTarget{}, err
}
if existingAuto != nil && existingAuto.ID != target.ID {
s.migrateCloudAutoCategoryLibrary(ctx, existingAuto, target, root)
}
return cloudAutoCategoryTarget{Library: target, RootID: libraryRootID(root)}, nil
}
if existingAuto != nil {
root, err := s.ensureCloudLibraryRoot(ctx, existingAuto.ID, name, path)
if err != nil {
return cloudAutoCategoryTarget{}, err
}
return cloudAutoCategoryTarget{Library: existingAuto, RootID: libraryRootID(root)}, nil
}
lib := &model.Library{
Name: name,
Path: path,
Type: kind,
Enabled: true,
}
root := model.LibraryRoot{Name: name, Path: path, Enabled: true}
if err := s.repo.Library.CreateWithRoots(ctx, lib, []model.LibraryRoot{root}); err != nil {
_, existing := s.findCloudAutoCategoryTarget(ctx, rootLib.ID, provider, displayDir, name, kind)
if existing != nil {
ensuredRoot, rootErr := s.ensureCloudLibraryRoot(ctx, existing.ID, name, path)
if rootErr != nil {
return cloudAutoCategoryTarget{}, rootErr
}
return cloudAutoCategoryTarget{Library: existing, RootID: libraryRootID(ensuredRoot)}, nil
}
return cloudAutoCategoryTarget{}, err
}
if s.log != nil {
s.log.Info("created cloud auto category library",
zap.String("root_library_id", rootLib.ID),
zap.String("library_id", lib.ID),
zap.String("provider", provider),
zap.String("display_dir", displayDir))
}
if len(lib.Roots) > 0 {
return cloudAutoCategoryTarget{Library: lib, RootID: lib.Roots[0].ID}, nil
}
return cloudAutoCategoryTarget{Library: lib}, nil
}
func (s *ScannerService) findCloudAutoCategoryTarget(ctx context.Context, rootLibraryID, provider, displayDir, name, kind string) (*model.Library, *model.Library) {
if s == nil || s.repo == nil || s.repo.Library == nil {
return nil, nil
}
libs, err := s.repo.Library.List(ctx)
if err != nil {
if s.log != nil {
s.log.Warn("list libraries for cloud auto category failed", zap.Error(err))
}
return nil, nil
}
displayDir = normalizeCloudMountDir(provider, displayDir)
targetKey, _ := CloudLibraryMergeKey(model.Library{Name: name, Type: kind})
var target *model.Library
var existingAuto *model.Library
for _, lib := range libs {
info, ok := ParseCloudLibraryMount(lib.Path)
if ok && info.Provider == provider && normalizeCloudMountDir(provider, info.DisplayDir) == displayDir && CloudLibraryAutoCategory(lib) {
copy := lib
existingAuto = &copy
continue
}
if lib.ID == rootLibraryID || CloudLibraryAutoCategory(lib) || !lib.Enabled || targetKey == "" {
continue
}
key, ok := CloudLibraryMergeKey(lib)
if target == nil && ok && key == targetKey {
copy := lib
target = &copy
}
}
return target, existingAuto
}
func (s *ScannerService) ensureCloudLibraryRoot(ctx context.Context, libraryID, name, pathValue string) (*model.LibraryRoot, error) {
if s == nil || s.repo == nil || s.repo.Library == nil {
return nil, nil
}
roots, err := s.repo.Library.ListRoots(ctx, libraryID)
if err != nil {
return nil, err
}
targetKey := libraryRootPathKey(pathValue)
for i := range roots {
if libraryRootPathKey(roots[i].Path) == targetKey {
if strings.TrimSpace(roots[i].Name) == "" && strings.TrimSpace(name) != "" {
_ = s.repo.Library.UpdateRoot(ctx, &roots[i], map[string]any{"name": strings.TrimSpace(name)})
roots[i].Name = strings.TrimSpace(name)
}
return &roots[i], nil
}
}
root := &model.LibraryRoot{
LibraryID: libraryID,
Name: strings.TrimSpace(name),
Path: pathValue,
Enabled: true,
SortOrder: len(roots),
}
if err := s.repo.Library.CreateRoot(ctx, root); err != nil {
return nil, err
}
return root, nil
}
func (s *ScannerService) migrateCloudAutoCategoryLibrary(ctx context.Context, source, target *model.Library, root *model.LibraryRoot) {
if s == nil || s.repo == nil || s.repo.DB == nil || source == nil || target == nil || source.ID == "" || target.ID == "" {
return
}
updates := map[string]any{"library_id": target.ID}
if rootID := libraryRootID(root); rootID != "" {
updates["library_root_id"] = rootID
}
if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("library_id = ?", source.ID).Updates(updates).Error; err != nil {
if s.log != nil {
s.log.Warn("migrate cloud auto category media failed",
zap.String("from_library_id", source.ID),
zap.String("to_library_id", target.ID),
zap.Error(err))
}
return
}
_ = hardDeleteLibraryRoots(ctx, s.repo.DB, source.ID)
if err := s.repo.Library.Delete(ctx, source.ID); err != nil && s.log != nil {
s.log.Warn("remove migrated cloud auto category library failed",
zap.String("library_id", source.ID),
zap.Error(err))
}
}
func (s *ScannerService) cloudScanLibraryScopeIDs(ctx context.Context, lib *model.Library, mount CloudMountInfo) []string {
if lib == nil {
return nil
}
ids := []string{lib.ID}
if !cloudRootMountNeedsAutoCategory(mount) || s == nil || s.repo == nil || s.repo.Library == nil {
return ids
}
libs, err := s.repo.Library.List(ctx)
if err != nil {
if s.log != nil {
s.log.Warn("list libraries for cloud scan scope failed", zap.String("library_id", lib.ID), zap.Error(err))
}
return ids
}
for _, candidate := range libs {
if candidate.ID == lib.ID || !CloudLibraryAutoCategory(candidate) {
continue
}
info, ok := ParseCloudLibraryMount(candidate.Path)
if ok && info.Provider == mount.Provider {
ids = appendUniqueLibraryIDs(ids, candidate.ID)
}
}
return ids
}
@@ -0,0 +1,248 @@
package service
import (
"context"
"net/url"
"path"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// CloudMountInfo 是云盘挂载库的规范标识(因网盘后端已移除,仅保留类型以兼容
// 既有调用点;实际不会有 cloud:// 路径)。
type CloudMountInfo struct {
Provider string
DisplayDir string
ScanDir string
Path string
}
// ParseCloudLibraryMount 原用于解析 cloud:// 挂载库路径。网盘后端已移除,恒
// 返回 (空, false)。
func ParseCloudLibraryMount(_ string) (CloudMountInfo, bool) {
return CloudMountInfo{}, false
}
func cloudMountAncestor(_, _ string) bool {
return false
}
func cloudRootMountNeedsAutoCategory(_ CloudMountInfo) bool {
return false
}
func appendUniqueLibraryIDs(ids []string, values ...string) []string {
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
exists := false
for _, id := range ids {
if id == value {
exists = true
break
}
}
if !exists {
ids = append(ids, value)
}
}
return ids
}
func compactLibraryIDs(ids ...string) []string {
out := make([]string, 0, len(ids))
for _, id := range ids {
out = appendUniqueLibraryIDs(out, id)
}
return out
}
// 云盘库展示/合并辅助函数。
//
// 网盘后端(存储配置/云盘扫描/云播放)已随「存储配置」功能整体移除,因此
// 库表里不再存在 cloud:// 挂载库。以下函数保留为空实现以保持媒体浏览、
// 检索与 Emby 兼容层对原有调用点的兼容;在没有云盘库的前提下它们的
// 语义等价于「不合并、不过滤、不开自动分类」。
// FilterDisplayCloudLibraries 原用于筛掉云盘库(展示时不单独列出)。现已无
// 云盘库,原样返回。
func FilterDisplayCloudLibraries(_ context.Context, _ *repository.Container, libs []model.Library) []model.Library {
return libs
}
// MergedLibraryIDsForLibrary 原用于把合并展示的云盘库 ID 集合展开。现已无云盘
// 库,仅返回目标库自身 ID。
func MergedLibraryIDsForLibrary(_ context.Context, _ *repository.Container, libraryID string) ([]string, error) {
return []string{libraryID}, nil
}
// ExpandMediaVisibilityForMergedCloudLibraries 原用于把用户可见范围展开到合并
// 的云盘库。现已无云盘库,原样返回。
func ExpandMediaVisibilityForMergedCloudLibraries(_ context.Context, _ *repository.Container, visibility MediaVisibility) MediaVisibility {
return visibility
}
// CloudLibraryAutoCategory 原用于识别「自动分类」云盘库。现已无云盘库,恒为
// false。
func CloudLibraryAutoCategory(_ model.Library) bool {
return false
}
// CloudLibraryMergeKey 原用于计算两个云盘库的合并键。现已无云盘库,返回
// (空, false)。
func CloudLibraryMergeKey(_ model.Library) (string, bool) {
return "", false
}
// ShadowedCloudLibraryIDSet 原用于返回被合并/遮蔽的云盘库 ID 集合。现已无云盘
// 库,返回空集合。
func ShadowedCloudLibraryIDSet(_ []model.Library) map[string]bool {
return map[string]bool{}
}
// NormalizeCloudLibraryDisplay 原用于归一化云盘库的展示名/类型。现已无云盘库,
// 原样返回。
func NormalizeCloudLibraryDisplay(libs []model.Library) []model.Library {
return libs
}
// normalizeRemotePath 归一化远程(云盘/STRM 目标)路径。属通用路径处理辅助,
// 保留供 STRM 目标路径使用。
func normalizeRemotePath(p string) string {
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
if p == "" || p == "." {
return "/"
}
if !strings.HasPrefix(p, "/") {
p = "/" + p
}
return path.Clean(p)
}
// scanHasImportChanges 报告一次扫描是否产生了入库/变更。属通用扫描辅助。
func scanHasImportChanges(res *ScanResult) bool {
return res != nil && (res.Added > 0 || res.Updated > 0 || res.Removed > 0)
}
// cloneLocalMetadata 深拷贝 LocalMetadata 值(值类型浅拷贝即可)。属通用辅助。
func cloneLocalMetadata(src *LocalMetadata) *LocalMetadata {
if src == nil {
return nil
}
cp := *src
return &cp
}
// joinRemotePath 拼接远程路径片段(供 STRM 目标路径使用)。属通用路径辅助。
func joinRemotePath(base, rel string) string {
parts := []string{normalizeRemotePath(base)}
for _, part := range strings.Split(strings.ReplaceAll(rel, "\\", "/"), "/") {
part = strings.TrimSpace(part)
if part != "" && part != "." {
parts = append(parts, part)
}
}
return path.Clean(path.Join(parts...))
}
// 通用库路径辅助(原云盘库/STRM 生成逻辑使用;网盘后端移除后保留为纯工具函数,
// 供库路径构建与既有测试作为稳定夹具使用)。
const LegacyQuarkProvider = "quark"
func BuildCloudLibraryPath(provider, scanDir, displayDir string) string {
provider = strings.TrimSpace(provider)
scanDir = normalizeCloudMountDir(provider, scanDir)
displayDir = normalizeCloudMountDir(provider, firstNonEmpty(displayDir, scanDir))
if provider == "" {
return ""
}
base := "cloud://" + provider
if displayDir == "" {
if scanDir != "" {
return base + "?dir=" + url.QueryEscape(scanDir)
}
return base
}
pathStr := base + "/" + url.PathEscape(displayDir)
if scanDir != "" && scanDir != displayDir {
pathStr += "?dir=" + url.QueryEscape(scanDir)
}
return pathStr
}
func BuildCloudAutoCategoryLibraryPath(provider, displayDir string) string {
return BuildCloudAutoCategoryLibraryPathWithScanDir(provider, "", displayDir)
}
func BuildCloudAutoCategoryLibraryPathWithScanDir(provider, scanDir, displayDir string) string {
base := BuildCloudLibraryPath(provider, scanDir, displayDir)
if base == "" || strings.TrimSpace(displayDir) == "" {
return ""
}
sep := "?"
if strings.Contains(base, "?") {
sep = "&"
}
return base + sep + "auto_category=1"
}
func normalizeCloudMountDir(provider, value string) string {
value = strings.TrimSpace(value)
if decoded, err := url.PathUnescape(value); err == nil {
value = decoded
}
if decoded, err := url.QueryUnescape(value); err == nil {
value = decoded
}
value = strings.ReplaceAll(value, "\\", "/")
value = strings.Trim(strings.TrimSpace(value), "/")
if value == "." || ((provider == "115" || provider == LegacyQuarkProvider) && value == "0") {
return ""
}
return value
}
func cloudMountDirBase(dir string) string {
dir = strings.Trim(strings.TrimSpace(strings.ReplaceAll(dir, "\\", "/")), "/")
if dir == "" {
return ""
}
parts := strings.Split(dir, "/")
for i := len(parts) - 1; i >= 0; i-- {
if part := strings.TrimSpace(parts[i]); part != "" {
return part
}
}
return ""
}
func CloudMountProviderLabel(provider string) string {
switch strings.TrimSpace(provider) {
case LegacyQuarkProvider:
return "已停用网盘"
case "115":
return "115 网盘"
case "clouddrive2":
return "CloudDrive2"
case "openlist":
return "OpenList"
default:
if strings.TrimSpace(provider) == "" {
return "网盘"
}
return strings.TrimSpace(provider)
}
}
func CloudArtworkURL(typ, ref string) string {
typ = strings.Trim(strings.ReplaceAll(strings.TrimSpace(typ), "\\", "/"), "/")
ref = strings.TrimSpace(ref)
if typ == "" || ref == "" {
return ""
}
return "/api/img/cloud/" + url.PathEscape(typ) + "?ref=" + url.QueryEscape(ref)
}
-168
View File
@@ -1,168 +0,0 @@
package service
import (
"context"
"encoding/xml"
"path/filepath"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
type cloudSidecarSet struct {
nfoByName map[string]string
nfoByBase map[string]string
jsonByName map[string]string
jsonByBase map[string]string
imageByName map[string]string
imageByBase map[string]string
}
func newCloudSidecarSet(typ string, entries []cloud.FileEntry) cloudSidecarSet {
set := cloudSidecarSet{
nfoByName: make(map[string]string),
nfoByBase: make(map[string]string),
jsonByName: make(map[string]string),
jsonByBase: make(map[string]string),
imageByName: make(map[string]string),
imageByBase: make(map[string]string),
}
for _, entry := range entries {
if entry.IsDir {
continue
}
ref := cloudEntryRef(typ, entry.ID, entry.PickCode)
if ref == "" {
continue
}
name := strings.TrimSpace(entry.Name)
ext := strings.ToLower(filepath.Ext(name))
base := strings.ToLower(strings.TrimSpace(strings.TrimSuffix(name, ext)))
if name == "" || base == "" {
continue
}
switch ext {
case ".nfo":
set.nfoByName[strings.ToLower(name)] = ref
set.nfoByBase[base] = ref
case ".json":
set.jsonByName[strings.ToLower(name)] = ref
set.jsonByBase[base] = ref
case ".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn":
set.imageByName[strings.ToLower(name)] = ref
set.imageByBase[base] = ref
}
}
return set
}
func (s *ScannerService) cloudDirectoryMetadata(ctx context.Context, typ, displayDir string, sidecars cloudSidecarSet, inherited *LocalMetadata) *LocalMetadata {
meta := cloneLocalMetadata(inherited)
if hinted, _ := pathHintMetadata(displayDir, true); hinted != nil {
meta = mergeCloudMetadata(meta, hinted)
}
for _, name := range cloudShowNFOCandidates(displayDir) {
ref := sidecars.nfoByName[strings.ToLower(name)]
if ref == "" {
ref = sidecars.nfoByBase[strings.ToLower(strings.TrimSuffix(name, filepath.Ext(name)))]
}
if ref == "" {
continue
}
if local, doc, err := s.readCloudNFO(ctx, typ, ref, true); err == nil && local != nil {
local = applyCloudNFOArtwork(typ, sidecars, local, doc)
meta = mergeCloudMetadata(meta, local)
break
}
}
for _, name := range cloudDirectoryJSONCandidates(displayDir) {
ref := cloudJSONRefByName(sidecars, name)
if ref == "" {
continue
}
if local, err := s.readCloudJSONMetadata(ctx, typ, ref, sidecars); err == nil && local != nil {
meta = mergeCloudMetadata(meta, local)
break
}
}
meta = applyCloudDirectoryArtwork(typ, displayDir, sidecars, meta)
if !cloudMetadataUseful(meta) {
return nil
}
return meta
}
func (s *ScannerService) cloudFileMetadata(ctx context.Context, typ, displayPath, fileName string, sidecars cloudSidecarSet, inherited *LocalMetadata, seriesLike bool) *LocalMetadata {
season, episode := ParseEpisode(displayPath)
seriesLike = seriesLike || season > 0 || episode > 0
meta := cloneLocalMetadata(inherited)
if hinted, _ := pathHintMetadata(displayPath, seriesLike); hinted != nil {
meta = mergeCloudPathHintMetadata(meta, hinted)
}
base := strings.ToLower(strings.TrimSpace(strings.TrimSuffix(fileName, filepath.Ext(fileName))))
if ref := sidecars.nfoByBase[base]; ref != "" {
if local, doc, err := s.readCloudNFO(ctx, typ, ref, seriesLike); err == nil && local != nil {
local = applyCloudNFOArtwork(typ, sidecars, local, doc)
if seriesLike && doc != nil {
if meta == nil {
meta = &LocalMetadata{}
}
mergeEpisodeMetadata(meta, local, doc)
meta.HasNFO = true
} else {
meta = mergeCloudMetadata(meta, local)
}
}
}
for _, name := range cloudFileJSONCandidates(fileName, base) {
ref := cloudJSONRefByName(sidecars, name)
if ref == "" {
continue
}
if local, err := s.readCloudJSONMetadata(ctx, typ, ref, sidecars); err == nil && local != nil {
if cloudFileJSONIsEpisodeMetadata(seriesLike, season, episode, local) {
meta = mergeCloudEpisodeMetadata(meta, local)
} else {
meta = mergeCloudMetadata(meta, local)
}
break
}
}
meta = applyCloudFileArtwork(typ, sidecars, displayPath, fileName, base, meta)
if !cloudMetadataUseful(meta) {
return nil
}
return meta
}
func (s *ScannerService) readCloudJSONMetadata(ctx context.Context, typ, ref string, sidecars cloudSidecarSet) (*LocalMetadata, error) {
if s.storage == nil {
return nil, nil
}
body, err := s.storage.CloudReadText(ctx, typ, ref, 512<<10)
if err != nil {
return nil, err
}
meta, artwork := metadataFromCloudJSON([]byte(body))
if meta == nil {
return nil, nil
}
meta = applyCloudJSONArtwork(typ, sidecars, meta, artwork)
return meta, nil
}
func (s *ScannerService) readCloudNFO(ctx context.Context, typ, ref string, seriesLike bool) (*LocalMetadata, *nfoDocument, error) {
if s.storage == nil {
return nil, nil, nil
}
body, err := s.storage.CloudReadText(ctx, typ, ref, 512<<10)
if err != nil {
return nil, nil, err
}
var doc nfoDocument
if err := xml.Unmarshal([]byte(body), &doc); err != nil {
return nil, nil, err
}
meta := metadataFromDoc(&doc, "", seriesLike)
return meta, &doc, nil
}
-119
View File
@@ -1,119 +0,0 @@
package service
import (
"net/url"
"path"
"strings"
)
func applyCloudNFOArtwork(typ string, sidecars cloudSidecarSet, meta *LocalMetadata, doc *nfoDocument) *LocalMetadata {
if meta == nil {
meta = &LocalMetadata{}
}
if doc == nil {
return meta
}
if ref := cloudImageRefFromNFOValues(sidecars, nfoPosterValues(doc)...); ref != "" {
meta.PosterURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
if ref := cloudImageRefFromNFOValues(sidecars, nfoBackdropValues(doc)...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
return meta
}
func applyCloudDirectoryArtwork(typ, displayDir string, sidecars cloudSidecarSet, meta *LocalMetadata) *LocalMetadata {
if meta == nil {
meta = &LocalMetadata{}
}
if meta.PosterURL == "" {
if ref := firstCloudImageRef(sidecars, cloudPosterNameCandidates(cloudDirectoryArtworkBases(displayDir), "poster", "folder", "cover", "show", "tvshow")...); ref != "" {
meta.PosterURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
if meta.BackdropURL == "" {
if ref := firstCloudImageRef(sidecars, cloudBackdropNameCandidates(cloudDirectoryArtworkBases(displayDir), "fanart", "backdrop", "background", "landscape")...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
return meta
}
func applyCloudFileArtwork(typ string, sidecars cloudSidecarSet, displayPath, fileName, base string, meta *LocalMetadata) *LocalMetadata {
if meta == nil {
meta = &LocalMetadata{}
}
bases := cloudFileArtworkBases(displayPath, fileName, base)
if meta.PosterURL == "" {
if ref := firstCloudImageRef(sidecars, cloudPosterNameCandidates(bases, "poster", "folder", "cover", "movie", "show", "thumb")...); ref != "" {
meta.PosterURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
if meta.BackdropURL == "" {
if ref := firstCloudImageRef(sidecars, cloudBackdropNameCandidates(bases, "fanart", "backdrop", "background", "landscape")...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
return meta
}
func firstCloudImageRef(sidecars cloudSidecarSet, names ...string) string {
for _, name := range names {
if ref := cloudImageRefByName(sidecars, name); ref != "" {
return ref
}
}
return ""
}
func cloudImageRefFromNFOValues(sidecars cloudSidecarSet, values ...string) string {
for _, value := range values {
if ref := cloudImageRefByName(sidecars, value); ref != "" {
return ref
}
}
return ""
}
func cloudImageRefByName(sidecars cloudSidecarSet, value string) string {
name := normalizeCloudArtworkName(value)
if name == "" || isHTTPURL(name) {
return ""
}
if ref := sidecars.imageByName[strings.ToLower(name)]; ref != "" {
return ref
}
base := strings.TrimSuffix(name, path.Ext(name))
if ref := sidecars.imageByBase[strings.ToLower(base)]; ref != "" {
return ref
}
return ""
}
func normalizeCloudArtworkName(value string) string {
value = cleanXMLText(value)
if value == "" {
return ""
}
if isHTTPURL(value) {
return value
}
if unescaped, err := url.QueryUnescape(value); err == nil {
value = unescaped
}
value = strings.ReplaceAll(value, "\\", "/")
if idx := strings.IndexAny(value, "?#"); idx >= 0 {
value = value[:idx]
}
value = strings.Trim(strings.TrimSpace(value), "/")
if value == "" {
return ""
}
return path.Base(value)
}
@@ -1,122 +0,0 @@
package service
import (
"path/filepath"
"strconv"
"strings"
)
func cloudShowNFOCandidates(displayDir string) []string {
names := []string{"tvshow.nfo", "series.nfo", "show.nfo", "movie.nfo"}
base := strings.TrimSpace(pathBaseSlash(displayDir))
if base != "" {
names = append(names, base+".nfo")
}
return names
}
func cloudDirectoryJSONCandidates(displayDir string) []string {
names := []string{"movie.json", "metadata.json", "tvshow.json", "series.json", "show.json"}
base := strings.TrimSpace(pathBaseSlash(displayDir))
if base != "" {
names = append(names, base+".json", base+"-metadata.json", base+".metadata.json", base+"-mediainfo.json", base+".mediainfo.json")
}
return names
}
func cloudFileJSONCandidates(fileName, base string) []string {
if base == "" {
base = strings.ToLower(strings.TrimSpace(strings.TrimSuffix(fileName, filepath.Ext(fileName))))
}
cleanBases := cloudCleanArtworkBases(fileName)
bases := uniqueCloudArtworkNames(append([]string{base}, cleanBases...)...)
out := make([]string, 0, len(bases)*5+2)
for _, value := range bases {
out = append(out, value+".json", value+"-metadata.json", value+".metadata.json", value+"-mediainfo.json", value+".mediainfo.json")
}
return append(out, "movie.json", "metadata.json")
}
func cloudJSONRefByName(sidecars cloudSidecarSet, name string) string {
name = normalizeCloudArtworkName(name)
if name == "" || isHTTPURL(name) {
return ""
}
if ref := sidecars.jsonByName[strings.ToLower(name)]; ref != "" {
return ref
}
base := strings.TrimSuffix(name, filepath.Ext(name))
return sidecars.jsonByBase[strings.ToLower(base)]
}
func cloudFileArtworkBases(displayPath, fileName, base string) []string {
return uniqueCloudArtworkNames(append(
[]string{base},
append(cloudCleanArtworkBases(fileName), cloudDirectoryArtworkBases(pathDirSlash(displayPath))...)...,
)...)
}
func cloudDirectoryArtworkBases(displayDir string) []string {
base := pathBaseSlash(displayDir)
return uniqueCloudArtworkNames(append([]string{base}, cloudCleanArtworkBases(base)...)...)
}
func cloudCleanArtworkBases(value string) []string {
title, year := CleanQuery(value)
title = strings.TrimSpace(title)
if title == "" {
return nil
}
out := []string{title}
if year > 0 {
yearText := strconv.Itoa(year)
out = append(out,
title+" ("+yearText+")",
title+"."+yearText,
title+" "+yearText,
)
}
return out
}
func cloudPosterNameCandidates(bases []string, fallback ...string) []string {
out := make([]string, 0, len(bases)*7+len(fallback))
for _, base := range bases {
base = strings.TrimSpace(base)
if base == "" {
continue
}
out = append(out, base, base+"-poster", base+".poster", base+"-cover", base+".cover", base+"-thumb", base+".thumb")
}
return append(out, fallback...)
}
func cloudBackdropNameCandidates(bases []string, fallback ...string) []string {
out := make([]string, 0, len(bases)*6+len(fallback))
for _, base := range bases {
base = strings.TrimSpace(base)
if base == "" {
continue
}
out = append(out, base+"-fanart", base+".fanart", base+"-backdrop", base+".backdrop", base+"-background", base+".background")
}
return append(out, fallback...)
}
func uniqueCloudArtworkNames(values ...string) []string {
out := make([]string, 0, len(values))
seen := map[string]struct{}{}
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
key := strings.ToLower(value)
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
out = append(out, value)
}
return out
}
@@ -1,83 +0,0 @@
package service
import "strings"
func cloudFileJSONIsEpisodeMetadata(seriesLike bool, parsedSeason, parsedEpisode int, meta *LocalMetadata) bool {
if !seriesLike || meta == nil {
return false
}
return parsedSeason > 0 ||
parsedEpisode > 0 ||
meta.SeasonNum > 0 ||
meta.EpisodeNum > 0 ||
strings.TrimSpace(meta.EpisodeTitle) != ""
}
func mergeCloudEpisodeMetadata(dst, episode *LocalMetadata) *LocalMetadata {
if episode == nil {
return dst
}
if dst == nil {
dst = &LocalMetadata{}
}
mergeCloudEpisodeIdentity(dst, episode)
mergeCloudEpisodeDisplay(dst, episode)
mergeCloudEpisodeNumbersAndTaxonomy(dst, episode)
dst.NSFW = dst.NSFW || episode.NSFW
dst.HasNFO = dst.HasNFO || episode.HasNFO
dst.HasArtwork = dst.HasArtwork || episode.HasArtwork
return dst
}
func mergeCloudEpisodeIdentity(dst, episode *LocalMetadata) {
showTitle := ""
if episode.EpisodeTitle != "" && episode.Title != "" && !strings.EqualFold(episode.Title, episode.EpisodeTitle) {
showTitle = episode.Title
}
if showTitle != "" {
dst.Title = showTitle
}
episodeTitle := strings.TrimSpace(episode.EpisodeTitle)
if episodeTitle == "" && (episode.SeasonNum > 0 || episode.EpisodeNum > 0) {
episodeTitle = strings.TrimSpace(episode.Title)
}
if episodeTitle != "" && !strings.EqualFold(episodeTitle, strings.TrimSpace(dst.Title)) {
dst.EpisodeTitle = episodeTitle
}
}
func mergeCloudEpisodeDisplay(dst, episode *LocalMetadata) {
if dst.Year == 0 && episode.Year > 0 {
dst.Year = episode.Year
}
if episode.Overview != "" {
dst.Overview = episode.Overview
}
if episode.Rating > 0 {
dst.Rating = episode.Rating
}
if episode.PosterURL != "" {
dst.PosterURL = episode.PosterURL
}
if episode.BackdropURL != "" {
dst.BackdropURL = episode.BackdropURL
}
}
func mergeCloudEpisodeNumbersAndTaxonomy(dst, episode *LocalMetadata) {
if episode.SeasonNum > 0 {
dst.SeasonNum = episode.SeasonNum
}
if episode.EpisodeNum > 0 {
dst.EpisodeNum = episode.EpisodeNum
}
if dst.Genres == "" && episode.Genres != "" {
dst.Genres = episode.Genres
}
if dst.Countries == "" && episode.Countries != "" {
dst.Countries = episode.Countries
}
if dst.Languages == "" && episode.Languages != "" {
dst.Languages = episode.Languages
}
}
-265
View File
@@ -1,265 +0,0 @@
package service
import (
"encoding/json"
"strconv"
"strings"
)
type cloudJSONArtwork struct {
posterValues []string
backdropValues []string
}
func metadataFromCloudJSON(body []byte) (*LocalMetadata, cloudJSONArtwork) {
var raw any
if err := json.Unmarshal(body, &raw); err != nil {
return nil, cloudJSONArtwork{}
}
obj := firstMetadataJSONObject(raw)
if len(obj) == 0 {
return nil, cloudJSONArtwork{}
}
meta := &LocalMetadata{
Title: firstJSONString(obj, "title", "name", "showtitle", "show_title"),
OriginalName: firstJSONString(obj, "original_title", "originaltitle", "original_name", "originalname", "sorttitle"),
EpisodeTitle: firstJSONString(obj, "episode_title", "episodetitle", "episode_name", "episodename"),
Year: firstJSONInt(obj, "year"),
ReleaseDate: normalizeReleaseDate(firstJSONString(obj, "release_date", "releasedate", "premiered", "aired", "date")),
Overview: firstJSONString(obj, "overview", "plot", "outline", "summary", "description"),
Rating: firstJSONFloat(obj, "rating", "vote_average", "score"),
TMDbID: firstJSONInt(obj, "tmdb_id", "tmdbid", "tmdb"),
BangumiID: firstJSONInt(obj, "bangumi_id", "bangumiid", "bgm_id"),
DoubanID: firstJSONString(obj, "douban_id", "doubanid"),
TheTVDBID: firstJSONString(obj, "thetvdb_id", "tvdb_id", "thetvdbid", "tvdbid"),
SeasonNum: firstJSONInt(obj, "season", "season_num", "season_number"),
EpisodeNum: firstJSONInt(obj, "episode", "episode_num", "episode_number"),
Genres: firstJSONList(obj, "genres", "genre", "tags"),
Countries: firstJSONList(obj, "countries", "country", "production_countries"),
Languages: firstJSONList(obj, "languages", "language", "spoken_languages"),
}
if showTitle := firstJSONString(obj, "showtitle", "show_title", "series_title", "series_name"); showTitle != "" {
if meta.EpisodeTitle == "" && meta.Title != "" && !strings.EqualFold(strings.TrimSpace(meta.Title), strings.TrimSpace(showTitle)) {
meta.EpisodeTitle = meta.Title
}
meta.Title = showTitle
}
if meta.Year == 0 {
meta.Year = yearFromDate(meta.ReleaseDate)
}
artwork := cloudJSONArtwork{
posterValues: firstJSONStrings(obj,
"poster_url", "poster", "poster_path", "cover", "cover_url", "thumb", "thumbnail", "image"),
backdropValues: firstJSONStrings(obj,
"backdrop_url", "backdrop", "backdrop_path", "fanart", "fanart_url", "background", "landscape"),
}
if images, ok := jsonObject(obj["images"]); ok {
artwork.posterValues = append(artwork.posterValues, firstJSONStrings(images, "poster", "large", "common", "medium", "small", "cover")...)
artwork.backdropValues = append(artwork.backdropValues, firstJSONStrings(images, "backdrop", "fanart", "background", "landscape")...)
}
if art, ok := jsonObject(obj["art"]); ok {
artwork.posterValues = append(artwork.posterValues, firstJSONStrings(art, "poster", "thumb", "cover")...)
artwork.backdropValues = append(artwork.backdropValues, firstJSONStrings(art, "fanart", "backdrop", "background", "landscape")...)
}
if len(artwork.posterValues) > 0 {
meta.PosterURL = firstHTTPJSONValue(artwork.posterValues)
}
if len(artwork.backdropValues) > 0 {
meta.BackdropURL = firstHTTPJSONValue(artwork.backdropValues)
}
if meta.PosterURL != "" || meta.BackdropURL != "" {
meta.HasArtwork = true
}
if localHasDescriptiveMetadata(meta) || meta.HasArtwork || len(artwork.posterValues) > 0 || len(artwork.backdropValues) > 0 {
meta.HasNFO = true
return meta, artwork
}
return nil, cloudJSONArtwork{}
}
func applyCloudJSONArtwork(typ string, sidecars cloudSidecarSet, meta *LocalMetadata, artwork cloudJSONArtwork) *LocalMetadata {
if meta == nil {
meta = &LocalMetadata{}
}
if meta.PosterURL == "" {
if ref := cloudImageRefFromNFOValues(sidecars, artwork.posterValues...); ref != "" {
meta.PosterURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
if meta.BackdropURL == "" {
if ref := cloudImageRefFromNFOValues(sidecars, artwork.backdropValues...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
}
return meta
}
func firstMetadataJSONObject(raw any) map[string]any {
obj, ok := jsonObject(raw)
if !ok {
return nil
}
for _, key := range []string{"movie", "media", "metadata", "item", "data"} {
if nested, ok := jsonObject(obj[key]); ok && jsonObjectLooksLikeMetadata(nested) {
return nested
}
}
return obj
}
func jsonObjectLooksLikeMetadata(obj map[string]any) bool {
for _, key := range []string{"title", "name", "overview", "plot", "tmdb_id", "tmdbid", "poster", "poster_url", "poster_path", "backdrop", "backdrop_path"} {
if _, ok := obj[key]; ok {
return true
}
}
return false
}
func jsonObject(raw any) (map[string]any, bool) {
obj, ok := raw.(map[string]any)
return obj, ok
}
func firstJSONString(obj map[string]any, keys ...string) string {
values := firstJSONStrings(obj, keys...)
if len(values) == 0 {
return ""
}
return values[0]
}
func firstJSONStrings(obj map[string]any, keys ...string) []string {
out := []string{}
for _, key := range keys {
value, ok := lookupJSONKey(obj, key)
if !ok {
continue
}
out = append(out, jsonStrings(value)...)
if len(out) > 0 {
return out
}
}
return out
}
func firstJSONInt(obj map[string]any, keys ...string) int {
for _, key := range keys {
value, ok := lookupJSONKey(obj, key)
if !ok {
continue
}
if i := jsonInt(value); i > 0 {
return i
}
}
return 0
}
func firstJSONFloat(obj map[string]any, keys ...string) float32 {
for _, key := range keys {
value, ok := lookupJSONKey(obj, key)
if !ok {
continue
}
if f := jsonFloat(value); f > 0 {
return f
}
}
return 0
}
func firstJSONList(obj map[string]any, keys ...string) string {
seen := map[string]struct{}{}
out := []string{}
for _, key := range keys {
value, ok := lookupJSONKey(obj, key)
if !ok {
continue
}
for _, part := range jsonStrings(value) {
for _, item := range strings.Split(part, ",") {
item = strings.TrimSpace(item)
if item == "" {
continue
}
dedupeKey := strings.ToLower(item)
if _, exists := seen[dedupeKey]; exists {
continue
}
seen[dedupeKey] = struct{}{}
out = append(out, item)
}
}
if len(out) > 0 {
return strings.Join(out, ",")
}
}
return ""
}
func lookupJSONKey(obj map[string]any, key string) (any, bool) {
for existing, value := range obj {
if strings.EqualFold(strings.TrimSpace(existing), key) {
return value, true
}
}
return nil, false
}
func jsonStrings(value any) []string {
switch v := value.(type) {
case string:
if text := strings.TrimSpace(v); text != "" {
return []string{text}
}
case []any:
out := make([]string, 0, len(v))
for _, item := range v {
out = append(out, jsonStrings(item)...)
}
return out
case map[string]any:
return firstJSONStrings(v, "name", "title", "value", "iso_3166_1", "iso_639_1")
case float64:
if v > 0 {
return []string{strconv.Itoa(int(v))}
}
}
return nil
}
func jsonInt(value any) int {
switch v := value.(type) {
case float64:
return int(v)
case string:
i, _ := strconv.Atoi(strings.TrimSpace(v))
return i
}
return 0
}
func jsonFloat(value any) float32 {
switch v := value.(type) {
case float64:
return float32(v)
case string:
f, _ := strconv.ParseFloat(strings.TrimSpace(v), 32)
return float32(f)
}
return 0
}
func firstHTTPJSONValue(values []string) string {
for _, value := range values {
value = strings.TrimSpace(value)
if isHTTPURL(value) {
return value
}
}
return ""
}
-159
View File
@@ -1,159 +0,0 @@
package service
import "strings"
func mergeCloudMetadata(dst, src *LocalMetadata) *LocalMetadata {
if src == nil {
return dst
}
if dst == nil {
return cloneLocalMetadata(src)
}
if src.Title != "" {
dst.Title = src.Title
}
if src.OriginalName != "" {
dst.OriginalName = src.OriginalName
}
if src.EpisodeTitle != "" {
dst.EpisodeTitle = src.EpisodeTitle
}
if src.AdultCode != "" {
dst.AdultCode = src.AdultCode
}
if src.Year > 0 {
dst.Year = src.Year
}
if src.ReleaseDate != "" {
dst.ReleaseDate = src.ReleaseDate
}
if src.Overview != "" {
dst.Overview = src.Overview
}
if src.Rating > 0 {
dst.Rating = src.Rating
}
if src.PosterURL != "" {
dst.PosterURL = src.PosterURL
}
if src.BackdropURL != "" {
dst.BackdropURL = src.BackdropURL
}
if src.TMDbID > 0 {
dst.TMDbID = src.TMDbID
}
if src.BangumiID > 0 {
dst.BangumiID = src.BangumiID
}
if src.DoubanID != "" {
dst.DoubanID = src.DoubanID
}
if src.TheTVDBID != "" {
dst.TheTVDBID = src.TheTVDBID
}
if src.SeasonNum > 0 || src.EpisodeNum > 0 {
dst.SeasonNum = src.SeasonNum
}
if src.EpisodeNum > 0 {
dst.EpisodeNum = src.EpisodeNum
}
if src.Genres != "" {
dst.Genres = src.Genres
}
if src.Countries != "" {
dst.Countries = src.Countries
}
if src.Languages != "" {
dst.Languages = src.Languages
}
dst.NSFW = dst.NSFW || src.NSFW
dst.HasNFO = dst.HasNFO || src.HasNFO
dst.HasArtwork = dst.HasArtwork || src.HasArtwork
dst.PathHint = dst.PathHint || src.PathHint
return dst
}
func mergeCloudPathHintMetadata(dst, hint *LocalMetadata) *LocalMetadata {
if hint == nil {
return dst
}
if dst == nil || !dst.HasNFO {
return mergeCloudMetadata(dst, hint)
}
if dst.Title == "" {
dst.Title = hint.Title
}
if dst.OriginalName == "" {
dst.OriginalName = hint.OriginalName
}
if dst.Year == 0 {
dst.Year = hint.Year
}
if dst.ReleaseDate == "" {
dst.ReleaseDate = hint.ReleaseDate
}
if dst.TMDbID == 0 {
dst.TMDbID = hint.TMDbID
}
if dst.BangumiID == 0 {
dst.BangumiID = hint.BangumiID
}
if dst.DoubanID == "" {
dst.DoubanID = hint.DoubanID
}
if dst.TheTVDBID == "" {
dst.TheTVDBID = hint.TheTVDBID
}
dst.PathHint = dst.PathHint || hint.PathHint
return dst
}
func cloneLocalMetadata(src *LocalMetadata) *LocalMetadata {
if src == nil {
return nil
}
cp := *src
return &cp
}
func cloudMetadataUseful(meta *LocalMetadata) bool {
return meta != nil && (meta.HasNFO || meta.HasArtwork || localHasDescriptiveMetadata(meta))
}
func cloudPlaybackURL(typ, ref string) string {
return CloudArtworkURL(typ, ref)
}
func joinCloudDisplayPath(parent, child string) string {
parent = strings.Trim(strings.ReplaceAll(strings.TrimSpace(parent), "\\", "/"), "/")
child = strings.Trim(strings.ReplaceAll(strings.TrimSpace(child), "\\", "/"), "/")
switch {
case parent == "":
return child
case child == "":
return parent
default:
return parent + "/" + child
}
}
func pathBaseSlash(value string) string {
value = strings.Trim(strings.ReplaceAll(strings.TrimSpace(value), "\\", "/"), "/")
if value == "" {
return ""
}
parts := strings.Split(value, "/")
return parts[len(parts)-1]
}
func pathDirSlash(value string) string {
value = strings.Trim(strings.ReplaceAll(strings.TrimSpace(value), "\\", "/"), "/")
if value == "" {
return ""
}
idx := strings.LastIndex(value, "/")
if idx < 0 {
return ""
}
return value[:idx]
}
-143
View File
@@ -1,143 +0,0 @@
package service
import (
"net/url"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// CloudMountInfo is the canonical identity of a mounted cloud library. ScanDir
// is the provider id/path used for listing. DisplayDir is a hierarchical path
// used to prevent mounting both a parent and its child as separate libraries.
type CloudMountInfo struct {
Provider string
DisplayDir string
ScanDir string
Path string
}
type CloudMountConflict struct {
Library model.Library `json:"library"`
Exact bool `json:"exact"`
Nested bool `json:"nested"`
ExistingIsAncestor bool `json:"existing_is_ancestor"`
}
func BuildCloudLibraryPath(provider, scanDir, displayDir string) string {
provider = strings.TrimSpace(provider)
scanDir = normalizeCloudMountDir(provider, scanDir)
displayDir = normalizeCloudMountDir(provider, firstNonEmpty(displayDir, scanDir))
if provider == "" {
return ""
}
base := "cloud://" + provider
if displayDir == "" {
if scanDir != "" {
return base + "?dir=" + url.QueryEscape(scanDir)
}
return base
}
path := base + "/" + url.PathEscape(displayDir)
if scanDir != "" && scanDir != displayDir {
path += "?dir=" + url.QueryEscape(scanDir)
}
return path
}
func ParseCloudLibraryMount(raw string) (CloudMountInfo, bool) {
raw = strings.TrimSpace(raw)
if !strings.HasPrefix(strings.ToLower(raw), "cloud://") {
return CloudMountInfo{}, false
}
u, err := url.Parse(raw)
if err != nil || strings.ToLower(u.Scheme) != "cloud" {
return CloudMountInfo{}, false
}
provider := strings.TrimSpace(u.Host)
if provider == "" {
return CloudMountInfo{}, false
}
displayDir := strings.Trim(strings.TrimSpace(u.Path), "/")
if decoded, err := url.PathUnescape(displayDir); err == nil {
displayDir = decoded
}
scanDir := displayDir
if qDir := strings.TrimSpace(u.Query().Get("dir")); qDir != "" {
if decoded, err := url.QueryUnescape(qDir); err == nil {
qDir = decoded
}
scanDir = qDir
}
displayDir = normalizeCloudMountDir(provider, displayDir)
scanDir = normalizeCloudMountDir(provider, scanDir)
return CloudMountInfo{
Provider: provider,
DisplayDir: displayDir,
ScanDir: scanDir,
Path: raw,
}, true
}
func FindCloudMountConflict(libs []model.Library, provider, scanDir, displayDir string) *CloudMountConflict {
candidate := CloudMountInfo{
Provider: strings.TrimSpace(provider),
DisplayDir: normalizeCloudMountDir(provider, firstNonEmpty(displayDir, scanDir)),
ScanDir: normalizeCloudMountDir(provider, scanDir),
}
for _, lib := range libs {
if CloudLibraryAutoCategory(lib) {
continue
}
existing, ok := ParseCloudLibraryMount(lib.Path)
if !ok || existing.Provider != candidate.Provider {
continue
}
if existing.DisplayDir == candidate.DisplayDir {
return &CloudMountConflict{Library: lib, Exact: true}
}
if existing.ScanDir != "" && candidate.ScanDir != "" && existing.ScanDir == candidate.ScanDir {
return &CloudMountConflict{Library: lib, Exact: true}
}
if cloudMountAncestor(candidate.DisplayDir, existing.DisplayDir) {
return &CloudMountConflict{Library: lib, Nested: true}
}
}
return nil
}
func CloudLibraryShadowed(libs []model.Library, lib model.Library) *CloudMountConflict {
current, ok := ParseCloudLibraryMount(lib.Path)
if !ok {
return nil
}
for _, existing := range libs {
if existing.ID == lib.ID || !existing.Enabled {
continue
}
if CloudLibraryAutoCategory(existing) {
continue
}
info, ok := ParseCloudLibraryMount(existing.Path)
if !ok || info.Provider != current.Provider {
continue
}
if info.DisplayDir == current.DisplayDir && existing.CreatedAt.Before(lib.CreatedAt) {
return &CloudMountConflict{Library: existing, Exact: true}
}
if cloudMountAncestor(current.DisplayDir, info.DisplayDir) {
return &CloudMountConflict{Library: existing, Nested: true}
}
}
return nil
}
func FilterShadowedCloudLibraries(libs []model.Library) []model.Library {
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if CloudLibraryShadowed(libs, lib) == nil {
out = append(out, lib)
}
}
return out
}
-38
View File
@@ -1,38 +0,0 @@
package service
import (
"context"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func cloudLibraryMediaCounts(ctx context.Context, repo *repository.Container, libs []model.Library) map[string]int64 {
counts := make(map[string]int64, len(libs))
if repo == nil || repo.DB == nil || len(libs) == 0 {
return counts
}
ids := make([]string, 0, len(libs))
for _, lib := range libs {
ids = append(ids, lib.ID)
}
if len(ids) == 0 {
return counts
}
var rows []struct {
LibraryID string
Count int64
}
if err := repo.DB.WithContext(ctx).
Model(&model.Media{}).
Select("library_id, COUNT(*) AS count").
Where("library_id IN ? AND deleted_at IS NULL", ids).
Group("library_id").
Scan(&rows).Error; err != nil {
return counts
}
for _, row := range rows {
counts[row.LibraryID] = row.Count
}
return counts
}
-122
View File
@@ -1,122 +0,0 @@
package service
import (
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func betterDisplayCloudLibrary(candidate, current model.Library, counts map[string]int64) bool {
candidateCount := counts[candidate.ID]
currentCount := counts[current.ID]
if (candidateCount > 0) != (currentCount > 0) {
return candidateCount > 0
}
if candidate.Enabled != current.Enabled {
return candidate.Enabled
}
candidateCanonical := cloudLibraryPathIsCanonical(candidate)
currentCanonical := cloudLibraryPathIsCanonical(current)
if candidateCanonical != currentCanonical {
return candidateCanonical
}
if !candidate.CreatedAt.Equal(current.CreatedAt) {
return candidate.CreatedAt.After(current.CreatedAt)
}
return candidate.ID > current.ID
}
func mergeDisplayCloudLibraries(libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
localByKey := make(map[string]struct{}, len(libs))
for _, lib := range libs {
if _, ok := ParseCloudLibraryMount(lib.Path); ok || !lib.Enabled {
continue
}
if key, ok := CloudLibraryMergeKey(lib); ok {
localByKey[key] = struct{}{}
}
}
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if displayName, ok := CloudLibraryDisplayName(lib); ok && displayName != "" {
lib.Name = displayName
if key, ok := CloudLibraryMergeKey(lib); ok {
if _, exists := localByKey[key]; exists && !CloudLibraryAutoCategory(lib) {
continue
}
}
} else if displayName := CanonicalLibraryDisplayName(lib); displayName != "" {
lib.Name = displayName
}
out = append(out, lib)
}
return out
}
func dedupeDisplayLibrariesByMergeKey(libs []model.Library, counts map[string]int64) []model.Library {
if len(libs) == 0 {
return libs
}
out := make([]model.Library, 0, len(libs))
byKey := make(map[string]int, len(libs))
for _, lib := range libs {
if displayName := CanonicalLibraryDisplayName(lib); displayName != "" {
lib.Name = displayName
}
if CloudLibraryAutoCategory(lib) {
out = append(out, lib)
continue
}
key, ok := CloudLibraryMergeKey(lib)
if !ok {
out = append(out, lib)
continue
}
if prev, exists := byKey[key]; exists {
if betterCanonicalDisplayLibrary(lib, out[prev], counts) {
out[prev] = lib
}
continue
}
byKey[key] = len(out)
out = append(out, lib)
}
return out
}
func betterCanonicalDisplayLibrary(candidate, current model.Library, counts map[string]int64) bool {
candidateScore := canonicalDisplayLibraryScore(candidate)
currentScore := canonicalDisplayLibraryScore(current)
if candidateScore != currentScore {
return candidateScore > currentScore
}
candidateCount := counts[candidate.ID]
currentCount := counts[current.ID]
if (candidateCount > 0) != (currentCount > 0) {
return candidateCount > 0
}
if candidate.Enabled != current.Enabled {
return candidate.Enabled
}
if !candidate.CreatedAt.Equal(current.CreatedAt) {
return candidate.CreatedAt.After(current.CreatedAt)
}
return candidate.ID > current.ID
}
func canonicalDisplayLibraryScore(lib model.Library) int {
score := 0
if canonical := CanonicalLibraryDisplayName(lib); canonical == "" || strings.EqualFold(strings.TrimSpace(lib.Name), canonical) {
score += 4
}
if canonical := canonicalLibraryCategoryName(lib.Type, pathBaseSlash(lib.Path)); canonical == "" {
score += 2
}
if _, ok := ParseCloudLibraryMount(lib.Path); !ok {
score++
}
return score
}
-234
View File
@@ -1,234 +0,0 @@
package service
import (
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func NormalizeCloudLibraryDisplayNames(libs []model.Library) []model.Library {
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if displayName, ok := CloudLibraryDisplayName(lib); ok && displayName != "" {
lib.Name = displayName
} else if displayName := CanonicalLibraryDisplayName(lib); displayName != "" {
lib.Name = displayName
}
out = append(out, lib)
}
return out
}
func NormalizeCloudLibraryDisplay(libs []model.Library) []model.Library {
return normalizeDisplayLibraries(libs)
}
func normalizeDisplayLibraries(libs []model.Library) []model.Library {
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if displayName, ok := CloudLibraryDisplayName(lib); ok && displayName != "" {
lib.Name = displayName
} else if displayName := CanonicalLibraryDisplayName(lib); displayName != "" {
lib.Name = displayName
}
if displayType := CanonicalLibraryDisplayType(lib); displayType != "" {
lib.Type = displayType
}
if displayPath := CanonicalLibraryDisplayPath(lib); displayPath != "" {
lib.Path = displayPath
}
out = append(out, lib)
}
return out
}
func CloudLibraryDisplayName(lib model.Library) (string, bool) {
info, ok := ParseCloudLibraryMount(lib.Path)
if !ok {
return "", false
}
name := stripCloudProviderDisplayPrefix(strings.TrimSpace(lib.Name), info.Provider)
dir := firstNonEmpty(info.DisplayDir, info.ScanDir)
if name == "" || strings.EqualFold(name, CloudMountProviderLabel(info.Provider)) {
if base := cloudMountDirBase(dir); base != "" {
name = base
}
}
if name == "" {
name = CloudMountProviderLabel(info.Provider)
}
if canonical := canonicalLibraryCategoryName(lib.Type, name); canonical != "" {
name = canonical
} else if canonical := canonicalLibraryCategoryNameAny(name); canonical != "" {
name = canonical
}
return name, true
}
func CanonicalLibraryDisplayName(lib model.Library) string {
if canonical := canonicalLibraryCategoryName(lib.Type, lib.Name); canonical != "" {
return canonical
}
return canonicalLibraryCategoryNameAny(lib.Name)
}
func CanonicalLibraryDisplayType(lib model.Library) string {
if displayName, ok := CloudLibraryDisplayName(lib); ok {
if typ := canonicalLibraryCategoryDisplayType(displayName); typ != "" {
return typ
}
}
if typ := canonicalLibraryCategoryDisplayType(lib.Name); typ != "" {
return typ
}
return canonicalLibraryCategoryDisplayType(pathBaseSlash(lib.Path))
}
func CanonicalLibraryDisplayPath(lib model.Library) string {
raw := strings.TrimSpace(lib.Path)
if raw == "" {
return ""
}
if info, ok := ParseCloudLibraryMount(raw); ok {
dir := firstNonEmpty(info.DisplayDir, info.ScanDir)
displayDir := canonicalLibraryDisplayDir(dir)
if displayDir == "" {
return raw
}
if CloudLibraryAutoCategory(lib) {
return BuildCloudAutoCategoryLibraryPathWithScanDir(info.Provider, info.ScanDir, displayDir)
}
return BuildCloudLibraryPath(info.Provider, info.ScanDir, displayDir)
}
return canonicalLocalLibraryDisplayPath(raw)
}
func canonicalLibraryCategoryName(libraryType, name string) string {
typeKey := cloudLibraryMergeTypeKey(libraryType)
name = normalizeLibraryMergeName(name)
switch typeKey {
case "movie":
switch name {
case "国产电影", "大陆电影":
return "华语电影"
case "外语电影", "外国电影":
return "欧美电影"
case "日本电影", "韩国电影":
return "日韩电影"
case "音乐会", "concert":
return "演唱会"
case "纪录":
return "纪录片"
case "动漫电影":
return "动画电影"
}
case "tvshows":
switch name {
case "国剧", "大陆剧", "华语剧", "国产电视剧", "大陆电视剧", "华语电视剧", "港剧", "台剧", "港台剧":
return "国产剧"
case "欧美电视剧", "美剧", "英剧", "未分类", "uncategorized":
return "欧美剧"
case "日韩电视剧", "日剧", "韩剧", "泰剧":
return "日韩剧"
case "真人秀":
return "综艺"
case "纪录":
return "纪录片"
case "少儿":
return "儿童"
case "国产动漫", "国产动画":
return "国漫"
case "日漫", "番剧", "日本动漫", "日本动画":
return "日番"
case "韩国动漫", "韩国动画":
return "韩漫"
case "欧美动漫", "欧美动画", "西方动画":
return "美漫"
case "其他动漫", "其它动漫", "other":
return "其他"
}
case "adult":
switch name {
case "9kg", "番号", "jav", "nsfw", "adult":
return "成人"
}
}
return ""
}
func canonicalLibraryCategoryNameAny(name string) string {
for _, libraryType := range []string{"movie", "tv", "anime", "adult"} {
if canonical := canonicalLibraryCategoryName(libraryType, name); canonical != "" {
return canonical
}
}
return ""
}
func canonicalLibraryDisplayDir(raw string) string {
parts := strmSlashParts(raw)
if len(parts) == 0 {
return ""
}
return strings.Join(canonicalLibraryDisplayParts(parts), "/")
}
func canonicalLocalLibraryDisplayPath(raw string) string {
value := strings.TrimSpace(raw)
if value == "" {
return ""
}
sep := "/"
if strings.Contains(value, "\\") {
sep = "\\"
}
slash := strings.ReplaceAll(value, "\\", "/")
prefix := ""
for strings.HasPrefix(slash, "/") {
prefix += "/"
slash = strings.TrimPrefix(slash, "/")
}
parts := strings.Split(slash, "/")
canonical := canonicalLibraryDisplayParts(parts)
if len(canonical) == 0 {
return raw
}
out := prefix + strings.Join(canonical, "/")
if sep == "\\" {
out = strings.ReplaceAll(out, "/", "\\")
}
return out
}
func canonicalLibraryDisplayParts(parts []string) []string {
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part == "" || part == "." {
continue
}
if canonical := canonicalLibraryCategoryNameAny(part); canonical != "" {
part = canonical
}
if len(out) > 0 && normalizeLibraryMergeName(out[len(out)-1]) == normalizeLibraryMergeName(part) {
continue
}
out = append(out, part)
}
return out
}
func canonicalLibraryCategoryDisplayType(name string) string {
switch normalizeLibraryMergeName(name) {
case "演唱会", "音乐会", "动画电影", "动漫电影", "华语电影", "国产电影", "大陆电影", "欧美电影", "外语电影", "外国电影", "日韩电影", "日本电影", "韩国电影":
return "movie"
case "国产剧", "国剧", "大陆剧", "华语剧", "国产电视剧", "大陆电视剧", "华语电视剧", "港剧", "台剧", "港台剧", "欧美剧", "欧美电视剧", "美剧", "英剧", "未分类", "uncategorized", "日韩剧", "日韩电视剧", "日剧", "韩剧", "泰剧", "综艺", "真人秀", "儿童", "少儿":
return "tv"
case "国漫", "国产动漫", "国产动画", "日番", "日漫", "番剧", "日本动漫", "日本动画", "韩漫", "韩国动漫", "韩国动画", "美漫", "欧美动漫", "欧美动画", "西方动画", "其他", "其他动漫", "其它动漫", "other":
return "anime"
case "成人", "9kg", "番号", "jav", "nsfw", "adult":
return "adult"
default:
return ""
}
}
-123
View File
@@ -1,123 +0,0 @@
package service
import (
"context"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func FilterDisplayCloudLibraries(ctx context.Context, repo *repository.Container, libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
libs = FilterDeprecatedNativeCloudLibraries(libs)
libs = FilterMergedCloudAutoCategoryLibraries(libs)
counts := cloudLibraryMediaCounts(ctx, repo, libs)
collapsed := make([]model.Library, 0, len(libs))
byKey := make(map[string]int, len(libs))
for _, lib := range libs {
key, ok := cloudLibraryDisplayKey(lib)
if !ok {
collapsed = append(collapsed, lib)
continue
}
if prevIndex, exists := byKey[key]; exists {
if betterDisplayCloudLibrary(lib, collapsed[prevIndex], counts) {
collapsed[prevIndex] = lib
}
continue
}
byKey[key] = len(collapsed)
collapsed = append(collapsed, lib)
}
collapsed = FilterShadowedCloudLibraries(collapsed)
return normalizeDisplayLibraries(dedupeDisplayLibrariesByMergeKey(mergeDisplayCloudLibraries(collapsed), counts))
}
func FilterInternalCloudAutoCategoryLibraries(libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if CloudLibraryAutoCategory(lib) {
continue
}
out = append(out, lib)
}
return out
}
func FilterMergedCloudAutoCategoryLibraries(libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
nonAutoKeys := make(map[string]struct{}, len(libs))
for _, lib := range libs {
if CloudLibraryAutoCategory(lib) {
continue
}
if key, ok := CloudLibraryMergeKey(lib); ok {
nonAutoKeys[key] = struct{}{}
}
}
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
if CloudLibraryAutoCategory(lib) {
if key, ok := CloudLibraryMergeKey(lib); ok {
if _, merged := nonAutoKeys[key]; merged {
continue
}
}
}
out = append(out, lib)
}
return out
}
func FilterScannableCloudLibraries(ctx context.Context, repo *repository.Container, libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
counts := cloudLibraryMediaCounts(ctx, repo, libs)
collapsed := make([]model.Library, 0, len(libs))
byKey := make(map[string]int, len(libs))
for _, lib := range libs {
if CloudLibraryAutoCategory(lib) {
continue
}
if info, ok := ParseCloudLibraryMount(lib.Path); ok && IsDeprecatedNativeCloudProvider(info.Provider) {
continue
}
key, ok := cloudLibraryDisplayKey(lib)
if !ok {
collapsed = append(collapsed, lib)
continue
}
if prevIndex, exists := byKey[key]; exists {
if betterDisplayCloudLibrary(lib, collapsed[prevIndex], counts) {
collapsed[prevIndex] = lib
}
continue
}
byKey[key] = len(collapsed)
collapsed = append(collapsed, lib)
}
return FilterShadowedCloudLibraries(collapsed)
}
func FilterDeprecatedNativeCloudLibraries(libs []model.Library) []model.Library {
if len(libs) == 0 {
return libs
}
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
info, ok := ParseCloudLibraryMount(lib.Path)
if ok && IsDeprecatedNativeCloudProvider(info.Provider) {
continue
}
out = append(out, lib)
}
return out
}
-532
View File
@@ -1,532 +0,0 @@
package service
import (
"path/filepath"
"slices"
"strings"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/config"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestFilterDisplayCloudLibrariesPrefersPopulatedCanonicalDuplicate(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
now := time.Now()
oldEmpty := model.Library{
Base: model.Base{ID: "old-empty", CreatedAt: now.Add(-time.Hour)},
Name: "OpenList · 国产剧",
Path: "cloud://openlist/%2F国产剧",
Type: "tv",
Enabled: true,
}
newPopulated := model.Library{
Base: model.Base{ID: "new-populated", CreatedAt: now},
Name: "OpenList · 国产剧",
Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"),
Type: "tv",
Enabled: true,
}
if err := repos.Library.Create(t.Context(), &oldEmpty); err != nil {
t.Fatal(err)
}
if err := repos.Library.Create(t.Context(), &newPopulated); err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.Media{
LibraryID: newPopulated.ID,
Title: "剧集",
Path: "cloud://openlist/国产剧/剧集.mkv",
}).Error; err != nil {
t.Fatal(err)
}
filtered := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{oldEmpty, newPopulated})
if len(filtered) != 1 || filtered[0].ID != newPopulated.ID {
t.Fatalf("filtered = %#v, want only populated canonical duplicate", filtered)
}
scanner := NewScannerService(nil, zap.NewNop(), repos, nil, nil, nil)
if conflict := scanner.shadowedCloudLibrary(t.Context(), &oldEmpty); conflict == nil || conflict.Library.ID != newPopulated.ID {
t.Fatalf("old duplicate scan conflict = %#v, want populated canonical library", conflict)
}
}
func TestFilterDisplayCloudLibrariesMergesCloudMountIntoExistingLibrary(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
local := model.Library{Name: "国产剧", Path: "/media/国产剧", Type: "tv", Enabled: true}
cloud := model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv", Enabled: true}
movieCloud := model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/电影/国产剧", "/电影/国产剧"), Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&local, &cloud, &movieCloud} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
filtered := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{local, cloud, movieCloud})
if got := libraryNames(filtered); !slices.Equal(got, []string{"国产剧", "国产剧"}) {
t.Fatalf("filtered names = %#v, want local tv plus stripped movie cloud", got)
}
if filtered[0].ID != local.ID {
t.Fatalf("first filtered library = %s, want existing local library %s", filtered[0].ID, local.ID)
}
if filtered[1].ID != movieCloud.ID {
t.Fatalf("movie cloud library should stay separate when type differs: %#v", filtered)
}
merged := MergedLibraryIDs([]model.Library{local, cloud, movieCloud}, local)
if !slices.Equal(merged, []string{local.ID, cloud.ID}) {
t.Fatalf("merged ids = %#v, want local+same-type cloud", merged)
}
}
func TestFilterDisplayCloudLibrariesMergesEpisodicTypeAliases(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
local := model.Library{Name: "国漫", Path: "/media/动漫/国漫", Type: "tv", Enabled: true}
cloud := model.Library{Name: "OpenList · 国漫", Path: BuildCloudLibraryPath("openlist", "/国漫", "/国漫"), Type: "anime", Enabled: true}
movie := model.Library{Name: "国漫", Path: BuildCloudLibraryPath("openlist", "/电影/国漫", "/电影/国漫"), Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&local, &cloud, &movie} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
filtered := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{local, cloud, movie})
if got := libraryNames(filtered); !slices.Equal(got, []string{"国漫", "国漫"}) {
t.Fatalf("filtered names = %#v, want local episodic plus separate movie library", got)
}
if filtered[0].ID != local.ID || filtered[1].ID != movie.ID {
t.Fatalf("filtered libraries = %#v, want anime cloud merged into local tv but movie kept", filtered)
}
merged := MergedLibraryIDs([]model.Library{local, cloud, movie}, local)
if !slices.Equal(merged, []string{local.ID, cloud.ID}) {
t.Fatalf("merged ids = %#v, want local tv + cloud anime only", merged)
}
}
func TestFilterDisplayCloudLibrariesMergesCategoryNameAliases(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
foreignMovie := model.Library{Name: "外语电影", Path: "/media/电影/外语电影", Type: "movie", Enabled: true}
westernMovie := model.Library{Name: "OpenList · 欧美电影", Path: BuildCloudLibraryPath("openlist", "/欧美电影", "/欧美电影"), Type: "movie", Enabled: true}
eastAsianMovie := model.Library{Name: "OpenList · 日韩电影", Path: BuildCloudLibraryPath("openlist", "/日韩电影", "/日韩电影"), Type: "movie", Enabled: true}
jpAnime := model.Library{Name: "日番", Path: "/media/动漫/日番", Type: "tv", Enabled: true}
jpAnimeCloud := model.Library{Name: "OpenList · 日漫", Path: BuildCloudLibraryPath("openlist", "/日漫", "/日漫"), Type: "anime", Enabled: true}
for _, lib := range []*model.Library{&foreignMovie, &westernMovie, &eastAsianMovie, &jpAnime, &jpAnimeCloud} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
filtered := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{foreignMovie, westernMovie, eastAsianMovie, jpAnime, jpAnimeCloud})
if got := libraryNames(filtered); !slices.Equal(got, []string{"欧美电影", "日韩电影", "日番"}) {
t.Fatalf("filtered names = %#v, want legacy foreign movie merged into western movie plus anime aliases", got)
}
movieMerged := MergedLibraryIDs([]model.Library{foreignMovie, westernMovie, eastAsianMovie, jpAnime, jpAnimeCloud}, foreignMovie)
if !slices.Equal(movieMerged, []string{foreignMovie.ID, westernMovie.ID}) {
t.Fatalf("movie merged ids = %#v, want legacy foreign movie merged with western movie", movieMerged)
}
animeMerged := MergedLibraryIDs([]model.Library{foreignMovie, westernMovie, eastAsianMovie, jpAnime, jpAnimeCloud}, jpAnime)
if !slices.Equal(animeMerged, []string{jpAnime.ID, jpAnimeCloud.ID}) {
t.Fatalf("anime merged ids = %#v, want jp anime aliases", animeMerged)
}
}
func TestFilterDisplayCloudLibrariesCanonicalizesLegacyDisplayPaths(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
westernAnimation := model.Library{Name: "欧美动漫", Path: `F:\media\动漫\欧美动漫`, Type: "tv", Enabled: true}
uncategorizedCloud := model.Library{Name: "OpenList · 未分类", Path: BuildCloudLibraryPath("openlist", "/未分类", "/未分类"), Type: "movie", Enabled: true}
adult := model.Library{Name: "9KG", Path: `F:\media\成人\9KG`, Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&westernAnimation, &uncategorizedCloud, &adult} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
filtered := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{westernAnimation, uncategorizedCloud, adult})
if got := libraryNames(filtered); !slices.Equal(got, []string{"美漫", "欧美剧", "成人"}) {
t.Fatalf("filtered names = %#v, want canonical category names", got)
}
if got := []string{filtered[0].Type, filtered[1].Type, filtered[2].Type}; !slices.Equal(got, []string{"anime", "tv", "adult"}) {
t.Fatalf("filtered types = %#v, want canonical display types", got)
}
combined := strings.Join([]string{filtered[0].Path, filtered[1].Path, filtered[2].Path}, "\n")
for _, legacy := range []string{"欧美动漫", "未分类", "9KG"} {
if strings.Contains(combined, legacy) {
t.Fatalf("display paths contain legacy category %q: %s", legacy, combined)
}
}
}
func TestCanonicalLibraryDisplayPathPreservesAutoCategoryScanDir(t *testing.T) {
raw := BuildCloudAutoCategoryLibraryPathWithScanDir("openlist", "国漫", "动漫/国产动漫")
got := CanonicalLibraryDisplayPath(model.Library{Name: "国漫", Path: raw, Type: "anime", Enabled: true})
info, ok := ParseCloudLibraryMount(got)
if !ok {
t.Fatalf("canonical path did not parse: %q", got)
}
if !CloudLibraryAutoCategory(model.Library{Path: got}) {
t.Fatalf("canonical path lost auto_category flag: %q", got)
}
if info.ScanDir != "国漫" || info.DisplayDir != "动漫/国漫" {
t.Fatalf("canonical path info = %#v, want scan 国漫 and canonical display 动漫/国漫", info)
}
}
func TestListMediaVisibleDoesNotMergeDistinctMovieRegionLibraries(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
foreignMovie := model.Library{Name: "外语电影", Path: "/media/电影/外语电影", Type: "movie", Enabled: true}
westernMovie := model.Library{Name: "OpenList · 欧美电影", Path: BuildCloudLibraryPath("openlist", "/欧美电影", "/欧美电影"), Type: "movie", Enabled: true}
eastAsianMovie := model.Library{Name: "OpenList · 日韩电影", Path: BuildCloudLibraryPath("openlist", "/日韩电影", "/日韩电影"), Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&foreignMovie, &westernMovie, &eastAsianMovie} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
if err := repos.DB.Create(&model.Media{
LibraryID: westernMovie.ID,
Title: "Western Movie",
Path: "cloud://openlist/欧美电影/Western.Movie.2026.mkv",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
items, total, err := svc.ListMediaVisible(t.Context(), foreignMovie.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 1 || !slices.Equal(mediaTitles(items), []string{"Western Movie"}) {
t.Fatalf("legacy foreign movie items total=%d items=%#v, want merged western media", total, mediaTitles(items))
}
items, total, err = svc.ListMediaVisible(t.Context(), eastAsianMovie.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 0 || len(items) != 0 {
t.Fatalf("east asian movie items total=%d items=%#v, want empty isolated library", total, mediaTitles(items))
}
items, total, err = svc.ListMediaVisible(t.Context(), westernMovie.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 1 || !slices.Equal(mediaTitles(items), []string{"Western Movie"}) {
t.Fatalf("western movie items total=%d items=%#v, want own media only", total, mediaTitles(items))
}
}
func TestFilterDeprecatedNativeCloudLibrariesHidesPopulatedHistory(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
emptyQuark := model.Library{Name: "旧 Quark 空库", Path: "cloud://quark/0", Type: "movie", Enabled: true}
populatedQuark := model.Library{Name: "旧 Quark 有数据", Path: "cloud://quark/archive", Type: "movie", Enabled: true}
openList := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&emptyQuark, &populatedQuark, &openList} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
if err := repos.DB.Create(&model.Media{
LibraryID: populatedQuark.ID,
Title: "历史媒体",
Path: "cloud://quark/archive/movie.mkv",
}).Error; err != nil {
t.Fatal(err)
}
filtered := FilterDeprecatedNativeCloudLibraries([]model.Library{emptyQuark, populatedQuark, openList})
if got := libraryNames(filtered); !slices.Equal(got, []string{"OpenList"}) {
t.Fatalf("filtered names = %#v, want only supported cloud libraries", got)
}
displayed := FilterDisplayCloudLibraries(t.Context(), repos, []model.Library{emptyQuark, populatedQuark, openList})
if got := libraryNames(displayed); !slices.Equal(got, []string{"OpenList"}) {
t.Fatalf("display names = %#v, want deprecated cloud hidden", got)
}
}
func TestListMediaVisibleIncludesMergedCloudLibraryItems(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
local := model.Library{Name: "国产剧", Path: "/media/国产剧", Type: "tv", Enabled: true}
cloud := model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv", Enabled: true}
other := model.Library{Name: "欧美剧", Path: BuildCloudLibraryPath("openlist", "/欧美剧", "/欧美剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&local, &cloud, &other} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
if err := repos.DB.Create(&[]model.Media{
{LibraryID: local.ID, Title: "本地剧", Path: "/media/国产剧/local.mkv"},
{LibraryID: cloud.ID, Title: "云盘剧", Path: "cloud://openlist/国产剧/cloud.mkv"},
{LibraryID: other.ID, Title: "其他剧", Path: "cloud://openlist/欧美剧/other.mkv"},
}).Error; err != nil {
t.Fatal(err)
}
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
items, total, err := svc.ListMediaVisible(t.Context(), local.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 2 {
t.Fatalf("total = %d, want merged local+cloud items", total)
}
if got := mediaTitles(items); !slices.Equal(got, []string{"云盘剧", "本地剧"}) {
t.Fatalf("items = %#v, want local+cloud only", got)
}
if cloudItem := mediaByTitle(items, "云盘剧"); cloudItem == nil || cloudItem.DisplayLibraryID != local.ID {
t.Fatalf("cloud item display library = %#v, want merged local library %s", cloudItem, local.ID)
}
items, total, err = svc.ListMediaVisible(t.Context(), local.ID, 1, 20, MediaVisibility{
IncludeNSFW: true,
AllowedLibraryIDs: []string{local.ID},
})
if err != nil {
t.Fatal(err)
}
if total != 2 || !slices.Equal(mediaTitles(items), []string{"云盘剧", "本地剧"}) {
t.Fatalf("profile-limited merged list total=%d items=%#v", total, mediaTitles(items))
}
searchItems, err := svc.SearchMediaVisible(t.Context(), "剧", 20, MediaVisibility{
IncludeNSFW: true,
AllowedLibraryIDs: []string{local.ID},
})
if err != nil {
t.Fatal(err)
}
if got := mediaTitles(searchItems); !slices.Equal(got, []string{"云盘剧", "本地剧"}) {
t.Fatalf("profile-limited merged search items=%#v, want local+hidden cloud", got)
}
if cloudItem := mediaByTitle(searchItems, "云盘剧"); cloudItem == nil || cloudItem.DisplayLibraryID != local.ID {
t.Fatalf("search cloud item display library = %#v, want merged local library %s", cloudItem, local.ID)
}
}
func TestListMediaVisibleUsesSpecificCloudChildLibraryAsDisplayTarget(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
root := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "tv", Enabled: true}
child := model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&root, &child} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
if err := repos.DB.Create(&model.Media{
LibraryID: root.ID,
Title: "折腰",
Path: "cloud://openlist/国产剧/折腰 (2025)/Season 1/折腰.S01E01.mkv",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
items, total, err := svc.ListMediaVisible(t.Context(), root.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 1 || len(items) != 1 {
t.Fatalf("items total=%d len=%d, want one root cloud item", total, len(items))
}
if items[0].DisplayLibraryID != child.ID {
t.Fatalf("display library = %q, want child cloud library %q", items[0].DisplayLibraryID, child.ID)
}
if items[0].DisplayLibraryPath != child.Path {
t.Fatalf("display library path = %q, want %q", items[0].DisplayLibraryPath, child.Path)
}
}
func TestGetMediaUsesMappedLocalCategoryLibraryAsDisplayTarget(t *testing.T) {
containerRoot := filepath.Join(t.TempDir(), "media")
t.Setenv("MEDIASTATION_MEDIA_CONTAINER_DIR", containerRoot)
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
parent := model.Library{Name: "电视剧", Path: containerRoot, Type: "tv", Enabled: true}
child := model.Library{Name: "国产剧", Path: filepath.Join("media", "电视剧", "国产剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&parent, &child} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
mediaPath := filepath.Join(containerRoot, "电视剧", "国产剧", "剧集", "Season 01", "剧集 - S01E01.mkv")
if err := repos.DB.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: parent.ID,
Title: "剧集",
Path: mediaPath,
SeasonNum: 1,
EpisodeNum: 1,
}).Error; err != nil {
t.Fatal(err)
}
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
got, err := svc.GetMedia(t.Context(), "media-1")
if err != nil {
t.Fatal(err)
}
if got.DisplayLibraryID != child.ID {
t.Fatalf("display library = %q, want mapped child library %q", got.DisplayLibraryID, child.ID)
}
if got.DisplayLibraryPath != filepath.Clean(filepath.Join(containerRoot, "电视剧", "国产剧")) {
t.Fatalf("display library path = %q, want mapped child path", got.DisplayLibraryPath)
}
}
func TestStartAllCloudLibraryScansIncludesMergedCloudMounts(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
local := model.Library{Name: "国产剧", Path: "/media/国产剧", Type: "tv", Enabled: true}
cloud := model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&local, &cloud} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
statuses, err := scanner.StartAllCloudLibraryScans()
if err != nil {
t.Fatal(err)
}
if len(statuses) != 1 || statuses[0].LibraryID != cloud.ID {
t.Fatalf("scan-all statuses = %#v, want merged cloud library queued", statuses)
}
}
func TestAutoCategoryCloudLibrariesMergeIntoExistingDisplayLibrary(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
local := model.Library{Name: "欧美剧", Path: "/media/电视剧/欧美剧", Type: "tv", Enabled: true}
root := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
auto := model.Library{Name: "欧美剧", Path: BuildCloudAutoCategoryLibraryPath("openlist", "电视剧/欧美剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&local, &root, &auto} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
libs, err := repos.Library.List(t.Context())
if err != nil {
t.Fatal(err)
}
if shadow := CloudLibraryShadowed(libs, root); shadow != nil {
t.Fatalf("auto category should not shadow root scan: %#v", shadow)
}
display := FilterDisplayCloudLibraries(t.Context(), repos, libs)
if got := libraryNames(display); !slices.Equal(got, []string{"欧美剧", "OpenList"}) {
t.Fatalf("display libraries = %#v, want local library and user-mounted root only", got)
}
scannable := FilterScannableCloudLibraries(t.Context(), repos, libs)
if got := libraryNames(scannable); !slices.Equal(got, []string{"欧美剧", "OpenList"}) {
t.Fatalf("scannable libraries = %#v, want local library and root only", got)
}
scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
statuses, err := scanner.StartAllCloudLibraryScans()
if err != nil {
t.Fatal(err)
}
if len(statuses) != 1 || statuses[0].LibraryID != root.ID {
t.Fatalf("scan-all statuses = %#v, want only cloud root queued", statuses)
}
}
func TestRootCloudLibraryIncludesAutoCategoryMedia(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
root := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
auto := model.Library{Name: "欧美剧", Path: BuildCloudAutoCategoryLibraryPath("openlist", "电视剧/欧美剧"), Type: "tv", Enabled: true}
for _, lib := range []*model.Library{&root, &auto} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
if err := repos.DB.Create(&model.Media{
LibraryID: auto.ID,
Title: "The Show",
Path: "cloud://openlist/电视剧/欧美剧/The Show/The.Show.S01E01.mkv",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
items, total, err := svc.ListMediaVisible(t.Context(), root.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
if total != 1 || len(items) != 1 {
t.Fatalf("root cloud items total=%d len=%d, want auto-category media", total, len(items))
}
if items[0].LibraryName != auto.Name || items[0].LibraryPath != auto.Path {
t.Fatalf("media library metadata = (%q, %q), want auto category", items[0].LibraryName, items[0].LibraryPath)
}
if items[0].DisplayLibraryID != auto.ID || items[0].DisplayLibraryPath != auto.Path {
t.Fatalf("display library = (%q, %q), want auto category", items[0].DisplayLibraryID, items[0].DisplayLibraryPath)
}
}
func TestStartAllCloudLibraryScansSkipsDeprecatedQuarkMounts(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
quark := model.Library{Name: "旧 Quark", Path: "cloud://quark/0", Type: "movie", Enabled: true}
openList := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&quark, &openList} {
if err := repos.Library.Create(t.Context(), lib); err != nil {
t.Fatal(err)
}
}
scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
statuses, err := scanner.StartAllCloudLibraryScans()
if err != nil {
t.Fatal(err)
}
if len(statuses) != 1 || statuses[0].Provider != "openlist" {
t.Fatalf("scan-all statuses = %#v, want only openlist", statuses)
}
}
func libraryNames(libs []model.Library) []string {
out := make([]string, 0, len(libs))
for _, lib := range libs {
out = append(out, lib.Name)
}
return out
}
func mediaTitles(items []model.Media) []string {
out := make([]string, 0, len(items))
for _, item := range items {
out = append(out, item.Title)
}
slices.Sort(out)
return out
}
func mediaByTitle(items []model.Media, title string) *model.Media {
for i := range items {
if items[i].Title == title {
return &items[i]
}
}
return nil
}
-136
View File
@@ -1,136 +0,0 @@
package service
import (
"net/url"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
func CloudMountProviderLabel(provider string) string {
switch strings.TrimSpace(provider) {
case LegacyQuarkProvider:
return "已停用网盘"
case cloud.Type115:
return "115 网盘"
case cloud.TypeCloudDrive2:
return "CloudDrive2"
case cloud.TypeOpenList:
return "OpenList"
default:
if strings.TrimSpace(provider) == "" {
return "网盘"
}
return strings.TrimSpace(provider)
}
}
func stripCloudProviderDisplayPrefix(name, provider string) string {
name = strings.TrimSpace(name)
if name == "" {
return ""
}
for _, label := range []string{CloudMountProviderLabel(provider), strings.TrimSpace(provider)} {
label = strings.TrimSpace(label)
if label == "" || len(name) < len(label) || !strings.EqualFold(name[:len(label)], label) {
continue
}
rest := strings.TrimSpace(name[len(label):])
rest = strings.TrimLeft(rest, " \t\r\n·・-—–||:/\\")
if rest != "" {
return strings.TrimSpace(rest)
}
if strings.EqualFold(name, label) {
return ""
}
}
return name
}
func cloudMountDirBase(dir string) string {
dir = strings.Trim(strings.TrimSpace(strings.ReplaceAll(dir, "\\", "/")), "/")
if dir == "" {
return ""
}
parts := strings.Split(dir, "/")
for i := len(parts) - 1; i >= 0; i-- {
if part := strings.TrimSpace(parts[i]); part != "" {
return part
}
}
return ""
}
func normalizeLibraryMergeName(name string) string {
name = strings.ToLower(strings.TrimSpace(name))
if name == "" {
return ""
}
return strings.Join(strings.Fields(name), " ")
}
func ShadowedCloudLibraryIDSet(libs []model.Library) map[string]bool {
out := make(map[string]bool)
for _, lib := range libs {
if CloudLibraryShadowed(libs, lib) != nil {
out[lib.ID] = true
}
}
return out
}
func InferCloudMountMediaType(dir, name string) string {
text := strings.ToLower(dir + " " + name)
switch {
case strings.Contains(text, "成人") || strings.Contains(text, "adult") || strings.Contains(text, "jav") || strings.Contains(text, "9kg"):
return "adult"
case containsAny(text, "动画电影", "华语电影", "外语电影", "外国电影", "欧美电影", "日韩电影", "韩国电影", "日本电影", "港台电影", "香港电影", "台湾电影", "大陆电影", "国产电影", "纪录片", "演唱会", "音乐会", "电影", "movie", "movies", "film", "films", "documentary", "concert"):
return "movie"
case containsAny(text, "综艺", "真人秀", "脱口秀", "晚会", "variety"):
return "variety"
case containsAny(text, "国漫", "日漫", "日番", "韩漫", "美漫", "番剧", "动漫", "欧美动漫", "动画剧集", "anime"):
return "anime"
case containsAny(text, "国产剧", "大陆剧", "华语剧", "欧美剧", "日韩剧", "韩剧", "日剧", "港剧", "台剧", "泰剧", "英剧", "美剧", "短剧", "电视剧", "剧集", "连续剧", "series", "tv", "shows"):
return "tv"
default:
return "movie"
}
}
func containsAny(text string, values ...string) bool {
for _, value := range values {
if strings.Contains(text, value) {
return true
}
}
return false
}
func cloudMountAncestor(parent, child string) bool {
parent = strings.Trim(parent, "/")
child = strings.Trim(child, "/")
if parent == child {
return false
}
if parent == "" {
return child != ""
}
return strings.HasPrefix(child, parent+"/")
}
func normalizeCloudMountDir(provider, value string) string {
value = strings.TrimSpace(value)
if decoded, err := url.PathUnescape(value); err == nil {
value = decoded
}
if decoded, err := url.QueryUnescape(value); err == nil {
value = decoded
}
value = strings.ReplaceAll(value, "\\", "/")
value = strings.Trim(strings.TrimSpace(value), "/")
if value == "." || ((provider == cloud.Type115 || provider == LegacyQuarkProvider) && value == "0") {
return ""
}
return value
}
-90
View File
@@ -1,90 +0,0 @@
package service
import (
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func cloudLibraryDisplayKey(lib model.Library) (string, bool) {
info, ok := ParseCloudLibraryMount(lib.Path)
if !ok {
return "", false
}
dir := firstNonEmpty(info.DisplayDir, info.ScanDir)
return info.Provider + "\x00" + dir, true
}
func cloudLibraryPathIsCanonical(lib model.Library) bool {
info, ok := ParseCloudLibraryMount(lib.Path)
if !ok {
return false
}
return BuildCloudLibraryPath(info.Provider, info.ScanDir, info.DisplayDir) == strings.TrimSpace(lib.Path)
}
func CloudLibraryMergeKey(lib model.Library) (string, bool) {
name := strings.TrimSpace(lib.Name)
if displayName, ok := CloudLibraryDisplayName(lib); ok {
name = displayName
}
name = normalizeLibraryMergeName(name)
if name == "" {
return "", false
}
typeKey := cloudLibraryMergeTypeKey(lib.Type)
return typeKey + "\x00" + cloudLibraryMergeNameKey(typeKey, name), true
}
func cloudLibraryMergeTypeKey(libraryType string) string {
switch strings.ToLower(strings.TrimSpace(libraryType)) {
case "tv", "anime", "variety":
return "tvshows"
default:
return strings.ToLower(strings.TrimSpace(libraryType))
}
}
func cloudLibraryMergeNameKey(typeKey, name string) string {
switch typeKey {
case "movie":
switch name {
case "国产电影", "大陆电影":
return "华语电影"
case "华语电影":
return "华语电影"
case "外语电影", "外国电影", "欧美电影":
return "欧美电影"
case "日韩电影", "日本电影", "韩国电影":
return "日韩电影"
case "纪录", "纪录片":
return "纪录片"
case "演唱会", "concert":
return "演唱会"
case "动画电影", "动漫电影":
return "动画电影"
}
case "tvshows":
switch name {
case "国产剧", "大陆剧", "华语剧", "国剧", "国产电视剧", "大陆电视剧", "华语电视剧", "港剧", "台剧", "港台剧":
return "国产剧"
case "欧美剧", "欧美电视剧", "美剧", "英剧":
return "欧美剧"
case "日韩剧", "日韩电视剧", "日剧", "韩剧", "泰剧":
return "日韩剧"
case "国漫", "国产动漫", "国产动画":
return "国漫"
case "日番", "日漫", "番剧", "日本动漫", "日本动画":
return "日番"
case "韩漫", "韩国动漫", "韩国动画":
return "韩漫"
case "美漫", "欧美动漫", "欧美动画", "西方动画":
return "美漫"
case "其他", "其他动漫", "其它动漫", "other":
return "其他"
case "纪录", "纪录片":
return "纪录片"
}
}
return name
}
-152
View File
@@ -1,152 +0,0 @@
package service
import (
"context"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func MergedLibraryIDsForLibrary(ctx context.Context, repo *repository.Container, libraryID string) ([]string, error) {
libraryID = strings.TrimSpace(libraryID)
if libraryID == "" || repo == nil || repo.Library == nil {
return []string{libraryID}, nil
}
lib, err := repo.Library.FindByID(ctx, libraryID)
if err != nil {
return nil, err
}
if lib == nil {
return []string{libraryID}, nil
}
libs, err := repo.Library.List(ctx)
if err != nil {
return nil, err
}
return MergedLibraryIDs(libs, *lib), nil
}
func MergedLibraryIDs(libs []model.Library, target model.Library) []string {
ids := appendUniqueLibraryIDs(nil, target.ID)
if rootAutoIDs := cloudRootAutoCategoryLibraryIDs(libs, target); len(rootAutoIDs) > 0 {
ids = appendUniqueLibraryIDs(ids, rootAutoIDs...)
}
targetKey, hasTargetKey := CloudLibraryMergeKey(target)
if !hasTargetKey {
return ids
}
_, targetIsCloud := ParseCloudLibraryMount(target.Path)
for _, candidate := range libs {
if candidate.ID == target.ID || strings.TrimSpace(candidate.ID) == "" || !candidate.Enabled {
continue
}
key, ok := CloudLibraryMergeKey(candidate)
if ok && key == targetKey {
_, candidateIsCloud := ParseCloudLibraryMount(candidate.Path)
if !targetIsCloud && !candidateIsCloud {
continue
}
ids = appendUniqueLibraryIDs(ids, candidate.ID)
}
}
return ids
}
func cloudRootAutoCategoryLibraryIDs(libs []model.Library, lib model.Library) []string {
mount, ok := ParseCloudLibraryMount(lib.Path)
if !ok || !cloudRootMountNeedsAutoCategory(mount) {
return nil
}
ids := make([]string, 0)
for _, candidate := range libs {
if candidate.ID == lib.ID || !candidate.Enabled || !CloudLibraryAutoCategory(candidate) {
continue
}
info, ok := ParseCloudLibraryMount(candidate.Path)
if ok && info.Provider == mount.Provider {
ids = appendUniqueLibraryIDs(ids, candidate.ID)
}
}
return ids
}
func ExpandMediaVisibilityForMergedCloudLibraries(ctx context.Context, repo *repository.Container, visibility MediaVisibility) MediaVisibility {
if repo == nil || repo.Library == nil {
return visibility
}
libs, err := repo.Library.List(ctx)
if err != nil {
return visibility
}
if len(visibility.AllowedLibraryIDs) > 0 {
visibility.AllowedLibraryIDs = expandMergedLibraryIDsFromLibraries(libs, visibility.AllowedLibraryIDs)
}
if len(visibility.HiddenLibraryIDs) > 0 {
visibility.HiddenLibraryIDs = expandMergedLibraryIDsFromLibraries(libs, visibility.HiddenLibraryIDs)
}
visibility.HiddenLibraryIDs = appendUniqueLibraryIDs(visibility.HiddenLibraryIDs, DeprecatedNativeCloudLibraryIDs(libs)...)
return visibility
}
func expandMergedLibraryIDs(ctx context.Context, repo *repository.Container, ids []string) []string {
if len(ids) == 0 || repo == nil || repo.Library == nil {
return ids
}
libs, err := repo.Library.List(ctx)
if err != nil {
return ids
}
return expandMergedLibraryIDsFromLibraries(libs, ids)
}
func expandMergedLibraryIDsFromLibraries(libs []model.Library, ids []string) []string {
byID := make(map[string]model.Library, len(libs))
for _, lib := range libs {
byID[lib.ID] = lib
}
out := make([]string, 0, len(ids))
for _, id := range ids {
id = strings.TrimSpace(id)
if id == "" {
continue
}
if lib, ok := byID[id]; ok {
out = appendUniqueLibraryIDs(out, MergedLibraryIDs(libs, lib)...)
continue
}
out = appendUniqueLibraryIDs(out, id)
}
return out
}
func DeprecatedNativeCloudLibraryIDs(libs []model.Library) []string {
ids := make([]string, 0)
for _, lib := range libs {
info, ok := ParseCloudLibraryMount(lib.Path)
if ok && IsDeprecatedNativeCloudProvider(info.Provider) {
ids = appendUniqueLibraryIDs(ids, lib.ID)
}
}
return ids
}
func appendUniqueLibraryIDs(ids []string, values ...string) []string {
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
exists := false
for _, id := range ids {
if id == value {
exists = true
break
}
}
if !exists {
ids = append(ids, value)
}
}
return ids
}
-122
View File
@@ -1,122 +0,0 @@
package service
import (
"context"
"strings"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// RepairCloudPathMetadata backfills external IDs from media paths such as
// "Movie (2025) {tmdb-123}" so existing placeholder rows can be scraped
// without requiring another successful filesystem or cloud provider traversal.
//
// 传入 libraryID 时只修复这些媒体库的行;为空则修复全库。
func (c *Container) RepairCloudPathMetadata(ctx context.Context, libraryID ...string) (int, error) {
if c == nil || c.Repo == nil || c.Repo.DB == nil {
return 0, nil
}
libraryIDs := compactLibraryIDs(libraryID...)
var repaired int
var rows []model.Media
query := c.Repo.DB.WithContext(ctx).
Model(&model.Media{}).
Select("id, title, path, year, season_num, episode_num, scrape_status, tm_db_id, bangumi_id, douban_id, thetvdb_id").
Where("("+strings.Join([]string{
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
"LOWER(path) LIKE ?",
}, " OR ")+")",
"%tmdb%", "%tmdbid%", "%douban%", "%db%", "%bangumi%", "%bgm%", "%thetvdb%", "%tvdb%")
if len(libraryIDs) > 0 {
query = query.Where("library_id IN ?", libraryIDs)
}
err := query.FindInBatches(&rows, 500, func(_ *gorm.DB, _ int) error {
for _, row := range rows {
meta, hints := pathHintMetadata(row.Path, row.SeasonNum > 0 || row.EpisodeNum > 0)
if meta == nil || !hints.useful() {
continue
}
updates := map[string]any{}
status := strings.TrimSpace(row.ScrapeStatus)
enrichable := status == "" || status == "pending" || status == "no_match"
changedExternalID := false
if meta.TMDbID > 0 && row.TMDbID != meta.TMDbID {
updates["tm_db_id"] = meta.TMDbID
changedExternalID = true
}
if meta.BangumiID > 0 && row.BangumiID != meta.BangumiID {
updates["bangumi_id"] = meta.BangumiID
changedExternalID = true
}
if strings.TrimSpace(meta.DoubanID) != "" && strings.TrimSpace(row.DoubanID) != strings.TrimSpace(meta.DoubanID) {
updates["douban_id"] = strings.TrimSpace(meta.DoubanID)
changedExternalID = true
}
if strings.TrimSpace(meta.TheTVDBID) != "" && strings.TrimSpace(row.TheTVDBID) != strings.TrimSpace(meta.TheTVDBID) {
updates["thetvdb_id"] = strings.TrimSpace(meta.TheTVDBID)
changedExternalID = true
}
if meta.Year > 0 && row.Year <= 0 {
updates["year"] = meta.Year
}
if enrichable && strings.TrimSpace(meta.Title) != "" && cloudPathRepairShouldReplaceTitle(row.Title, meta.Title) {
updates["title"] = strings.TrimSpace(meta.Title)
}
if changedExternalID && (status == "" || status == "no_match" || status == "matched") {
updates["scrape_status"] = "pending"
}
if len(updates) == 0 {
continue
}
if err := c.Repo.DB.WithContext(ctx).Model(&model.Media{}).Where("id = ?", row.ID).Updates(updates).Error; err != nil {
return err
}
repaired++
}
return nil
}).Error
if err != nil {
return repaired, err
}
if repaired > 0 && c.Log != nil {
c.Log.Info("cloud path metadata repaired", zap.Int("media_count", repaired))
}
return repaired, nil
}
func cloudPathRepairShouldReplaceTitle(current, hinted string) bool {
current = strings.TrimSpace(current)
hinted = strings.TrimSpace(hinted)
if hinted == "" || strings.EqualFold(current, hinted) {
return false
}
if current == "" {
return true
}
noise := []string{"web-dl", "bluray", "hdtv", "2160p", "1080p", "720p", "ddp", "aac", "h.264", "h.265", "x264", "x265", "adweb", "mweb", "cmctv", "bit"}
lower := strings.ToLower(current)
for _, token := range noise {
if strings.Contains(lower, token) {
return true
}
}
return len([]rune(current)) > len([]rune(hinted))*2
}
func compactLibraryIDs(ids ...string) []string {
out := make([]string, 0, len(ids))
for _, id := range ids {
out = appendUniqueLibraryIDs(out, id)
}
return out
}
-233
View File
@@ -1,233 +0,0 @@
package service
import (
"os"
"path/filepath"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestRepairAndRescrapeLibraryForceRematchesThenReclassifies(t *testing.T) {
scraper, repos, closeServer := newTestScraper(t)
defer closeServer()
root := t.TempDir()
wrongRoot := filepath.Join(root, "media", "电视剧", "国产剧")
mediaPath := filepath.Join(wrongRoot, "Spy Family", "Season 01", "Spy Family - S01E01.mkv")
if err := os.MkdirAll(filepath.Dir(mediaPath), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(mediaPath, []byte("episode"), 0o644); err != nil {
t.Fatal(err)
}
lib := model.Library{Name: "国产剧", Path: wrongRoot, Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
media := model.Media{
LibraryID: lib.ID,
Title: "错误旧匹配",
Path: mediaPath,
SeasonNum: 1,
EpisodeNum: 1,
TMDbID: 999,
Countries: "CN",
Languages: "zh",
Genres: "Drama",
ScrapeStatus: "matched",
}
if err := repos.DB.Create(&media).Error; err != nil {
t.Fatal(err)
}
cfg := &config.Config{}
cfg.Organizer.SmartClassify = true
organizer := NewOrganizerService(cfg, zap.NewNop(), repos)
organizer.SetScraper(scraper)
container := &Container{Cfg: cfg, Log: zap.NewNop(), Repo: repos, Scraper: scraper, Organizer: organizer}
result, err := container.RepairAndRescrapeLibrary(t.Context(), lib.ID)
if err != nil {
t.Fatalf("repair and rescrape: %v", err)
}
if result.Reclassified != 1 {
t.Fatalf("result=%+v, want one corrected classification", result)
}
want := filepath.Join(root, "media", "动漫", "日番", "间谍过家家", "Season 01", "间谍过家家 - S01E01.mkv")
if _, err := os.Stat(want); err != nil {
t.Fatalf("corrected media missing at %q: %v", want, err)
}
var got model.Media
if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
t.Fatal(err)
}
if got.TMDbID != 12345 || got.Countries != "JP" || got.Path != want {
t.Fatalf("repaired media=%#v, want rematched Japanese anime at corrected path", got)
}
}
func TestRepairRescrapeOptionsDefaultSkipsEpisodeArtwork(t *testing.T) {
options := repairRescrapeOptions()
if !options.RetryNoMatch {
t.Fatal("repair rescrape should retry no_match rows")
}
if !options.IncludeMatched {
t.Fatal("repair rescrape should refresh already matched rows")
}
if !options.ForceRematch {
t.Fatal("repair rescrape should ignore stale external IDs and rematch by path/title")
}
if options.EpisodeArtwork == nil {
t.Fatal("repair rescrape should set an explicit episode artwork option")
}
if *options.EpisodeArtwork {
t.Fatal("repair rescrape should skip episode artwork by default")
}
}
func TestRepairRescrapeOptionsCanEnableEpisodeArtwork(t *testing.T) {
episodeArtwork := true
options := repairRescrapeOptions(ScrapeOptions{EpisodeArtwork: &episodeArtwork})
if !options.RetryNoMatch {
t.Fatal("repair rescrape should force retry no_match rows")
}
if !options.IncludeMatched {
t.Fatal("repair rescrape should force refreshing already matched rows")
}
if options.EpisodeArtwork == nil || !*options.EpisodeArtwork {
t.Fatal("repair rescrape should keep explicit episode artwork=true")
}
}
func TestRepairRescrapeOptionsKeepsExplicitEpisodeArtworkFalse(t *testing.T) {
episodeArtwork := false
options := repairRescrapeOptions(ScrapeOptions{EpisodeArtwork: &episodeArtwork})
if !options.RetryNoMatch {
t.Fatal("repair rescrape should force retry no_match rows")
}
if !options.IncludeMatched {
t.Fatal("repair rescrape should force refreshing already matched rows")
}
if options.EpisodeArtwork == nil {
t.Fatal("repair rescrape should keep explicit episode artwork option")
}
if *options.EpisodeArtwork {
t.Fatal("repair rescrape should keep explicit episode artwork=false")
}
}
// TestResetEpisodicMatchedForRescrape 验证「修复+重刮」会把脏的 matched 剧集行
// 重置为 pending(让 EnrichLibrary 能重新刮削),而电影行与其它库不受影响。
func TestResetEpisodicMatchedForRescrape(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
container := &Container{Repo: repository.New(db), Log: zap.NewNop()}
rows := []model.Media{
// 目标库的剧集行(matched, 有季集号)→ 应被重置。
{Base: model.Base{ID: "ep1"}, LibraryID: "lib-a", SeasonNum: 1, EpisodeNum: 1, ScrapeStatus: "matched", Path: "/a/show/S01/ep1.mkv"},
{Base: model.Base{ID: "ep2"}, LibraryID: "lib-a", SeasonNum: 1, EpisodeNum: 2, ScrapeStatus: "matched", Path: "/a/show/S01/ep2.mkv"},
// 目标库的电影行(无季集号)→ 不应被重置。
{Base: model.Base{ID: "movie1"}, LibraryID: "lib-a", ScrapeStatus: "matched", Path: "/a/movie.mkv"},
// 目标库已是 pending 的剧集行 → 不计入重置数。
{Base: model.Base{ID: "ep3"}, LibraryID: "lib-a", SeasonNum: 1, EpisodeNum: 3, ScrapeStatus: "pending", Path: "/a/show/S01/ep3.mkv"},
// 其它库的剧集行(matched)→ 单库重置时不应受影响。
{Base: model.Base{ID: "ep-other"}, LibraryID: "lib-b", SeasonNum: 1, EpisodeNum: 1, ScrapeStatus: "matched", Path: "/b/show/S01/ep1.mkv"},
}
for i := range rows {
if err := db.Create(&rows[i]).Error; err != nil {
t.Fatalf("create media %s: %v", rows[i].ID, err)
}
}
reset, err := container.resetEpisodicMatchedForRescrape(t.Context(), "lib-a")
if err != nil {
t.Fatalf("reset: %v", err)
}
if reset != 2 {
t.Fatalf("reset = %d, want 2 (only matched episodic rows in lib-a)", reset)
}
status := func(id string) string {
var m model.Media
if err := db.First(&m, "id = ?", id).Error; err != nil {
t.Fatalf("load %s: %v", id, err)
}
return m.ScrapeStatus
}
if status("ep1") != "pending" || status("ep2") != "pending" {
t.Fatalf("episodic matched rows should be pending, got ep1=%q ep2=%q", status("ep1"), status("ep2"))
}
if status("movie1") != "matched" {
t.Fatalf("movie row should stay matched, got %q", status("movie1"))
}
if status("ep-other") != "matched" {
t.Fatalf("other library row should stay matched, got %q", status("ep-other"))
}
}
func TestRepairAndRescrapeLibraryExpandsMergedCloudLibraries(t *testing.T) {
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
container := &Container{Repo: repository.New(db), Log: zap.NewNop()}
local := model.Library{Name: "国产剧", Path: "/media/电视剧/国产剧", Type: "tv", Enabled: true}
cloud := model.Library{
Name: "OpenList · 国产剧",
Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"),
Type: "tv",
Enabled: true,
}
if err := container.Repo.Library.Create(t.Context(), &local); err != nil {
t.Fatal(err)
}
if err := container.Repo.Library.Create(t.Context(), &cloud); err != nil {
t.Fatal(err)
}
repairMedia := model.Media{
LibraryID: cloud.ID,
Title: "主角",
Path: "cloud://openlist/国产剧/主角 (2026) {tmdb-284110}/Season 1/主角.S01E01.mkv",
SeasonNum: 1,
EpisodeNum: 1,
ScrapeStatus: "matched",
}
resetMedia := model.Media{
LibraryID: cloud.ID,
Title: "无占位符剧集",
Path: "cloud://openlist/国产剧/无占位符剧集/Season 1/无占位符剧集.S01E01.mkv",
SeasonNum: 1,
EpisodeNum: 1,
ScrapeStatus: "matched",
}
if err := db.Create(&repairMedia).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&resetMedia).Error; err != nil {
t.Fatal(err)
}
result, err := container.RepairAndRescrapeLibrary(t.Context(), local.ID)
if err != nil {
t.Fatal(err)
}
if result.Repaired != 1 || result.Reset != 1 {
t.Fatalf("result = %+v, want repaired/reset for merged cloud row", result)
}
var repaired model.Media
if err := db.First(&repaired, "id = ?", repairMedia.ID).Error; err != nil {
t.Fatal(err)
}
if repaired.TMDbID != 284110 || repaired.ScrapeStatus != "pending" {
t.Fatalf("merged cloud row not repaired/reset: tmdb=%d status=%q", repaired.TMDbID, repaired.ScrapeStatus)
}
var reset model.Media
if err := db.First(&reset, "id = ?", resetMedia.ID).Error; err != nil {
t.Fatal(err)
}
if reset.ScrapeStatus != "pending" {
t.Fatalf("merged cloud row not reset: status=%q", reset.ScrapeStatus)
}
}
-217
View File
@@ -1,217 +0,0 @@
package service
import (
"context"
"strings"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// RepairAndRescrapeResult 汇总一次「全库修复+重刮」的结果。
type RepairAndRescrapeResult struct {
Repaired int `json:"repaired"` // 从路径占位符回填外部 ID 的媒体数
Reclassified int `json:"reclassified"` // 按元数据纠偏到正确分类/媒体库的媒体数
Libraries int `json:"libraries"` // 参与重刮的媒体库数
Matched int `json:"matched"` // 重刮后成功匹配的媒体数
Processed int `json:"processed"` // 实际完成刮削处理的媒体数
Errors int `json:"errors"` // 单条媒体刮削失败数
Reset int `json:"reset"` // 被重置为 pending 以便重刮的剧集行数
}
// resetEpisodicMatchedForRescrape 把剧集类(有季集号)且已 matched 的行重置为
// pending,使 EnrichLibrary(只处理 pending/no_match)能重新刮削它们。
//
// 背景: 历史版本(commit b44c7f8)曾把【单集 episode id】写进整剧 tm_db_id、把
// 单集名写进 original_name,污染了合集分组键 —— 同一部剧每集 id/原名各不相同,
// 被前端 / Emby 拆成 N 张单集卡。这些行 scrape_status 多为 matched,常规「修复+
// 重刮」会跳过,导致「无法修复」。源头已在 local_metadata.go 修正,这里把脏的
// matched 剧集行放回 pending,借重刮写回正确的整剧 ID / 原名。
//
// libraryIDs 为空时处理全库;非空时仅这些库。返回被重置的行数。
func (c *Container) resetEpisodicMatchedForRescrape(ctx context.Context, libraryIDs ...string) (int, error) {
if c == nil || c.Repo == nil || c.Repo.DB == nil {
return 0, nil
}
ids := compactLibraryIDs(libraryIDs...)
q := c.Repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("(season_num > 0 OR episode_num > 0)").
Where("LOWER(scrape_status) = ?", "matched")
if len(ids) > 0 {
q = q.Where("library_id IN ?", ids)
}
res := q.Update("scrape_status", "pending")
if res.Error != nil {
return 0, res.Error
}
reset := int(res.RowsAffected)
if reset > 0 && c.Log != nil {
c.Log.Info("episodic matched rows reset to pending for rescrape",
zap.String("libraries", strings.Join(ids, ",")),
zap.Int("reset", reset))
}
return reset, nil
}
// RepairAndRescrapeAllLibraries 修复并重刮所有媒体库:先从媒体路径中的
// {tmdb-123}/{bangumi-456} 等占位符回填缺失或错误的外部 ID(回填后会把相关
// 行的 scrape_status 重置为 pending),随后逐个媒体库重刮(含 no_match 重试),
// 让此前因空 ID / 脏 ID 无法刮削的媒体重新匹配到正确数据。
func repairRescrapeOptions(values ...ScrapeOptions) ScrapeOptions {
options := ScrapeOptions{RetryNoMatch: true, IncludeMatched: true, ForceRematch: true}
if len(values) > 0 {
options = values[0]
options.RetryNoMatch = true
options.IncludeMatched = true
options.ForceRematch = true
}
if options.EpisodeArtwork == nil {
episodeArtwork := false
options.EpisodeArtwork = &episodeArtwork
}
return options
}
func (c *Container) RepairAndRescrapeAllLibraries(ctx context.Context, options ...ScrapeOptions) (RepairAndRescrapeResult, error) {
var result RepairAndRescrapeResult
if c == nil || c.Repo == nil || c.Repo.DB == nil {
return result, nil
}
scrapeOptions := repairRescrapeOptions(options...)
repaired, err := c.RepairCloudPathMetadata(ctx)
if err != nil {
return result, err
}
result.Repaired = repaired
// 重置全库脏的 matched 剧集行(单集 id 污染整剧字段),让其下方重刮一并修正。
if reset, err := c.resetEpisodicMatchedForRescrape(ctx); err != nil {
return result, err
} else {
result.Reset = reset
}
if c.Scraper == nil || c.Repo.Library == nil {
return result, nil
}
libraries, err := c.Repo.Library.List(ctx)
if err != nil {
return result, err
}
for i := range libraries {
select {
case <-ctx.Done():
return result, ctx.Err()
default:
}
lib := libraries[i]
if !lib.Enabled {
continue
}
result.Libraries++
// retryNoMatch=true:连之前匹配失败的也再试一次,因为这次可能已回填到正确 ID。
scrapeResult, err := c.Scraper.EnrichLibraryDetailedWithOptions(ctx, lib.ID, scrapeOptions)
if err != nil {
if c.Log != nil {
c.Log.Warn("repair rescrape library failed", zap.String("library", lib.ID), zap.Error(err))
}
result.Errors++
continue
}
result.Matched += scrapeResult.Matched
result.Processed += scrapeResult.Processed
result.Errors += scrapeResult.Failed
}
if c.Organizer != nil {
reclassifyResult, err := c.Organizer.ReclassifyMisclassifiedMedia(ctx, MediaCategoryReclassifyOptions{})
if err != nil {
return result, err
}
if reclassifyResult != nil {
result.Reclassified = reclassifyResult.Reclassified
result.Errors += len(reclassifyResult.Errors)
c.invalidateRepairReclassifyCache(ctx, reclassifyResult.Reclassified)
}
}
if c.Log != nil {
c.Log.Info("repair and rescrape all libraries done",
zap.Int("repaired", result.Repaired),
zap.Int("reclassified", result.Reclassified),
zap.Int("libraries", result.Libraries),
zap.Int("matched", result.Matched),
zap.Int("processed", result.Processed),
zap.Int("errors", result.Errors))
}
return result, nil
}
// RepairAndRescrapeLibrary 修复并重刮单个媒体库:先从该库媒体路径中的占位符
// 回填缺失/错误的外部 ID(重置相关行 scrape_status=pending),再对该库重刮
// (含 no_match 重试)。用于「按媒体库」单独触发修复,不影响其它库。
func (c *Container) RepairAndRescrapeLibrary(ctx context.Context, libraryID string, options ...ScrapeOptions) (RepairAndRescrapeResult, error) {
var result RepairAndRescrapeResult
libraryID = strings.TrimSpace(libraryID)
if c == nil || c.Repo == nil || c.Repo.DB == nil || libraryID == "" {
return result, nil
}
scrapeOptions := repairRescrapeOptions(options...)
libraryIDs, err := MergedLibraryIDsForLibrary(ctx, c.Repo, libraryID)
if err != nil {
return result, err
}
repaired, err := c.RepairCloudPathMetadata(ctx, libraryIDs...)
if err != nil {
return result, err
}
result.Repaired = repaired
// 重置该库脏的 matched 剧集行,让下方重刮修正被单集 id 污染的整剧字段。
if reset, err := c.resetEpisodicMatchedForRescrape(ctx, libraryIDs...); err != nil {
return result, err
} else {
result.Reset = reset
}
if c.Scraper == nil {
return result, nil
}
result.Libraries = 1
// retryNoMatch=true:连之前匹配失败的也再试一次,因为这次可能已回填到正确 ID。
scrapeResult, err := c.Scraper.EnrichLibraryDetailedWithOptions(ctx, libraryID, scrapeOptions)
if err != nil {
return result, err
}
result.Matched = scrapeResult.Matched
result.Processed = scrapeResult.Processed
result.Errors = scrapeResult.Failed
if c.Organizer != nil {
reclassifyResult, err := c.Organizer.ReclassifyMisclassifiedMedia(ctx, MediaCategoryReclassifyOptions{LibraryIDs: libraryIDs})
if err != nil {
return result, err
}
if reclassifyResult != nil {
result.Reclassified = reclassifyResult.Reclassified
result.Errors += len(reclassifyResult.Errors)
c.invalidateRepairReclassifyCache(ctx, reclassifyResult.Reclassified)
}
}
if c.Log != nil {
c.Log.Info("repair and rescrape library done",
zap.String("library", libraryID),
zap.Int("repaired", result.Repaired),
zap.Int("reclassified", result.Reclassified),
zap.Int("matched", result.Matched),
zap.Int("processed", result.Processed),
zap.Int("errors", result.Errors))
}
return result, nil
}
func (c *Container) invalidateRepairReclassifyCache(ctx context.Context, changed int) {
if c == nil || c.Cache == nil || changed <= 0 {
return
}
c.Cache.DeletePrefix(ctx, "media:")
c.Cache.DeletePrefix(ctx, "stats:")
}
-202
View File
@@ -1,202 +0,0 @@
// Package service — TMDb discovery (trending / popular).
//
// DiscoverService surfaces curated lists from TMDb so the React home
// page can show "Trending" and "Popular" rails alongside the user's own
// library. All methods gracefully no-op when the TMDb provider is
// disabled.
package service
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"go.uber.org/zap"
)
// DiscoverService talks to TMDb's /trending and /movie/popular endpoints.
type DiscoverService struct {
log *zap.Logger
tmdb *TMDbProvider
client *http.Client
images *ImageProxy
sectionCache *DiscoverSectionCache
}
// NewDiscoverService is the constructor.
func NewDiscoverService(log *zap.Logger, tmdb *TMDbProvider) *DiscoverService {
return &DiscoverService{
log: log,
tmdb: tmdb,
client: NewExternalHTTPClient(15 * time.Second),
sectionCache: NewDiscoverSectionCache(6 * time.Hour),
}
}
// Trending returns the daily trending movies (TMDb /trending/movie/day).
func (d *DiscoverService) Trending(ctx context.Context) ([]Match, error) {
return d.fetch(ctx, "/trending/movie/day")
}
// Popular returns the popular movies list (TMDb /movie/popular).
func (d *DiscoverService) Popular(ctx context.Context) ([]Match, error) {
return d.fetch(ctx, "/movie/popular")
}
// TMDbSection returns one TMDb rail converted to the common external
// discovery shape used by the multi-source Discover page.
func (d *DiscoverService) TMDbSection(ctx context.Context, key string, pages ...int) ([]ExternalMediaResult, error) {
path := tmdbDiscoverPath(key)
if path == "" {
return []ExternalMediaResult{}, nil
}
matches, err := d.Fetch(ctx, path, pages...)
if err != nil {
return nil, err
}
mediaType := "movie"
if strings.Contains(path, "/tv/") {
mediaType = "tv"
}
out := make([]ExternalMediaResult, 0, len(matches))
for _, item := range matches {
out = append(out, ExternalMediaResult{
Source: "tmdb",
MediaType: mediaType,
Title: item.Title,
OriginalName: item.OriginalName,
Overview: item.Overview,
PosterURL: item.PosterURL,
BackdropURL: item.BackdropURL,
Year: item.Year,
Rating: item.Rating,
TMDbID: item.TMDbID,
SubscribeKeyword: buildSubscribeKeyword(item.Title, item.Year),
SubscribeAliases: buildSubscribeAliases(item.Title, item.OriginalName, item.Year),
})
}
return out, nil
}
// fetch is the shared helper that paginates page=1 only — that's all the
// home page needs and it keeps us under TMDb's 50 rps limit.
func (d *DiscoverService) fetch(ctx context.Context, path string) ([]Match, error) {
return d.Fetch(ctx, path)
}
// Fetch is the public entry point used by the multi-section handler.
// It paginates page=1 only — that's all the home page needs and it
// keeps us under TMDb's 50 rps limit.
func (d *DiscoverService) Fetch(ctx context.Context, path string, pages ...int) ([]Match, error) {
if d.tmdb == nil {
return nil, nil
}
// Resolve API key from config or database
apiKey := d.tmdb.resolveAPIKey(ctx)
if apiKey == "" {
return nil, nil
}
base := d.tmdb.resolveBaseURL(ctx)
q := url.Values{}
q.Set("api_key", apiKey)
q.Set("language", "zh-CN")
pageNumber := 1
if len(pages) > 0 && pages[0] > 0 {
pageNumber = pages[0]
}
q.Set("page", strconv.Itoa(pageNumber))
u := base + path + "?" + q.Encode()
type result struct {
ID int `json:"id"`
Title string `json:"title"`
Name string `json:"name"`
OriginalTitle string `json:"original_title"`
OriginalName string `json:"original_name"`
Overview string `json:"overview"`
PosterPath string `json:"poster_path"`
BackdropPath string `json:"backdrop_path"`
ReleaseDate string `json:"release_date"`
FirstAirDate string `json:"first_air_date"`
VoteAverage float32 `json:"vote_average"`
}
type page struct {
Results []result `json:"results"`
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
if err != nil {
return nil, err
}
resp, err := d.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return nil, fmt.Errorf("tmdb %s: %d", path, resp.StatusCode)
}
var p page
if err := json.NewDecoder(resp.Body).Decode(&p); err != nil {
return nil, err
}
out := make([]Match, 0, len(p.Results))
for _, r := range p.Results {
title := r.Title
if title == "" {
title = r.Name
}
m := Match{
TMDbID: r.ID,
Title: title,
OriginalName: firstNonEmpty(r.OriginalTitle, r.OriginalName),
Overview: r.Overview,
Rating: r.VoteAverage,
}
if r.PosterPath != "" {
m.PosterURL = d.tmdb.imgCDN + "/w500" + r.PosterPath
}
if r.BackdropPath != "" {
m.BackdropURL = d.tmdb.imgCDN + "/w1280" + r.BackdropPath
}
date := r.ReleaseDate
if date == "" {
date = r.FirstAirDate
}
if len(date) >= 4 {
_, _ = fmt.Sscanf(date[:4], "%d", &m.Year)
}
out = append(out, m)
}
return out, nil
}
func tmdbDiscoverPath(key string) string {
switch key {
case "tmdb_trending_day", "trending_day":
return "/trending/movie/day"
case "tmdb_trending_week", "trending_week":
return "/trending/movie/week"
case "tmdb_latest_movie", "latest_movie":
return "/movie/now_playing"
case "tmdb_latest_tv", "latest_tv":
return "/tv/on_the_air"
case "tmdb_popular_movie", "popular_movie":
return "/movie/popular"
case "tmdb_popular_tv", "popular_tv":
return "/tv/popular"
case "tmdb_top_rated_movie", "top_rated_movie":
return "/movie/top_rated"
case "tmdb_upcoming_movie", "upcoming_movie":
return "/movie/upcoming"
default:
return ""
}
}
-115
View File
@@ -1,115 +0,0 @@
package service
import (
"context"
"strings"
"sync"
"time"
"go.uber.org/zap"
)
const (
discoverArtworkPrefetchLimit = 384
discoverArtworkPrefetchConcurrency = 8
discoverArtworkPrefetchTimeout = 90 * time.Second
)
func (d *DiscoverService) SetImageProxy(images *ImageProxy) *DiscoverService {
if d != nil {
d.images = images
}
return d
}
func (d *DiscoverService) WarmMatchArtwork(items []Match) int {
return d.warmArtworkURLs(matchArtworkURLs(items))
}
func matchArtworkURLs(items []Match) []string {
urls := make([]string, 0, len(items)*2)
for _, item := range items {
urls = append(urls, item.PosterURL)
}
for _, item := range items {
urls = append(urls, item.BackdropURL)
}
return urls
}
func (d *DiscoverService) WarmExternalArtwork(items []ExternalMediaResult) int {
return d.warmArtworkURLs(externalArtworkURLs(items))
}
func externalArtworkURLs(items []ExternalMediaResult) []string {
urls := make([]string, 0, len(items)*2)
for _, item := range items {
urls = append(urls, item.PosterURL)
}
for _, item := range items {
urls = append(urls, item.BackdropURL)
}
return urls
}
func (d *DiscoverService) warmArtworkURLs(urls []string) int {
if d == nil || d.images == nil || len(urls) == 0 {
return 0
}
pending := uniqueDiscoverArtworkURLs(urls, discoverArtworkPrefetchLimit)
if len(pending) == 0 {
return 0
}
if d.log != nil {
d.log.Debug("discover artwork prefetch scheduled", zap.Int("count", len(pending)))
}
go d.prefetchArtworkURLs(pending)
return len(pending)
}
func uniqueDiscoverArtworkURLs(urls []string, limit int) []string {
if limit <= 0 {
return nil
}
seen := map[string]struct{}{}
out := make([]string, 0, min(len(urls), limit))
for _, raw := range urls {
raw = strings.TrimSpace(raw)
if raw == "" || !isHTTPish(raw) {
continue
}
if _, ok := seen[raw]; ok {
continue
}
seen[raw] = struct{}{}
out = append(out, raw)
if len(out) >= limit {
break
}
}
return out
}
func (d *DiscoverService) prefetchArtworkURLs(urls []string) {
ctx, cancel := context.WithTimeout(context.Background(), discoverArtworkPrefetchTimeout)
defer cancel()
sem := make(chan struct{}, discoverArtworkPrefetchConcurrency)
var wg sync.WaitGroup
for _, raw := range urls {
select {
case <-ctx.Done():
return
case sem <- struct{}{}:
}
wg.Add(1)
go func(raw string) {
defer wg.Done()
defer func() { <-sem }()
if err := d.images.PrefetchRemote(ctx, raw); err != nil && d.log != nil {
d.log.Debug("discover artwork prefetch failed", zap.String("url", raw), zap.Error(err))
}
}(raw)
}
wg.Wait()
}
-135
View File
@@ -1,135 +0,0 @@
package service
import (
"bytes"
"io"
"net/http"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
)
func TestUniqueDiscoverArtworkURLsFiltersDuplicatesAndLimits(t *testing.T) {
urls := uniqueDiscoverArtworkURLs([]string{
"",
"/local/poster.jpg",
"https://image.tmdb.org/t/p/w500/a.jpg",
"https://image.tmdb.org/t/p/w500/a.jpg",
"https://image.tmdb.org/t/p/w500/b.jpg",
"https://image.tmdb.org/t/p/w500/c.jpg",
}, 2)
if len(urls) != 2 {
t.Fatalf("len = %d, want 2: %v", len(urls), urls)
}
if urls[0] != "https://image.tmdb.org/t/p/w500/a.jpg" || urls[1] != "https://image.tmdb.org/t/p/w500/b.jpg" {
t.Fatalf("urls = %v", urls)
}
}
func TestDiscoverArtworkURLsPrioritizePosters(t *testing.T) {
items := []ExternalMediaResult{
{PosterURL: "https://img.example/a-poster.jpg", BackdropURL: "https://img.example/a-backdrop.jpg"},
{PosterURL: "https://img.example/b-poster.jpg", BackdropURL: "https://img.example/b-backdrop.jpg"},
}
urls := externalArtworkURLs(items)
want := []string{
"https://img.example/a-poster.jpg",
"https://img.example/b-poster.jpg",
"https://img.example/a-backdrop.jpg",
"https://img.example/b-backdrop.jpg",
}
if len(urls) != len(want) {
t.Fatalf("len = %d, want %d: %v", len(urls), len(want), urls)
}
for i := range want {
if urls[i] != want[i] {
t.Fatalf("url[%d] = %q, want %q", i, urls[i], want[i])
}
}
}
func TestDiscoverWarmExternalArtworkPrefetchesAndCaches(t *testing.T) {
proxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: filepath.Join(t.TempDir(), "cache")}}, zap.NewNop())
var calls int32
proxy.client = &http.Client{Transport: imageRoundTripFunc(func(req *http.Request) (*http.Response, error) {
atomic.AddInt32(&calls, 1)
return &http.Response{
StatusCode: http.StatusOK,
Status: "200 OK",
Header: http.Header{"Content-Type": []string{"image/jpeg"}},
Body: io.NopCloser(bytes.NewReader(testJPEG)),
Request: req,
}, nil
})}
discover := NewDiscoverService(zap.NewNop(), nil).SetImageProxy(proxy)
poster := "https://image.tmdb.org/t/p/w500/discover-poster.jpg"
backdrop := "https://image.tmdb.org/t/p/w1280/discover-backdrop.jpg"
queued := discover.WarmExternalArtwork([]ExternalMediaResult{
{Title: "A", PosterURL: poster, BackdropURL: backdrop},
{Title: "B", PosterURL: poster},
{Title: "Local", PosterURL: "/media/poster.jpg"},
})
if queued != 2 {
t.Fatalf("queued = %d, want 2", queued)
}
for _, raw := range []string{poster, backdrop} {
_, cachePath, _, err := proxy.remoteImageCachePaths(raw)
if err != nil {
t.Fatal(err)
}
waitForDiscoverArtworkCache(t, &calls, 2, cachePath, raw)
}
callsAfterCache := atomic.LoadInt32(&calls)
queued = discover.WarmExternalArtwork([]ExternalMediaResult{{Title: "Cached", PosterURL: poster, BackdropURL: backdrop}})
if queued != 2 {
t.Fatalf("queued cached = %d, want 2", queued)
}
time.Sleep(150 * time.Millisecond)
if got := atomic.LoadInt32(&calls); got != callsAfterCache {
t.Fatalf("cached prefetch should not call upstream again: got %d want %d", got, callsAfterCache)
}
}
func TestDiscoverWarmArtworkNoImageProxyIsNoop(t *testing.T) {
discover := NewDiscoverService(zap.NewNop(), nil)
if got := discover.WarmMatchArtwork([]Match{{PosterURL: "https://image.tmdb.org/t/p/w500/a.jpg"}}); got != 0 {
t.Fatalf("queued = %d, want 0 without image proxy", got)
}
}
func TestTMDbDiscoverPathIncludesLatestSections(t *testing.T) {
cases := map[string]string{
"tmdb_latest_movie": "/movie/now_playing",
"latest_movie": "/movie/now_playing",
"tmdb_latest_tv": "/tv/on_the_air",
"latest_tv": "/tv/on_the_air",
"tmdb_upcoming_movie": "/movie/upcoming",
}
for key, want := range cases {
if got := tmdbDiscoverPath(key); got != want {
t.Fatalf("tmdbDiscoverPath(%q) = %q, want %q", key, got, want)
}
}
}
func waitForDiscoverArtworkCache(t *testing.T, calls *int32, wantCalls int32, cachePath, raw string) {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if atomic.LoadInt32(calls) >= wantCalls {
if _, err := os.Stat(cachePath); err == nil {
return
}
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("expected cached artwork %q after %d upstream calls: %v", raw, atomic.LoadInt32(calls), os.ErrNotExist)
}
-104
View File
@@ -1,104 +0,0 @@
package service
import (
"fmt"
"sync"
"time"
)
// DiscoverSectionCache keeps the last good discover rail in memory so a slow
// upstream provider does not turn a populated page into empty rows.
type DiscoverSectionCache struct {
ttl time.Duration
mu sync.RWMutex
entries map[string]discoverSectionCacheEntry
}
type discoverSectionCacheEntry struct {
items []ExternalMediaResult
storedAt time.Time
}
func NewDiscoverSectionCache(ttl time.Duration) *DiscoverSectionCache {
if ttl <= 0 {
ttl = 6 * time.Hour
}
return &DiscoverSectionCache{
ttl: ttl,
entries: map[string]discoverSectionCacheEntry{},
}
}
func (d *DiscoverService) RememberSection(key string, page int, items []ExternalMediaResult) {
if d == nil || d.sectionCache == nil || len(items) == 0 {
return
}
d.sectionCache.Set(key, page, items)
}
func (d *DiscoverService) CachedSection(key string, page int) ([]ExternalMediaResult, bool) {
if d == nil || d.sectionCache == nil {
return nil, false
}
return d.sectionCache.Get(key, page)
}
func (c *DiscoverSectionCache) Set(key string, page int, items []ExternalMediaResult) {
if c == nil || key == "" || page < 1 || len(items) == 0 {
return
}
c.mu.Lock()
defer c.mu.Unlock()
c.entries[discoverSectionCacheKey(key, page)] = discoverSectionCacheEntry{
items: cloneExternalMediaResults(items),
storedAt: time.Now(),
}
}
func (c *DiscoverSectionCache) Get(key string, page int) ([]ExternalMediaResult, bool) {
if c == nil || key == "" || page < 1 {
return nil, false
}
c.mu.RLock()
entry, ok := c.entries[discoverSectionCacheKey(key, page)]
c.mu.RUnlock()
if !ok || time.Since(entry.storedAt) > c.ttl || len(entry.items) == 0 {
return nil, false
}
return cloneExternalMediaResults(entry.items), true
}
func discoverSectionCacheKey(key string, page int) string {
return fmt.Sprintf("%s:%d", key, page)
}
func cloneExternalMediaResults(items []ExternalMediaResult) []ExternalMediaResult {
out := make([]ExternalMediaResult, len(items))
for i, item := range items {
out[i] = item
out[i].SubscribeAliases = cloneStrings(item.SubscribeAliases)
out[i].MissingEpisodes = cloneInts(item.MissingEpisodes)
out[i].Languages = cloneStrings(item.Languages)
out[i].Countries = cloneStrings(item.Countries)
out[i].Genres = cloneStrings(item.Genres)
}
return out
}
func cloneStrings(items []string) []string {
if len(items) == 0 {
return nil
}
out := make([]string, len(items))
copy(out, items)
return out
}
func cloneInts(items []int) []int {
if len(items) == 0 {
return nil
}
out := make([]int, len(items))
copy(out, items)
return out
}
@@ -1,44 +0,0 @@
package service
import (
"testing"
"time"
)
func TestDiscoverSectionCacheReturnsClone(t *testing.T) {
cache := NewDiscoverSectionCache(time.Hour)
cache.Set("douban_hot_movie", 1, []ExternalMediaResult{{
Title: "第一部",
SubscribeAliases: []string{"别名"},
MissingEpisodes: []int{1},
Languages: []string{"zh"},
}})
got, ok := cache.Get("douban_hot_movie", 1)
if !ok || len(got) != 1 || got[0].Title != "第一部" {
t.Fatalf("cached section = %#v, %v", got, ok)
}
got[0].Title = "被修改"
got[0].SubscribeAliases[0] = "别名被改"
got[0].MissingEpisodes[0] = 9
got[0].Languages[0] = "en"
again, ok := cache.Get("douban_hot_movie", 1)
if !ok ||
again[0].Title != "第一部" ||
again[0].SubscribeAliases[0] != "别名" ||
again[0].MissingEpisodes[0] != 1 ||
again[0].Languages[0] != "zh" {
t.Fatalf("cache should return a clone, got %#v", again)
}
}
func TestDiscoverSectionCacheExpires(t *testing.T) {
cache := NewDiscoverSectionCache(time.Nanosecond)
cache.Set("tmdb_latest_movie", 1, []ExternalMediaResult{{Title: "旧数据"}})
time.Sleep(time.Millisecond)
if got, ok := cache.Get("tmdb_latest_movie", 1); ok || len(got) != 0 {
t.Fatalf("expired cache should miss, got %#v", got)
}
}
-90
View File
@@ -1,90 +0,0 @@
// Package service — Douban discovery rails.
package service
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
)
// Discover returns public Douban movie/TV rails. Douban does not require a
// formal API key here; these are the same public web endpoints the site uses.
func (d *DoubanProvider) Discover(ctx context.Context, key string, pages ...int) ([]ExternalMediaResult, error) {
doubanType := "movie"
tag := "热门"
switch key {
case "douban_hot_movie":
doubanType = "movie"
tag = "热门"
case "douban_top_movie":
doubanType = "movie"
tag = "高分"
case "douban_hot_tv":
doubanType = "tv"
tag = "热门"
default:
return []ExternalMediaResult{}, nil
}
q := url.Values{}
q.Set("type", doubanType)
q.Set("tag", tag)
q.Set("sort", "recommend")
q.Set("page_limit", "24")
pageNumber := 1
if len(pages) > 0 && pages[0] > 0 {
pageNumber = pages[0]
}
q.Set("page_start", strconv.Itoa((pageNumber-1)*24))
u := "https://movie.douban.com/j/search_subjects?" + q.Encode()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
if err != nil {
return nil, err
}
d.setHeaders(req)
resp, err := d.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return nil, fmt.Errorf("douban discover: %d", resp.StatusCode)
}
var page struct {
Subjects []struct {
ID string `json:"id"`
Title string `json:"title"`
Rate string `json:"rate"`
Cover string `json:"cover"`
URL string `json:"url"`
} `json:"subjects"`
}
if err := json.NewDecoder(resp.Body).Decode(&page); err != nil {
return nil, err
}
out := make([]ExternalMediaResult, 0, len(page.Subjects))
mediaType := "movie"
if doubanType == "tv" {
mediaType = "tv"
}
for _, subject := range page.Subjects {
if strings.TrimSpace(subject.Title) == "" {
continue
}
rating, _ := strconv.ParseFloat(subject.Rate, 32)
out = append(out, ExternalMediaResult{
Source: "douban",
MediaType: mediaType,
Title: subject.Title,
PosterURL: subject.Cover,
Rating: float32(rating),
DoubanID: subject.ID,
SubscribeKeyword: subject.Title,
SubscribeAliases: buildSubscribeAliases(subject.Title, "", 0),
})
}
return out, nil
}
-36
View File
@@ -1,36 +0,0 @@
package service
import (
"context"
"time"
"go.uber.org/zap"
)
const activeDownloadSnapshotFallbackAge = 2 * time.Minute
func (d *DownloadService) ActiveDownloadPaths(ctx context.Context) []string {
if d == nil {
return 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 download clients and has no fresh snapshot", zap.Error(err))
}
}
return activeDownloadPathCandidates(live, d.downloadPathMappings(ctx))
}
func (d *DownloadService) downloadPathMappings(ctx context.Context) map[string]string {
mappings := map[string]string{
"/var/apps/qBittorrent/shares/qBittorrent/Download": "/downloads",
"/data/qBittorrent/downloads": "/downloads",
"/downloads/qBittorrent": "/downloads",
}
for clientPrefix, localPrefix := range d.userPathMappings(ctx) {
mappings[clientPrefix] = localPrefix
}
return mappings
}
-90
View File
@@ -1,90 +0,0 @@
// Package service 定义下载适配器接口和通用数据结构。
package service
import (
"context"
"time"
)
// DownloadAdapter 定义下载客户端的统一接口。
// 所有下载客户端(qBittorrent / Transmission / Aria2)必须实现此接口。
type DownloadAdapter interface {
// Initialize 使用配置初始化客户端连接。
Initialize(ctx context.Context, cfg DownloadClientConfig) error
// Ping 测试客户端连接是否可用。
Ping(ctx context.Context) error
// AddTorrent 通过 URL(磁力链接或种子 URL)添加下载任务。
AddTorrent(ctx context.Context, url, savePath string) (string, error)
// AddMagnet 通过磁力链接添加下载任务。
AddMagnet(ctx context.Context, magnet, savePath string) (string, error)
// Pause 暂停指定下载任务。
Pause(ctx context.Context, hash string) error
// Resume 恢复指定下载任务。
Resume(ctx context.Context, hash string) error
// Remove 移除指定下载任务,deleteFiles 控制是否同时删除文件。
Remove(ctx context.Context, hash string, deleteFiles bool) error
// List 列出所有或过滤后的种子任务。filter 可为空字符串表示全部。
List(ctx context.Context, filter string) ([]TorrentInfo, error)
// GetInfo 获取指定种子的详细信息。
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"`
ContentPath string `json:"content_path"`
CompletionOn int64 `json:"completion_on"`
}
// DownloadClientConfig 是下载客户端的连接配置。
type DownloadClientConfig struct {
Host string `json:"host"`
Username string `json:"username"`
Password string `json:"password"`
Extra map[string]string `json:"extra,omitempty"`
}
// AdapterFactory 根据客户端类型创建适配器实例。
func AdapterFactory(clientType string) DownloadAdapter {
switch clientType {
case "qbittorrent":
return NewQBitAdapter()
case "transmission":
return NewTransmissionAdapter()
case "aria2":
return NewAria2Adapter()
default:
return nil
}
}
@@ -1,219 +0,0 @@
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)
}
}
-232
View File
@@ -1,232 +0,0 @@
package service
import (
"context"
"errors"
"path"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// DownloadTaskMeta carries public display metadata for a download. It is
// deliberately separate from the private torrent URL so API responses never
// need to expose tracker tokens.
type DownloadTaskMeta struct {
SubscriptionID string
Title string
PosterURL string
BackdropURL string
Overview string
MediaType string
MediaCategory string
SourceCategory string
OriginalName string
OriginalLanguage string
Year int
Rating float32
Genres string
AllowExistingLibrary bool
}
type downloadAddRequest struct {
title string
savePath string
qbitCategory string
meta DownloadTaskMeta
}
// AddDownload accepts a magnet URL / HTTP URL and persists a tracking row.
func (d *DownloadService) AddDownload(ctx context.Context, userID, urlStr, savePath string) (*model.DownloadTask, error) {
return d.AddDownloadWithMeta(ctx, userID, urlStr, savePath, DownloadTaskMeta{})
}
func (d *DownloadService) AddDownloadWithMeta(ctx context.Context, userID, urlStr, savePath string, meta DownloadTaskMeta) (*model.DownloadTask, error) {
req, err := d.prepareDownloadAdd(ctx, urlStr, savePath, meta)
if err != nil {
return nil, err
}
if !req.meta.AllowExistingLibrary && d.localMediaAlreadyExists(ctx, req.title) {
return nil, ErrMediaAlreadyInLibrary
}
if existing, ok := d.findExistingDownloadTask(ctx, req); ok {
d.linkExistingDownloadTaskToSubscription(ctx, existing, req)
return existing, ErrDownloadAlreadyExists
}
_ = d.ReloadConfig(ctx)
target, err := d.defaultDownloadTarget(ctx)
if err != nil {
return nil, d.defaultDownloaderNotConfiguredError(ctx)
}
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
}
if strings.TrimSpace(req.meta.SubscriptionID) != "" {
return task, nil
}
return task, ErrDownloadAlreadyExists
}
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, target, externalID)
}
return nil, err
}
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) {
if urlStr == "" {
return downloadAddRequest{}, errors.New("empty url")
}
title := strings.TrimSpace(meta.Title)
if title == "" {
title = publicDownloadTitle(urlStr)
meta.Title = title
}
autoClassify := downloadSmartClassifyEnabled(ctx, d.repo, d.organizer)
savePath, resolvedCategory := d.resolveDownloadSavePath(ctx, savePath, meta, autoClassify)
if !autoClassify {
meta.MediaCategory = ""
} else if strings.TrimSpace(meta.MediaCategory) == "" {
meta.MediaCategory = resolvedCategory
}
return downloadAddRequest{
title: title,
savePath: savePath,
qbitCategory: strings.TrimSpace(meta.MediaCategory),
meta: meta,
}, nil
}
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 {
return d.addTorrentFileToTarget(ctx, data, name, req, target)
} else {
siteFetchErr = err
}
}
externalID, err := d.addTorrentURLToTarget(ctx, urlStr, req, target)
if err != nil {
return externalID, joinTorrentFetchError(err, siteFetchErr)
}
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) {
if strings.TrimSpace(explicitSavePath) != "" {
if !autoClassify {
return explicitSavePath, ""
}
return explicitSavePath, strings.TrimSpace(meta.MediaCategory)
}
base := downloadDefaultSaveRoot(ctx, d.repo)
if strings.TrimSpace(base) == "" {
return "", strings.TrimSpace(meta.MediaCategory)
}
mediaType := normalizeMediaType(meta.MediaType, meta.Title, meta.SourceCategory)
category := strings.TrimSpace(meta.MediaCategory)
if category == "" {
category = classifyMediaCategory(mediaClassifyInput{
MediaType: mediaType,
Title: meta.Title,
Category: meta.SourceCategory,
}, downloadCategoryMap(d.organizer))
}
if !autoClassify || category == "" {
return base, ""
}
return downloadSavePathCategoryRoot(base, sanitizeFilename(category)), category
}
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)
}
t := &model.DownloadTask{
UserID: userID,
SubscriptionID: strings.TrimSpace(meta.SubscriptionID),
DownloadClientID: target.clientID,
ExternalID: strings.TrimSpace(externalID),
Source: target.typ,
URL: urlStr,
Title: title,
PosterURL: meta.PosterURL,
BackdropURL: meta.BackdropURL,
Overview: meta.Overview,
SavePath: savePath,
MediaType: meta.MediaType,
MediaCategory: meta.MediaCategory,
OriginalName: meta.OriginalName,
OriginalLanguage: meta.OriginalLanguage,
Year: meta.Year,
Rating: meta.Rating,
Genres: meta.Genres,
Status: "queued",
AllowExistingLibrary: meta.AllowExistingLibrary,
}
if err := d.repo.Download.Create(ctx, t); err != nil {
return nil, err
}
return t, nil
}
@@ -1,160 +0,0 @@
package service
import (
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestAddDownloadWithMetaAutoClassifiesSavePathAndQBitCategory(t *testing.T) {
var addCalls int32
var gotSavePath string
var gotCategory string
qb := 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/info":
if atomic.LoadInt32(&addCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"auto123","name":"声生不息 S01E01","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
if err := r.ParseMultipartForm(1024 * 1024); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
gotSavePath = r.FormValue("savepath")
gotCategory = r.FormValue("category")
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Setting.Set(t.Context(), "qbittorrent.savepath", "/downloads"); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=%E5%A3%B0%E7%94%9F%E4%B8%8D%E6%81%AF+S01E01", "", DownloadTaskMeta{
Title: "声生不息 S01E01",
SourceCategory: "综艺",
})
if err != nil {
t.Fatal(err)
}
wantPath := filepath.Join("/downloads", "综艺")
if task.SavePath != wantPath {
t.Fatalf("task save path = %q, want %q", task.SavePath, wantPath)
}
if gotSavePath != wantPath {
t.Fatalf("qb savepath = %q, want %q", gotSavePath, wantPath)
}
if gotCategory != "综艺" {
t.Fatalf("qb category = %q, want 综艺", gotCategory)
}
}
func TestDownloadSavePathCategoryRootKeepsWindowsClientSeparators(t *testing.T) {
if got := downloadSavePathCategoryRoot(`F:\downloads`, "国产剧"); got != `F:\downloads\国产剧` {
t.Fatalf("downloadSavePathCategoryRoot() = %q, want Windows qB path", got)
}
if got := downloadSavePathCategoryRoot(`F:\downloads\国产剧`, "国产剧"); got != `F:\downloads\国产剧` {
t.Fatalf("downloadSavePathCategoryRoot() duplicated category: %q", got)
}
if got := downloadSavePathCategoryRoot(`/downloads`, "国产剧"); got != filepath.Join(`/downloads`, "国产剧") {
t.Fatalf("downloadSavePathCategoryRoot() = %q, want local path", got)
}
}
func TestTranslateClientPathMapsWindowsQBitPathToContainerDownloadPath(t *testing.T) {
root := t.TempDir()
containerDownloads := filepath.Join(root, "downloads")
want := filepath.Join(containerDownloads, "国产剧", "Show.S01E01.mkv")
if err := os.MkdirAll(filepath.Dir(want), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(want, []byte("episode"), 0o644); err != nil {
t.Fatal(err)
}
got := translateClientPath(`F:\downloads\国产剧\Show.S01E01.mkv`, map[string]string{
`F:\downloads`: containerDownloads,
})
if got != want {
t.Fatalf("translateClientPath() = %q, want %q", got, want)
}
}
func TestAddDownloadWithMetaCanDisableAutoClassifiedSavePath(t *testing.T) {
var addCalls int32
var gotSavePath string
var gotCategory string
qb := 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/info":
if atomic.LoadInt32(&addCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"auto456","name":"声生不息 S01E01","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
if err := r.ParseMultipartForm(1024 * 1024); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
gotSavePath = r.FormValue("savepath")
gotCategory = r.FormValue("category")
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Setting.Set(t.Context(), "qbittorrent.savepath", "/downloads"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), DownloadSmartClassifySettingKey, "false"); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:ffffffffffffffffffffffffffffffffffffffff&dn=%E5%A3%B0%E7%94%9F%E4%B8%8D%E6%81%AF+S01E01", "", DownloadTaskMeta{
Title: "声生不息 S01E01",
SourceCategory: "综艺",
})
if err != nil {
t.Fatal(err)
}
if task.SavePath != "/downloads" {
t.Fatalf("task save path = %q, want /downloads", task.SavePath)
}
if gotSavePath != "/downloads" {
t.Fatalf("qb savepath = %q, want /downloads", gotSavePath)
}
if gotCategory != "" {
t.Fatalf("qb category = %q, want empty", gotCategory)
}
}
-196
View File
@@ -1,196 +0,0 @@
package service
import (
"context"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func (d *DownloadService) findExistingDownloadTask(ctx context.Context, req downloadAddRequest) (*model.DownloadTask, bool) {
key := downloadTaskIdentityKey(req.title)
if key == "" || d == nil || d.repo == nil || d.repo.Download == nil {
return nil, false
}
rows, err := d.repo.Download.List(ctx)
if err != nil {
return nil, false
}
subscriptionID := strings.TrimSpace(req.meta.SubscriptionID)
for i := range rows {
if subscriptionID != "" {
if !downloadTaskBlocksReadd(rows[i].Status) {
continue
}
if !downloadTaskInSubscriptionScope(rows[i], req) {
continue
}
if !d.subscriptionDownloadTaskStillLive(ctx, rows[i]) {
continue
}
} else if !downloadTaskBlocksDuplicate(rows[i].Status) {
continue
}
current := downloadTaskIdentityKey(rows[i].Title)
if downloadTaskCoversAddRequest(rows[i].Title, req) || current == key {
return &rows[i], true
}
}
return nil, false
}
func (d *DownloadService) subscriptionDownloadTaskStillLive(ctx context.Context, row model.DownloadTask) bool {
live, ok := d.liveTorrentSnapshot(30 * time.Second)
if !ok && d != nil {
var err error
live, err = d.listLiveTorrents(ctx, "")
if err != nil && len(live) == 0 {
return true
}
ok = true
}
if !ok {
return true
}
for _, torrent := range live {
if downloadTaskMatchesLiveTorrent(row, torrent) {
return true
}
}
return false
}
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
}
req := downloadAddRequest{
title: row.Title,
savePath: row.SavePath,
meta: DownloadTaskMeta{
SubscriptionID: row.SubscriptionID,
},
}
if downloadTaskCoversAddRequest(torrentName, req) {
return true
}
rowKey := downloadTaskIdentityKey(row.Title)
torrentKey := downloadTaskIdentityKey(torrentName)
if rowKey != "" && torrentKey != "" {
return rowKey == torrentKey
}
if len(episodeRefsFromTitle(row.Title)) > 0 || len(episodeRefsFromTitle(torrentName)) > 0 {
return false
}
rowTorrentKey := normalizeTorrentName(row.Title)
liveTorrentKey := normalizeTorrentName(torrentName)
return rowTorrentKey != "" && rowTorrentKey == liveTorrentKey
}
func downloadTaskCoversAddRequest(existing string, req downloadAddRequest) bool {
if subscriptionRequestHasExplicitEpisodes(req) {
return downloadExplicitEpisodesCoverRequest(existing, req.title)
}
return downloadTitleCoversRequest(existing, req.title)
}
func subscriptionRequestHasExplicitEpisodes(req downloadAddRequest) bool {
return strings.TrimSpace(req.meta.SubscriptionID) != "" && len(episodeRefsFromTitle(req.title)) > 0
}
func downloadExplicitEpisodesCoverRequest(existing, requested string) bool {
current := parseDownloadMediaIdentity(existing)
want := parseDownloadMediaIdentity(requested)
if current.TitleKey == "" || want.TitleKey == "" {
return false
}
if current.TitleKey != want.TitleKey {
return false
}
if current.Year > 0 && want.Year > 0 && current.Year != want.Year {
return false
}
if len(current.Episodes) == 0 || len(want.Episodes) == 0 {
return false
}
currentEpisodes := map[string]struct{}{}
for _, ref := range current.Episodes {
currentEpisodes[episodeKey(ref.Season, ref.Episode)] = struct{}{}
}
for _, ref := range want.Episodes {
if _, ok := currentEpisodes[episodeKey(ref.Season, ref.Episode)]; !ok {
return false
}
}
return true
}
func downloadTaskInSubscriptionScope(row model.DownloadTask, req downloadAddRequest) bool {
subscriptionID := strings.TrimSpace(req.meta.SubscriptionID)
if subscriptionID == "" {
return true
}
rowSubscriptionID := strings.TrimSpace(row.SubscriptionID)
if rowSubscriptionID != "" {
return rowSubscriptionID == subscriptionID
}
rowSavePath := strings.TrimSpace(row.SavePath)
requestSavePath := strings.TrimSpace(req.savePath)
if rowSavePath == "" || requestSavePath == "" {
return false
}
return sameOrChildPath(rowSavePath, requestSavePath) || sameOrChildPath(requestSavePath, rowSavePath)
}
func (d *DownloadService) findLiveTorrentByIdentity(ctx context.Context, downloadURL string, req downloadAddRequest) (QBitTorrent, bool) {
query := downloadTaskIdentityKey(req.title)
requestHash := torrentURLInfoHash(downloadURL)
if query == "" && requestHash == "" {
return QBitTorrent{}, false
}
live, err := d.listLiveTorrents(ctx, "")
if err != nil {
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 torrent, true
}
current := downloadTaskIdentityKey(torrent.Name)
if current == "" {
continue
}
if current == query {
return torrent, true
}
}
return QBitTorrent{}, false
}
func torrentInDownloadRequestScope(torrent QBitTorrent, req downloadAddRequest) bool {
if strings.TrimSpace(req.meta.SubscriptionID) == "" {
return true
}
requestSavePath := strings.TrimSpace(req.savePath)
torrentSavePath := strings.TrimSpace(torrent.SavePath)
if requestSavePath == "" || torrentSavePath == "" {
return false
}
return sameOrChildPath(torrentSavePath, requestSavePath) || sameOrChildPath(requestSavePath, torrentSavePath)
}
-332
View File
@@ -1,332 +0,0 @@
package service
import (
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestAddDownloadWithMetaDoesNotDedupRangeAgainstSingleEpisodeTask(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old",
Title: "Archives The Nanyang Mystery 2026 S01E07 2160p WEB-DL",
SavePath: "/downloads/tv",
Status: "completed",
Progress: 1,
}); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:cdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcd&dn=Archives+The+Nanyang+Mystery+2026+S01E07-S01E08", "/downloads", DownloadTaskMeta{
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because existing task covers only E07", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaDoesNotDedupRangeAgainstSeasonOnlyTask(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
UserID: "u1",
SubscriptionID: "sub-nanyang",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old-season",
Title: "Archives The Nanyang Mystery 2026 S01 2160p WEB-DL",
SavePath: "/downloads/tv",
Status: "completed",
Progress: 1,
}); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:abababababababababababababababababababab&dn=Archives+The+Nanyang+Mystery+2026+S01E09-E10", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "sub-nanyang",
Title: "Archives The Nanyang Mystery 2026 S01E09-E10 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because season-only task does not prove E09-E10 exists", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaDoesNotDedupSubscriptionRangeAgainstCompletePackTask(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
UserID: "u1",
SubscriptionID: "sub-nanyang",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old-complete",
Title: "Archives The Nanyang Mystery 2026 S01 Complete 2160p WEB-DL",
SavePath: "/downloads/tv",
Status: "completed",
Progress: 1,
}); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:fafafafafafafafafafafafafafafafafafafafa&dn=Archives+The+Nanyang+Mystery+2026+S01E29-E33", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "sub-nanyang",
Title: "Archives The Nanyang Mystery 2026 S01E29-E33 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because complete-pack history does not prove missing range exists", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaTracksExistingQBTorrentForSubscription(t *testing.T) {
torrentData := []byte("d4:infod4:name7:fixtureee")
hash := torrentInfoHash(torrentData)
if hash == "" {
t.Fatal("expected fixture info hash")
}
var addCalls int32
qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/fixture.torrent":
w.Header().Set("Content-Type", "application/x-bittorrent")
_, _ = w.Write(torrentData)
case "/api/v2/auth/login":
_, _ = w.Write([]byte("Ok."))
case "/api/v2/torrents/info":
_, _ = w.Write([]byte(`[{"hash":"` + hash + `","name":"Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL","save_path":"/downloads/tv"}]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", qb.URL+"/fixture.torrent", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "sub-nanyang",
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want tracking task for existing qB torrent", err)
}
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)
}
}
func TestDownloadTitleCoversRequestKeepsCompletePackDedup(t *testing.T) {
if !downloadTitleCoversRequest("Archives The Nanyang Mystery 2026 S01 Complete 2160p WEB-DL", "Archives The Nanyang Mystery 2026 S01E09-E10 2160p WEB-DL") {
t.Fatal("complete pack should cover requested episode range")
}
if downloadTitleCoversRequest("Archives The Nanyang Mystery 2026 S01 2160p WEB-DL", "Archives The Nanyang Mystery 2026 S01E09-E10 2160p WEB-DL") {
t.Fatal("season-only title must not cover requested episode range")
}
}
func TestAddDownloadWithMetaScopesSubscriptionDedupBySubscriptionOrSavePath(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
for _, existing := range []model.DownloadTask{
{
UserID: "u1",
SubscriptionID: "other-subscription",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old-sub",
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
SavePath: "/downloads/other",
Status: "completed",
Progress: 1,
},
{
UserID: "u1",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old-manual",
Title: "Archives The Nanyang Mystery 2026 S01E09-S01E10 2160p WEB-DL",
SavePath: "/downloads/archive",
Status: "completed",
Progress: 1,
},
} {
row := existing
if err := repos.Download.Create(t.Context(), &row); err != nil {
t.Fatal(err)
}
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:efefefefefefefefefefefefefefefefefefefef&dn=Archives+The+Nanyang+Mystery+2026+S01E07-S01E08", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "current-subscription",
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because old task is outside current subscription scope", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaRequeuesStaleSubscriptionTaskMissingFromQB(t *testing.T) {
var addCalls int32
qb := 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/info":
if atomic.LoadInt32(&addCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"newhash","name":"Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
UserID: "u1",
SubscriptionID: "sub-nanyang",
Source: "qbittorrent",
URL: "https://pt.example/download?id=stale",
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
SavePath: "/downloads/tv",
Status: "queued",
Progress: 0,
}); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
svc.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"})
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:bcbcbcbcbcbcbcbcbcbcbcbcbcbcbcbcbcbcbcbc&dn=Archives+The+Nanyang+Mystery+2026+S01E07-S01E08", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "sub-nanyang",
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want stale task ignored and candidate queued", err)
}
if task == nil {
t.Fatal("task = nil, want requeued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
@@ -1,220 +0,0 @@
package service
import (
"errors"
"fmt"
"net/http"
"net/http/httptest"
"path/filepath"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestAddDownloadWithMetaSkipsExistingLocalMovieBeforeQBAdd(t *testing.T) {
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
if err := db.Create(&model.Media{
Title: "Inception",
Path: "/media/movies/Inception (2010)/Inception (2010).mkv",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Inception+2010+1080p", "/downloads", DownloadTaskMeta{
Title: "Inception 2010 1080p WEB-DL",
})
if !errors.Is(err, ErrMediaAlreadyInLibrary) {
t.Fatalf("err = %v, want ErrMediaAlreadyInLibrary", err)
}
if task != nil {
t.Fatalf("task = %#v, want nil because local media already exists", task)
}
rows, err := repos.Download.List(t.Context())
if err != nil {
t.Fatal(err)
}
if len(rows) != 0 {
t.Fatalf("download rows = %d, want 0", len(rows))
}
}
func TestAddDownloadWithMetaSkipsExistingLocalEpisodeBeforeQBAdd(t *testing.T) {
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
if err := db.Create(&model.Media{
Title: "Some Show",
Path: "/media/tv/Some Show/Season 01/Some Show - S01E01.mkv",
SeasonNum: 1,
EpisodeNum: 1,
}).Error; err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Some+Show+S01E01", "/downloads", DownloadTaskMeta{
Title: "Some Show S01E01 2160p WEB-DL",
})
if !errors.Is(err, ErrMediaAlreadyInLibrary) {
t.Fatalf("err = %v, want ErrMediaAlreadyInLibrary", err)
}
if task != nil {
t.Fatalf("task = %#v, want nil because local episode already exists", task)
}
rows, err := repos.Download.List(t.Context())
if err != nil {
t.Fatal(err)
}
if len(rows) != 0 {
t.Fatalf("download rows = %d, want 0", len(rows))
}
}
func TestAddDownloadWithMetaQueuesEpisodeRangeWhenOnlyPartlyInLibrary(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := db.Create(&model.Media{
Title: "Archives The Nanyang Mystery",
Path: "/media/tv/Archives The Nanyang Mystery/Season 01/Archives The Nanyang Mystery - S01E07.mkv",
SeasonNum: 1,
EpisodeNum: 7,
}).Error; err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:abababababababababababababababababababab&dn=Archives+The+Nanyang+Mystery+2026+S01E07-S01E08", "/downloads", DownloadTaskMeta{
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because E08 is missing", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaSkipsEpisodeRangeWhenFullyInLibrary(t *testing.T) {
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
for _, episode := range []int{7, 8} {
if err := db.Create(&model.Media{
Title: "Archives The Nanyang Mystery",
Path: filepath.Join("/media/tv/Archives The Nanyang Mystery/Season 01", fmt.Sprintf("Archives The Nanyang Mystery - S01E%02d.mkv", episode)),
SeasonNum: 1,
EpisodeNum: episode,
}).Error; err != nil {
t.Fatal(err)
}
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:babababababababababababababababababababa&dn=Archives+The+Nanyang+Mystery+2026+S01E07-S01E08", "/downloads", DownloadTaskMeta{
Title: "Archives The Nanyang Mystery 2026 S01E07-S01E08 2160p WEB-DL",
})
if !errors.Is(err, ErrMediaAlreadyInLibrary) {
t.Fatalf("err = %v, want ErrMediaAlreadyInLibrary", err)
}
if task != nil {
t.Fatalf("task = %#v, want nil", task)
}
}
func TestAddDownloadWithMetaQueuesExplicitEpisodeWhenOnlySeriesPackInLibrary(t *testing.T) {
var addCalls int32
qb := 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/info":
http.Error(w, "temporary list unavailable", http.StatusInternalServerError)
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
if err := db.Create(&model.Media{
Title: "Archives The Nanyang Mystery S01 Complete",
Path: "/media/tv/Archives The Nanyang Mystery/Season 01/Archives The Nanyang Mystery S01 Complete.mkv",
}).Error; err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:fefefefefefefefefefefefefefefefefefefefe&dn=Archives+The+Nanyang+Mystery+2026+S01E29-E33", "/downloads/tv", DownloadTaskMeta{
SubscriptionID: "sub-nanyang",
Title: "Archives The Nanyang Mystery 2026 S01E29-E33 2160p WEB-DL",
})
if err != nil {
t.Fatalf("AddDownloadWithMeta returned %v, want queued because explicit missing episodes are not proven by a pack row", err)
}
if task == nil {
t.Fatal("task = nil, want queued task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
func TestAddDownloadWithMetaSkipsExistingLocalEpisodeWithReleaseGroup(t *testing.T) {
db := newServiceTestDB(t, &model.Media{}, &model.DownloadTask{}, &model.Setting{}, &model.DownloadClient{})
repos := repository.New(db)
if err := db.Create(&model.Media{
Title: "凡人修仙传",
Path: "/media/动漫/国漫/凡人修仙传/Season 01/凡人修仙传 - S01E146.mkv",
SeasonNum: 1,
EpisodeNum: 146,
}).Error; err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=%5BMagicStar%5D+%E5%87%A1%E4%BA%BA%E4%BF%AE%E4%BB%99%E4%BC%A0+%E5%B9%B4%E7%95%AA+-+146+%5B1080p%5D", "/downloads", DownloadTaskMeta{
Title: "[MagicStar] 凡人修仙传 年番 - 146 [1080p][WEB-DL]",
})
if !errors.Is(err, ErrMediaAlreadyInLibrary) {
t.Fatalf("err = %v, want ErrMediaAlreadyInLibrary", err)
}
if task != nil {
t.Fatalf("task = %#v, want nil", task)
}
rows, err := repos.Download.List(t.Context())
if err != nil {
t.Fatal(err)
}
if len(rows) != 0 {
t.Fatalf("download rows = %d, want 0", len(rows))
}
}
-92
View File
@@ -1,92 +0,0 @@
package service
import (
"context"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func (d *DownloadService) localMediaAlreadyExists(ctx context.Context, title string) bool {
rows, ok := d.localMediaAvailabilityRows(ctx, title)
if !ok {
return false
}
return localMediaRowsMatchDownloadTitle(title, rows)
}
func (d *DownloadService) localMediaAvailabilityRows(ctx context.Context, title string) ([]model.Media, bool) {
if d == nil || d.repo == nil || d.repo.DB == nil {
return nil, false
}
if !d.repo.DB.Migrator().HasTable(&model.Media{}) {
return nil, false
}
queries := localAvailabilityTitleCandidates(title)
if len(queries) == 0 {
return nil, false
}
var rows []model.Media
db := d.repo.DB.WithContext(ctx).Model(&model.Media{})
for i, query := range queries {
like := "%" + query + "%"
clause := "title LIKE ? OR original_name LIKE ? OR path LIKE ?"
if i == 0 {
db = db.Where(clause, like, like, like)
} else {
db = db.Or(clause, like, like, like)
}
}
if err := db.
Order("season_num asc, episode_num asc, created_at desc").
Limit(200).
Find(&rows).Error; err != nil || len(rows) == 0 {
return nil, false
}
return rows, true
}
func localMediaRowsMatchDownloadTitle(title string, rows []model.Media) bool {
wanted := episodeRefsFromTitle(title)
if len(wanted) == 0 {
return true
}
existing := map[string]struct{}{}
hasSeriesPack := false
for _, row := range rows {
rowSeason, rowEpisode := localMediaRowSeasonEpisode(row)
if rowEpisode > 0 {
existing[episodeKey(rowSeason, rowEpisode)] = struct{}{}
continue
}
if rowEpisode <= 0 && isSeriesPackTitle(row.Title+" "+row.OriginalName+" "+row.Path) {
hasSeriesPack = true
}
}
if hasSeriesPack {
return len(wanted) == 0
}
for _, ref := range wanted {
if _, ok := existing[episodeKey(ref.Season, ref.Episode)]; !ok {
return false
}
}
return true
}
func localMediaRowSeasonEpisode(row model.Media) (int, int) {
rowSeason := row.SeasonNum
rowEpisode := row.EpisodeNum
if rowSeason <= 0 || rowEpisode <= 0 {
parsedSeason, parsedEpisode := ParseEpisode(row.Path)
if rowSeason <= 0 {
rowSeason = parsedSeason
}
if rowEpisode <= 0 {
rowEpisode = parsedEpisode
}
}
if rowSeason <= 0 {
rowSeason = 1
}
return rowSeason, rowEpisode
}
@@ -1,48 +0,0 @@
package service
import (
"context"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func (d *DownloadService) linkExistingDownloadTaskToSubscription(ctx context.Context, task *model.DownloadTask, req downloadAddRequest) {
subscriptionID := strings.TrimSpace(req.meta.SubscriptionID)
if d == nil || d.repo == nil || d.repo.DB == nil || task == nil || subscriptionID == "" || strings.TrimSpace(task.ID) == "" {
return
}
updates := map[string]any{}
if strings.TrimSpace(task.SubscriptionID) == "" {
updates["subscription_id"] = subscriptionID
task.SubscriptionID = subscriptionID
}
if strings.TrimSpace(task.MediaType) == "" && strings.TrimSpace(req.meta.MediaType) != "" {
updates["media_type"] = req.meta.MediaType
task.MediaType = req.meta.MediaType
}
if strings.TrimSpace(task.MediaCategory) == "" && strings.TrimSpace(req.meta.MediaCategory) != "" {
updates["media_category"] = req.meta.MediaCategory
task.MediaCategory = req.meta.MediaCategory
}
if strings.TrimSpace(task.PosterURL) == "" && strings.TrimSpace(req.meta.PosterURL) != "" {
updates["poster_url"] = req.meta.PosterURL
task.PosterURL = req.meta.PosterURL
}
if strings.TrimSpace(task.BackdropURL) == "" && strings.TrimSpace(req.meta.BackdropURL) != "" {
updates["backdrop_url"] = req.meta.BackdropURL
task.BackdropURL = req.meta.BackdropURL
}
if strings.TrimSpace(task.Overview) == "" && strings.TrimSpace(req.meta.Overview) != "" {
updates["overview"] = req.meta.Overview
task.Overview = req.meta.Overview
}
if !task.AllowExistingLibrary && req.meta.AllowExistingLibrary {
updates["allow_existing_library"] = true
task.AllowExistingLibrary = true
}
if len(updates) == 0 {
return
}
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", task.ID).Updates(updates).Error
}
-149
View File
@@ -1,149 +0,0 @@
package service
import (
"errors"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestPublicDownloadTitleUsesMagnetDisplayName(t *testing.T) {
got := publicDownloadTitle("magnet:?xt=urn:btih:abc&dn=%E6%B5%8B%E8%AF%95%E5%BD%B1%E7%89%87")
if got != "测试影片" {
t.Fatalf("publicDownloadTitle = %q, want %q", got, "测试影片")
}
}
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{
Name: "qB test",
Type: "qbittorrent",
Host: baseURL,
Username: "admin",
Password: "admin",
IsDefault: true,
Enabled: true,
}); err != nil {
t.Fatalf("create default qB client: %v", err)
}
if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil {
t.Fatalf("mark download clients managed: %v", err)
}
}
func TestAddDownloadWithMetaSkipsExistingTaskBeforeQBAdd(t *testing.T) {
var addCalls int32
qb := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
existing := &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old&passkey=old",
Title: "Some Show S01E01 1080p",
SavePath: "/downloads/tv",
Status: "completed",
Progress: 1,
}
if err := repos.Download.Create(t.Context(), existing); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
svc.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"})
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "https://pt.example/download?id=new&passkey=new", "/downloads/tv", DownloadTaskMeta{
Title: "Some Show S01E01 2160p WEB-DL",
})
if !errors.Is(err, ErrDownloadAlreadyExists) {
t.Fatalf("err = %v, want ErrDownloadAlreadyExists", err)
}
if task == nil || task.ID != existing.ID {
t.Fatalf("task = %#v, want existing task %#v", task, existing)
}
if got := atomic.LoadInt32(&addCalls); got != 0 {
t.Fatalf("qb add calls = %d, want 0", got)
}
}
func TestAddDownloadWithMetaSkipsUserDeletedTaskBeforeQBAdd(t *testing.T) {
var addCalls int32
qb := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "qbittorrent.url", qb.URL); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.username", "admin"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.password", "admin"); err != nil {
t.Fatal(err)
}
existing := &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "https://pt.example/download?id=old&passkey=old",
Title: "User Deleted Show S01E01 1080p",
SavePath: "/downloads/tv",
Status: "deleted",
}
if err := repos.Download.Create(t.Context(), existing); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "https://pt.example/download?id=new&passkey=new", "/downloads/tv", DownloadTaskMeta{
Title: "User Deleted Show S01E01 1080p WEB-DL",
})
if !errors.Is(err, ErrDownloadAlreadyExists) {
t.Fatalf("err = %v, want ErrDownloadAlreadyExists", err)
}
if task == nil || task.ID != existing.ID {
t.Fatalf("task = %#v, want existing task %#v", task, existing)
}
if got := atomic.LoadInt32(&addCalls); got != 0 {
t.Fatalf("qb add calls = %d, want 0", got)
}
}
-83
View File
@@ -1,83 +0,0 @@
package service
import (
"context"
"os"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func downloadDefaultSaveRoot(ctx context.Context, repo *repository.Container) string {
if repo != nil && repo.Setting != nil {
if base, _ := repo.Setting.Get(ctx, "qbittorrent.savepath"); strings.TrimSpace(base) != "" {
return strings.TrimSpace(base)
}
}
for _, key := range []string{"MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "MEDIASTATION_DOWNLOAD_DIR"} {
if value := strings.TrimSpace(os.Getenv(key)); value != "" {
return value
}
}
return ""
}
func downloadSmartClassifyEnabled(ctx context.Context, repo *repository.Container, organizer *OrganizerService) bool {
if repo != nil && repo.Setting != nil {
val, err := repo.Setting.Get(ctx, DownloadSmartClassifySettingKey)
if err == nil && val != "" {
return parseBoolSetting(val, true)
}
val, err = repo.Setting.Get(ctx, "organizer.smart_classify")
if err == nil && parseBoolSetting(val, false) {
return true
}
}
if organizer != nil && organizer.cfg != nil && organizer.cfg.Organizer.SmartClassify {
return true
}
return true
}
func downloadCategoryMap(organizer *OrganizerService) map[string]string {
if organizer == nil {
return nil
}
return organizer.categoryMap()
}
func downloadSavePathCategoryRoot(root, category string) string {
root = strings.TrimSpace(root)
category = strings.TrimSpace(category)
if root == "" || category == "" {
return root
}
if isWindowsStyleClientPath(root) {
cleanRoot := strings.ReplaceAll(root, "/", `\`)
cleanRoot = strings.TrimRight(cleanRoot, `\`)
if windowsPathBaseEqual(cleanRoot, category) {
return cleanRoot
}
return cleanRoot + `\` + category
}
return categoryRoot(root, category)
}
func isWindowsStyleClientPath(path string) bool {
path = strings.TrimSpace(path)
return (len(path) >= 2 && isASCIIAlpha(path[0]) && path[1] == ':') ||
strings.HasPrefix(path, `\\`)
}
func windowsPathBaseEqual(path, base string) bool {
path = strings.TrimRight(strings.ReplaceAll(strings.TrimSpace(path), "/", `\`), `\`)
base = strings.Trim(strings.TrimSpace(base), `\/`)
if path == "" || base == "" {
return false
}
idx := strings.LastIndex(path, `\`)
if idx >= 0 {
path = path[idx+1:]
}
return strings.EqualFold(path, base)
}
@@ -1,115 +0,0 @@
package service
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"regexp"
"strconv"
"strings"
)
var downloadClientEndpointPattern = regexp.MustCompile(`^https?://(?:[A-Za-z0-9.-]+|\[[0-9A-Fa-f:.]+\])(?::[0-9]{1,5})?(?:/[A-Za-z0-9._~%!$&'()*+,;=:@/-]*)?$`)
func NormalizeDownloadClientHost(clientType, raw string) (string, error) {
return normalizeDownloadClientEndpoint(clientType, raw)
}
func normalizeDownloadClientEndpoint(clientType, raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", errors.New("host required")
}
if strings.ContainsAny(raw, "\r\n\t") {
return "", errors.New("host contains invalid control characters")
}
if !strings.Contains(raw, "://") {
raw = "http://" + raw
}
if !downloadClientEndpointPattern.MatchString(raw) {
return "", errors.New("host must be a valid http(s) URL without username, query, or fragment")
}
parsed, err := url.Parse(raw)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return "", errors.New("host must be a valid http(s) URL")
}
scheme := strings.ToLower(parsed.Scheme)
if scheme != "http" && scheme != "https" {
return "", errors.New("host only supports http or https")
}
if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
return "", errors.New("host must not include username, query, or fragment")
}
if strings.TrimSpace(parsed.Hostname()) == "" {
return "", errors.New("host must include a hostname")
}
if port := parsed.Port(); port != "" {
n, err := strconv.Atoi(port)
if err != nil || n < 1 || n > 65535 {
return "", errors.New("host port must be between 1 and 65535")
}
}
if err := validateDownloadClientPath(clientType, parsed.Path); err != nil {
return "", err
}
parsed.Scheme = scheme
parsed.Path = strings.TrimRight(parsed.Path, "/")
parsed.RawPath = ""
parsed.RawQuery = ""
parsed.Fragment = ""
return strings.TrimRight(parsed.String(), "/"), nil
}
func validateDownloadClientPath(clientType, rawPath string) error {
rawPath = strings.TrimSpace(rawPath)
if rawPath == "" || rawPath == "/" {
return nil
}
for _, segment := range strings.Split(rawPath, "/") {
if segment == "." || segment == ".." {
return errors.New("host path must not contain traversal segments")
}
}
switch clientType {
case "qbittorrent", "aria2", "transmission":
return nil
default:
return fmt.Errorf("unsupported client type %q", clientType)
}
}
func downloadClientRPCURL(clientType, host string) (string, error) {
base, err := normalizeDownloadClientEndpoint(clientType, host)
if err != nil {
return "", err
}
u, err := url.Parse(base)
if err != nil {
return "", err
}
switch clientType {
case "aria2":
if !strings.HasSuffix(strings.ToLower(u.Path), "/jsonrpc") {
u.Path = strings.TrimRight(u.Path, "/") + "/jsonrpc"
}
case "transmission":
if !strings.Contains(strings.ToLower(u.Path), "/rpc") {
u.Path = strings.TrimRight(u.Path, "/") + "/transmission/rpc"
}
case "qbittorrent":
default:
return "", fmt.Errorf("unsupported client type %q", clientType)
}
return u.String(), nil
}
func newDownloadClientHTTPRequest(ctx context.Context, method, endpoint string, body io.Reader) (*http.Request, error) {
endpoint = strings.TrimSpace(endpoint)
if !downloadClientEndpointPattern.MatchString(endpoint) {
return nil, errors.New("download client endpoint failed safety validation")
}
return http.NewRequestWithContext(ctx, method, endpoint, body)
}
-270
View File
@@ -1,270 +0,0 @@
// Package service — download client (qBittorrent / Aria2 / Transmission)
// configuration. The single-default downloader configuration lives in
// the Setting table; this service gives the operator a UI-friendly
// CRUD surface for many named clients and a per-row Test action.
package service
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// DownloadClientService persists model.DownloadClient rows.
type DownloadClientService struct {
log *zap.Logger
repo *repository.Container
client *http.Client
}
// NewDownloadClientService is the constructor.
func NewDownloadClientService(log *zap.Logger, repo *repository.Container) *DownloadClientService {
return &DownloadClientService{
log: log,
repo: repo,
client: NewInternalHTTPClient(10 * time.Second),
}
}
// DownloadClientInput is the create / update payload.
type DownloadClientInput struct {
Name string `json:"name" binding:"required"`
Type string `json:"type" binding:"required"`
Host string `json:"host" binding:"required"`
Username string `json:"username,omitempty"`
Password string `json:"password,omitempty"`
IsDefault bool `json:"is_default"`
Enabled bool `json:"enabled"`
}
// List returns every configured client.
func (s *DownloadClientService) List(ctx context.Context) ([]model.DownloadClient, error) {
return s.repo.DownloadClient.List(ctx)
}
// Create inserts a new client.
func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInput) (*model.DownloadClient, error) {
normalized, err := normalizeDownloadClientInput(in)
if err != nil {
return nil, err
}
s.markManaged(ctx)
c := &model.DownloadClient{
Name: normalized.Name,
Type: normalized.Type,
Host: normalized.Host,
Username: normalized.Username,
Password: normalized.Password,
IsDefault: normalized.IsDefault,
Enabled: normalized.Enabled,
}
if !c.IsDefault && c.Enabled {
if currentDefault, err := s.repo.DownloadClient.FindDefault(ctx); err == nil && currentDefault == nil {
if enabled, err := s.repo.DownloadClient.ListEnabled(ctx); err == nil && len(enabled) == 0 {
c.IsDefault = true
}
}
}
if normalized.IsDefault {
_ = s.repo.DownloadClient.ClearDefault(ctx)
}
if err := s.repo.DownloadClient.Create(ctx, c); err != nil {
return nil, err
}
return c, nil
}
// Update applies a patch.
func (s *DownloadClientService) Update(ctx context.Context, id string, in DownloadClientInput) (*model.DownloadClient, error) {
normalized, err := normalizeDownloadClientInput(in)
if err != nil {
return nil, err
}
s.markManaged(ctx)
patch := map[string]any{
"name": normalized.Name,
"type": normalized.Type,
"host": normalized.Host,
"username": normalized.Username,
"is_default": normalized.IsDefault,
"enabled": normalized.Enabled,
}
// Only overwrite the password when the caller actually sent one.
if normalized.Password != "" {
patch["password"] = normalized.Password
}
// Fetch existing row, apply patch via Save
existing, err := s.repo.DownloadClient.FindByID(ctx, id)
if err != nil {
return nil, err
}
if existing == nil {
return nil, errors.New("client not found")
}
if normalized.IsDefault {
_ = s.repo.DownloadClient.ClearDefault(ctx)
}
existing.Name = patch["name"].(string)
existing.Type = patch["type"].(string)
existing.Host = patch["host"].(string)
existing.Username = patch["username"].(string)
existing.IsDefault = patch["is_default"].(bool)
existing.Enabled = patch["enabled"].(bool)
if pw, ok := patch["password"]; ok {
existing.Password = pw.(string)
}
if err := s.repo.DownloadClient.Update(ctx, existing); err != nil {
return nil, err
}
s.clearLegacyQBitConnectionIfNoDefault(ctx)
return s.repo.DownloadClient.FindByID(ctx, id)
}
// Delete removes one client.
func (s *DownloadClientService) Delete(ctx context.Context, id string) error {
s.markManaged(ctx)
if err := s.repo.DownloadClient.Delete(ctx, id); err != nil {
return err
}
s.clearLegacyQBitConnectionIfNoDefault(ctx)
return nil
}
// Test verifies that the client's WebUI is reachable. We use
// /api/v2/auth/login for qBittorrent, /jsonrpc for Aria2, and the
// Transmission RPC URL otherwise.
func (s *DownloadClientService) Test(ctx context.Context, id string) error {
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
c, err := s.repo.DownloadClient.FindByID(ctx, id)
if err != nil {
return err
}
if c == nil {
return errors.New("client not found")
}
switch c.Type {
case "qbittorrent":
return qbitLogin(ctx, s.client, c.Host, c.Username, c.Password)
case "aria2", "transmission":
endpoint, err := downloadClientRPCURL(c.Type, c.Host)
if err != nil {
return err
}
req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return err
}
resp, err := s.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 500 {
return fmt.Errorf("%s returned %d", c.Type, resp.StatusCode)
}
return nil
}
return fmt.Errorf("unsupported client type %q", c.Type)
}
// Aria2GlobalStats issues a JSON-RPC `aria2.getGlobalStat` call against
// the first enabled aria2 client. Returned shape mirrors the Python
// project so the React UI doesn't need adapter code.
func (s *DownloadClientService) Aria2GlobalStats(ctx context.Context, clientID string) (map[string]any, error) {
c, err := s.repo.DownloadClient.FindByID(ctx, clientID)
if err != nil {
return nil, err
}
if c == nil || c.Type != "aria2" {
return nil, errors.New("aria2 client not found")
}
endpoint, err := downloadClientRPCURL("aria2", c.Host)
if err != nil {
return nil, err
}
payload := fmt.Sprintf(
`{"jsonrpc":"2.0","id":"x","method":"aria2.getGlobalStat","params":["token:%s"]}`,
c.Password,
)
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, endpoint,
strings.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
resp, err := s.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return nil, fmt.Errorf("aria2 returned %d", resp.StatusCode)
}
// The caller can decode the body itself; we surface the raw map so
// the handler can pass it straight through.
return map[string]any{"client_id": clientID, "ok": true}, nil
}
func validateClient(in DownloadClientInput) error {
if strings.TrimSpace(in.Name) == "" {
return errors.New("name required")
}
if strings.TrimSpace(in.Host) == "" {
return errors.New("host required")
}
switch in.Type {
case "qbittorrent", "aria2", "transmission":
default:
return fmt.Errorf("unsupported client type %q", in.Type)
}
return nil
}
func normalizeDownloadClientInput(in DownloadClientInput) (DownloadClientInput, error) {
in.Name = strings.TrimSpace(in.Name)
in.Type = strings.TrimSpace(in.Type)
in.Host = strings.TrimSpace(in.Host)
in.Username = strings.TrimSpace(in.Username)
if err := validateClient(in); err != nil {
return in, err
}
if !strings.Contains(in.Host, "://") {
in.Host = "http://" + in.Host
}
normalized, err := normalizeDownloadClientEndpoint(in.Type, in.Host)
if err != nil {
return in, err
}
in.Host = normalized
return in, nil
}
func (s *DownloadClientService) markManaged(ctx context.Context) {
if s == nil || s.repo == nil || s.repo.Setting == nil {
return
}
_ = s.repo.Setting.Set(ctx, settingDownloadClientsManaged, "true")
}
func (s *DownloadClientService) clearLegacyQBitConnectionIfNoDefault(ctx context.Context) {
if s == nil || s.repo == nil || s.repo.DownloadClient == nil || s.repo.Setting == nil {
return
}
defaultClient, err := s.repo.DownloadClient.FindDefault(ctx)
if err != nil || defaultClient != nil {
return
}
_ = s.repo.Setting.Set(ctx, "qbittorrent.url", "")
_ = s.repo.Setting.Set(ctx, "qbittorrent.username", "")
_ = s.repo.Setting.Set(ctx, "qbittorrent.password", "")
}
-248
View File
@@ -1,248 +0,0 @@
package service
import (
"net/http"
"net/http/httptest"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestDownloadClientCreateNormalizesHostAndClearsDefault(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
svc := NewDownloadClientService(zap.NewNop(), repos)
first, err := svc.Create(t.Context(), DownloadClientInput{
Name: "qB old",
Type: "qbittorrent",
Host: "http://127.0.0.1:8080/",
IsDefault: true,
Enabled: true,
})
if err != nil {
t.Fatal(err)
}
second, err := svc.Create(t.Context(), DownloadClientInput{
Name: "qB NAS",
Type: "qbittorrent",
Host: "172.17.0.1:8085",
IsDefault: true,
Enabled: true,
})
if err != nil {
t.Fatal(err)
}
if second.Host != "http://172.17.0.1:8085" {
t.Fatalf("host = %q, want normalized http URL", second.Host)
}
refreshedFirst, err := repos.DownloadClient.FindByID(t.Context(), first.ID)
if err != nil {
t.Fatal(err)
}
if refreshedFirst == nil || refreshedFirst.IsDefault {
t.Fatalf("old default should be cleared, got %#v", refreshedFirst)
}
refreshedSecond, err := repos.DownloadClient.FindByID(t.Context(), second.ID)
if err != nil {
t.Fatal(err)
}
if refreshedSecond == nil || !refreshedSecond.IsDefault {
t.Fatalf("new default should be active, got %#v", refreshedSecond)
}
}
func TestDownloadClientCreateMakesFirstEnabledClientDefault(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
svc := NewDownloadClientService(zap.NewNop(), repos)
client, err := svc.Create(t.Context(), DownloadClientInput{
Name: "qB",
Type: "qbittorrent",
Host: "127.0.0.1:8080",
Enabled: true,
})
if err != nil {
t.Fatal(err)
}
if !client.IsDefault {
t.Fatalf("first enabled client should become default: %#v", client)
}
}
func TestDownloadClientRejectsUnsupportedHostScheme(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
svc := NewDownloadClientService(zap.NewNop(), repository.New(db))
if _, err := svc.Create(t.Context(), DownloadClientInput{
Name: "bad",
Type: "qbittorrent",
Host: "ftp://127.0.0.1:8080",
Enabled: true,
}); err == nil {
t.Fatal("expected unsupported scheme error")
}
}
func TestDownloadClientRejectsUnsafeEndpointParts(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
svc := NewDownloadClientService(zap.NewNop(), repository.New(db))
for _, host := range []string{
"http://user:pass@127.0.0.1:6800",
"http://127.0.0.1:6800/jsonrpc?target=http://169.254.169.254",
"http://127.0.0.1:6800/jsonrpc#fragment",
"http://127.0.0.1:70000",
"file:///etc/passwd",
} {
if _, err := svc.Create(t.Context(), DownloadClientInput{
Name: "bad",
Type: "aria2",
Host: host,
Enabled: true,
}); err == nil {
t.Fatalf("Create allowed unsafe host %q", host)
}
}
}
func TestDownloadClientRPCURLAppendsExpectedPath(t *testing.T) {
cases := []struct {
clientType string
host string
want string
}{
{"aria2", "127.0.0.1:6800", "http://127.0.0.1:6800/jsonrpc"},
{"aria2", "http://nas.local:6800/rpc", "http://nas.local:6800/rpc/jsonrpc"},
{"transmission", "http://nas.local:9091", "http://nas.local:9091/transmission/rpc"},
{"transmission", "http://nas.local:9091/transmission/rpc", "http://nas.local:9091/transmission/rpc"},
}
for _, tc := range cases {
got, err := downloadClientRPCURL(tc.clientType, tc.host)
if err != nil {
t.Fatalf("downloadClientRPCURL(%q, %q) error: %v", tc.clientType, tc.host, err)
}
if got != tc.want {
t.Fatalf("downloadClientRPCURL(%q, %q) = %q, want %q", tc.clientType, tc.host, got, tc.want)
}
}
}
func TestAria2AdapterRejectsUnsafeHostBeforeHTTPRequest(t *testing.T) {
adapter := NewAria2Adapter()
called := false
adapter.client = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
called = true
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})}
if err := adapter.Initialize(t.Context(), DownloadClientConfig{
Host: "http://user:pass@127.0.0.1:6800",
Password: "secret",
}); err == nil {
t.Fatal("expected unsafe host error")
}
if called {
t.Fatal("unsafe aria2 host should be rejected before any HTTP request")
}
}
func TestAria2AdapterUsesNormalizedRPCURL(t *testing.T) {
var gotPath string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
adapter := NewAria2Adapter()
if err := adapter.Initialize(t.Context(), DownloadClientConfig{
Host: server.URL,
Password: "secret",
}); err != nil {
t.Fatal(err)
}
if gotPath != "/jsonrpc" {
t.Fatalf("aria2 request path = %q, want /jsonrpc", gotPath)
}
}
func TestDownloadClientDeleteClearsLegacyQBitConnectionWhenNoDefault(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
for key, value := range map[string]string{
"qbittorrent.url": "http://127.0.0.1:8080",
"qbittorrent.username": "admin",
"qbittorrent.password": "admin",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
svc := NewDownloadClientService(zap.NewNop(), repos)
row, err := svc.Create(t.Context(), DownloadClientInput{
Name: "qB",
Type: "qbittorrent",
Host: "http://127.0.0.1:8080",
IsDefault: true,
Enabled: true,
})
if err != nil {
t.Fatal(err)
}
if err := svc.Delete(t.Context(), row.ID); err != nil {
t.Fatal(err)
}
for _, key := range []string{"qbittorrent.url", "qbittorrent.username", "qbittorrent.password"} {
value, err := repos.Setting.Get(t.Context(), key)
if err != nil {
t.Fatal(err)
}
if value != "" {
t.Fatalf("%s = %q, want cleared", key, value)
}
}
}
func TestDownloadClientUpdateClearsLegacyQBitConnectionWhenDefaultDisabled(t *testing.T) {
db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "qbittorrent.url", "http://127.0.0.1:8080"); err != nil {
t.Fatal(err)
}
svc := NewDownloadClientService(zap.NewNop(), repos)
row, err := svc.Create(t.Context(), DownloadClientInput{
Name: "qB",
Type: "qbittorrent",
Host: "http://127.0.0.1:8080",
IsDefault: true,
Enabled: true,
})
if err != nil {
t.Fatal(err)
}
if _, err := svc.Update(t.Context(), row.ID, DownloadClientInput{
Name: "qB",
Type: "qbittorrent",
Host: "http://127.0.0.1:8080",
IsDefault: false,
Enabled: false,
}); err != nil {
t.Fatal(err)
}
value, err := repos.Setting.Get(t.Context(), "qbittorrent.url")
if err != nil {
t.Fatal(err)
}
if value != "" {
t.Fatalf("qbittorrent.url = %q, want cleared", value)
}
}
-216
View File
@@ -1,216 +0,0 @@
package service
import (
"context"
"errors"
"os"
"strings"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// onTorrentComplete handles a torrent that just finished downloading.
// It organizes the completed torrent payload directly. Relying on existing
// Media rows is too late for freshly-downloaded files: they usually have not
// been scanned into the library yet.
func (d *DownloadService) onTorrentComplete(ctx context.Context, torrent QBitTorrent) {
taskRow, hasTask := d.completedTorrentTask(ctx, torrent)
d.notifyDownloadComplete(ctx, torrent, taskRow)
if d.organizer == nil {
return
}
// 仅当显式开启 organizer.auto_after_download / organize.auto 时才在下载完成后整理。
// 之前的代码错误地把 organizer.smart_classify 也当成"自动整理"开关,
// 让操作员只想启用"分类子目录"就被动触发了文件 move。
autoOrganize := d.downloadAutoOrganizeEnabled(ctx)
if !autoOrganize {
d.log.Info("download completed, auto-organize disabled", zap.String("hash", torrent.Hash))
return
}
source := d.completedTorrentSource(ctx, torrent)
if source == "" {
d.log.Warn("download completed but payload path is not accessible",
zap.String("hash", torrent.Hash),
zap.String("name", torrent.Name),
zap.String("save_path", torrent.SavePath),
zap.String("content_path", torrent.ContentPath))
return
}
allowReplace := hasTask && taskRow.AllowExistingLibrary
d.runCompletedTorrentOrganize(ctx, torrent, taskRow, source, allowReplace)
}
func (d *DownloadService) runCompletedTorrentOrganize(ctx context.Context, torrent QBitTorrent, task *model.DownloadTask, source string, allowReplace bool) {
d.log.Info("download completed, triggering directory organize",
zap.String("hash", torrent.Hash),
zap.String("name", torrent.Name),
zap.String("source", source),
zap.Bool("allow_replace_existing", allowReplace))
resWrap, err := d.ensureOrganizePipeline().Run(ctx, OrganizePipelineRequest{
Scope: OrganizeScopeDirectory,
Trigger: OrganizeTriggerDownload,
TaskName: d.downloadOrganizeTaskName(torrent, allowReplace),
SourcePath: source,
MediaType: downloadTaskMediaType(task),
MediaCategory: firstNonEmpty(downloadTaskMediaCategory(task), torrent.Category),
AllowReplace: allowReplace,
})
if err != nil {
if errors.Is(err, ErrUnsupportedOrganizeSource) {
d.markCompletedTorrentCatchupRecorded(context.Background(), torrent)
d.log.Warn("auto organize skipped unsupported completed torrent",
zap.String("hash", torrent.Hash),
zap.String("source", source),
zap.Error(err))
return
}
d.log.Error("auto organize completed torrent failed",
zap.String("hash", torrent.Hash),
zap.String("source", source),
zap.Error(err))
return
}
res := resWrap.Result
if res == nil {
res = &OrganizeResult{}
}
d.markCompletedTorrentCatchupRecorded(context.Background(), torrent)
d.log.Info("auto organize completed torrent finished",
zap.String("hash", torrent.Hash),
zap.String("source", source),
zap.String("dest", firstNonEmpty(res.DestPath, "")),
zap.Int("organized", res.Organized),
zap.Int("replaced", res.Replaced),
zap.Int("skipped", res.Skipped),
zap.Int("scrapes", len(res.Scrapes)),
zap.Int("errors", len(res.Errors)))
}
func (d *DownloadService) downloadOrganizeTaskName(torrent QBitTorrent, allowReplace bool) string {
name := strings.TrimSpace(torrent.Name)
if name == "" {
name = "下载完成自动整理"
}
if allowReplace {
name += "(允许洗版)"
}
return name
}
func (d *DownloadService) ensureOrganizePipeline() *OrganizePipelineService {
if d.organizePipeline != nil {
return d.organizePipeline
}
return NewOrganizePipelineService(d.log, d.repo, d.organizer, d.scanner, d.tasks)
}
func (d *DownloadService) completedTorrentTask(ctx context.Context, torrent QBitTorrent) (*model.DownloadTask, bool) {
if d == nil || d.repo == nil || d.repo.Download == nil {
return nil, false
}
rows, err := d.repo.Download.List(ctx)
if err != nil || len(rows) == 0 {
return nil, false
}
taskByKey := tasksByTorrentIdentity(rows)
if task, ok := findMatchingTaskForTorrent(torrent, taskByKey); ok {
return &task, true
}
if strings.TrimSpace(torrent.ContentPath) != "" {
pathTorrent := torrent
pathTorrent.Name = downloaderPathBase(torrent.ContentPath)
if task, ok := findMatchingTaskForTorrent(pathTorrent, taskByKey); ok {
return &task, true
}
}
return nil, false
}
func downloadTaskMediaType(task *model.DownloadTask) string {
if task == nil {
return ""
}
return strings.TrimSpace(task.MediaType)
}
func downloadTaskMediaCategory(task *model.DownloadTask) string {
if task == nil {
return ""
}
return strings.TrimSpace(task.MediaCategory)
}
// DownloadPathMappingsSettingKey 允许用户自定义「下载器路径 → 本程序路径」
// 映射,每行一条,格式 `客户端路径=本地路径`(也接受 `=>` 或单个 `:` 分隔)。
// qBittorrent 与本程序常在不同容器/主机里,对同一份数据看到的路径不同;
// 此前映射表是写死的三条猜测,对不上时整理静默失败。
const DownloadPathMappingsSettingKey = "download.path_mappings"
func (d *DownloadService) completedTorrentSource(ctx context.Context, torrent QBitTorrent) string {
mappings := d.downloadPathMappings(ctx)
for _, candidate := range []string{
torrent.ContentPath,
downloaderPayloadPath(torrent.SavePath, torrent.Name),
} {
clean := strings.TrimSpace(candidate)
if clean == "" || clean == "." {
continue
}
// 尝试直接访问或路径映射
if translated := translateClientPath(clean, mappings); translated != "" {
return translated
}
// 复用 compose 注入的 MEDIASTATION_DOWNLOAD_DIR/MEDIA_DIR 宿主机↔容器
// 映射(与媒体库路径换算同一套规则),覆盖「qB 跑在宿主机、
// 本程序在容器里」的最常见部署形态。
for _, mapped := range mappedPathCandidates(clean) {
if mapped == clean {
continue
}
if _, err := os.Stat(mapped); err == nil {
return mapped
}
}
}
return ""
}
// userPathMappings 解析用户配置的下载器路径映射。
func (d *DownloadService) userPathMappings(ctx context.Context) map[string]string {
out := map[string]string{}
if d == nil || d.repo == nil || d.repo.Setting == nil {
return out
}
raw, err := d.repo.Setting.Get(ctx, DownloadPathMappingsSettingKey)
if err != nil {
return out
}
for _, line := range strings.Split(raw, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
var from, to string
switch {
case strings.Contains(line, "=>"):
parts := strings.SplitN(line, "=>", 2)
from, to = parts[0], parts[1]
case strings.Contains(line, "="):
parts := strings.SplitN(line, "=", 2)
from, to = parts[0], parts[1]
case strings.Count(line, ":") == 1:
parts := strings.SplitN(line, ":", 2)
from, to = parts[0], parts[1]
default:
continue
}
from = strings.TrimSpace(from)
to = strings.TrimSpace(to)
if from != "" && to != "" {
out[from] = to
}
}
return out
}
@@ -1,327 +0,0 @@
package service
import (
"context"
"crypto/sha1"
"fmt"
"math"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// completedTorrentCatchupWindow 限定重启补整理只覆盖最近完成的种子,
// 防止每次启动都把全部历史种子重新过一遍整理流程。
const completedTorrentCatchupWindow = 24 * time.Hour
const completedTorrentCatchupSettingPrefix = "download.auto_organized."
const completedTorrentNotifySettingPrefix = "download.completed_notified."
func (d *DownloadService) downloadAutoOrganizeEnabled(ctx context.Context) bool {
if d == nil || d.repo == nil || d.repo.Setting == nil {
return false
}
if v, err := d.repo.Setting.Get(ctx, "organizer.auto_after_download"); err == nil && parseBoolSetting(v, false) {
return true
}
if v, err := d.repo.Setting.Get(ctx, "organize.auto"); err == nil && parseBoolSetting(v, false) {
return true
}
return false
}
// recentlyCompletedTorrent 报告该种子是否在补整理时间窗内完成。
// qBittorrent 未提供 completion_on 时保守地返回 false。
func recentlyCompletedTorrent(torrent QBitTorrent, now time.Time) bool {
if torrent.CompletionOn <= 0 {
return false
}
completed := time.Unix(torrent.CompletionOn, 0)
return now.Sub(completed) <= completedTorrentCatchupWindow
}
func qbitTorrentCompleted(torrent QBitTorrent) bool {
if torrent.Progress < 1 {
return false
}
state := strings.ToLower(strings.TrimSpace(torrent.State))
switch state {
case "completed", "complete", "seeding", "uploading", "stalledup", "pausedup", "queuedup", "forcedup":
return true
default:
return false
}
}
func (d *DownloadService) completedTorrentCatchupRecorded(ctx context.Context, torrent QBitTorrent) bool {
if d == nil || d.repo == nil || d.repo.Setting == nil {
return false
}
key := completedTorrentCatchupSettingKey(torrent)
if key == "" {
return false
}
value, err := d.repo.Setting.Get(ctx, key)
if err != nil {
return false
}
return parseBoolSetting(value, false)
}
func (d *DownloadService) markCompletedTorrentCatchupRecorded(ctx context.Context, torrent QBitTorrent) {
if d == nil || d.repo == nil || d.repo.Setting == nil {
return
}
key := completedTorrentCatchupSettingKey(torrent)
if key == "" {
return
}
if err := d.repo.Setting.Set(ctx, key, "true"); err != nil && d.log != nil {
d.log.Debug("mark completed torrent catchup failed",
zap.String("hash", torrent.Hash),
zap.String("name", torrent.Name),
zap.Error(err))
}
}
func completedTorrentCatchupSettingKey(torrent QBitTorrent) string {
key := completedTorrentQueueKey(torrent)
if key == "" {
return ""
}
sum := sha1.Sum([]byte(key))
return completedTorrentCatchupSettingPrefix + fmt.Sprintf("%x", sum[:])
}
func (d *DownloadService) completedTorrentNotified(ctx context.Context, torrent QBitTorrent) bool {
if d == nil || d.repo == nil || d.repo.Setting == nil {
return false
}
key := completedTorrentNotifySettingKey(torrent)
if key == "" {
return false
}
value, err := d.repo.Setting.Get(ctx, key)
if err != nil {
return false
}
return parseBoolSetting(value, false)
}
func (d *DownloadService) markCompletedTorrentNotified(ctx context.Context, torrent QBitTorrent) {
if d == nil || d.repo == nil || d.repo.Setting == nil {
return
}
key := completedTorrentNotifySettingKey(torrent)
if key == "" {
return
}
if err := d.repo.Setting.Set(ctx, key, "true"); err != nil && d.log != nil {
d.log.Debug("mark completed torrent notification failed",
zap.String("hash", torrent.Hash),
zap.String("name", torrent.Name),
zap.Error(err))
}
}
func completedTorrentNotifySettingKey(torrent QBitTorrent) string {
key := completedTorrentQueueKey(torrent)
if key == "" {
return ""
}
sum := sha1.Sum([]byte(key))
return completedTorrentNotifySettingPrefix + fmt.Sprintf("%x", sum[:])
}
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])
}
key := strings.Join(parts, "|")
if strings.Trim(key, "|") == "" {
return ""
}
return strings.ToLower(key)
}
func (d *DownloadService) syncDownloadTaskProgress(ctx context.Context, torrent QBitTorrent, taskByKey map[string]model.DownloadTask) {
if d == nil || d.repo == nil || d.repo.DB == nil {
return
}
matched, ok := findMatchingTaskForTorrent(torrent, taskByKey)
if !ok {
return
}
status := torrent.State
if qbitTorrentCompleted(torrent) {
status = "completed"
}
if strings.TrimSpace(status) == "" {
status = matched.Status
}
updates := map[string]any{}
if math.Abs(float64(matched.Progress-torrent.Progress)) > 0.0001 {
updates["progress"] = torrent.Progress
}
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
}
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", matched.ID).Updates(updates).Error
}
func tasksByIdentity(rows []model.DownloadTask) map[string]model.DownloadTask {
out := make(map[string]model.DownloadTask, len(rows))
for _, row := range rows {
key := downloadTaskIdentityKey(row.Title)
if key != "" {
out[key] = row
}
}
return out
}
func tasksByTorrentIdentity(rows []model.DownloadTask) map[string]model.DownloadTask {
out := make(map[string]model.DownloadTask, len(rows)*4)
for _, row := range rows {
key := normalizeTorrentName(row.Title)
if key != "" {
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 == "" {
return model.DownloadTask{}, false
}
if row, ok := taskByKey[key]; ok {
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
}
}
return model.DownloadTask{}, false
}
func findMatchingTaskByTorrentIdentity(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.HasPrefix(currentKey, "\x00") {
continue
}
if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) {
return row, true
}
}
return model.DownloadTask{}, false
}
func downloadTaskNeedsCompletion(task model.DownloadTask) bool {
if task.Progress < 1 {
return true
}
return strings.ToLower(strings.TrimSpace(task.Status)) != "completed"
}
-390
View File
@@ -1,390 +0,0 @@
package service
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestReloadConfigDoesNotFallbackToLegacyAfterClientDeleted(t *testing.T) {
var addCalls int32
qb := 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/info":
if atomic.LoadInt32(&addCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"abc123","name":"Movie 2026 1080p","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "qbittorrent.url", qb.URL); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.username", "admin"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.password", "admin"); err != nil {
t.Fatal(err)
}
client := &model.DownloadClient{Name: "qB", Type: "qbittorrent", Host: qb.URL, Username: "admin", Password: "admin", IsDefault: true, Enabled: true}
if err := repos.DownloadClient.Create(t.Context(), client); err != nil {
t.Fatal(err)
}
if err := repos.DownloadClient.Delete(t.Context(), client.ID); 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)
}
_, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err == nil {
t.Fatal("expected add to fail when the configured downloader was deleted")
}
if got := atomic.LoadInt32(&addCalls); got != 0 {
t.Fatalf("qb add calls = %d, want 0", got)
}
}
func TestReloadConfigDoesNotFallbackToLegacyAfterClientDisabled(t *testing.T) {
var addCalls int32
qb := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "qbittorrent.url", qb.URL); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.username", "admin"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.password", "admin"); err != nil {
t.Fatal(err)
}
client := &model.DownloadClient{Name: "qB", Type: "qbittorrent", Host: qb.URL, Username: "admin", Password: "admin", IsDefault: true, Enabled: true}
if err := repos.DownloadClient.Create(t.Context(), client); err != nil {
t.Fatal(err)
}
client.Enabled = false
if err := repos.DownloadClient.Update(t.Context(), client); 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)
}
_, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err == nil {
t.Fatal("expected add to fail when the configured downloader was disabled")
}
if got := atomic.LoadInt32(&addCalls); got != 0 {
t.Fatalf("qb add calls = %d, want 0", got)
}
}
func TestReloadConfigUsesSoleEnabledQBitWhenNoExplicitDefault(t *testing.T) {
var addCalls int32
qb := 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/info":
if atomic.LoadInt32(&addCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"sole123","name":"Movie 2026 1080p","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.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)
}
client := &model.DownloadClient{Name: "qB", Type: "qbittorrent", Host: qb.URL, Username: "admin", Password: "admin", IsDefault: false, Enabled: true}
if err := repos.DownloadClient.Create(t.Context(), client); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:abababababababababababababababababababab&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err != nil {
t.Fatal(err)
}
if task == nil {
t.Fatal("expected task")
}
if got := atomic.LoadInt32(&addCalls); got != 1 {
t.Fatalf("qb add calls = %d, want 1", got)
}
}
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)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err == nil {
t.Fatal("expected no downloader configured error")
}
if !strings.Contains(err.Error(), "当前没有已启用的下载器") {
t.Fatalf("err = %v, want enabled downloader guidance", err)
}
if task != nil {
t.Fatalf("task = %#v, want nil", task)
}
rows, err := repos.Download.List(t.Context())
if err != nil {
t.Fatal(err)
}
if len(rows) != 0 {
t.Fatalf("download rows = %d, want 0", len(rows))
}
}
func TestAddDownloadSelectsFirstEnabledQBitWhenDefaultMissing(t *testing.T) {
var firstAddCalls int32
first := 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/info":
if atomic.LoadInt32(&firstAddCalls) > 0 {
_, _ = w.Write([]byte(`[{"hash":"abc123","name":"Movie 2026 1080p","state":"downloading","progress":0.1}]`))
return
}
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&firstAddCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer first.Close()
second := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
t.Fatal("second qB should not be selected before first enabled qB")
default:
http.NotFound(w, r)
}
}))
defer second.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)
}
firstClient := &model.DownloadClient{Name: "qB first", Type: "qbittorrent", Host: first.URL, Username: "admin", Password: "admin", IsDefault: false, Enabled: true}
secondClient := &model.DownloadClient{Name: "qB second", Type: "qbittorrent", Host: second.URL, Username: "admin", Password: "admin", IsDefault: false, Enabled: true}
if err := repos.DownloadClient.Create(t.Context(), firstClient); err != nil {
t.Fatal(err)
}
if err := repos.DownloadClient.Create(t.Context(), secondClient); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err != nil {
t.Fatal(err)
}
if task == nil {
t.Fatal("expected task")
}
if got := atomic.LoadInt32(&firstAddCalls); got != 1 {
t.Fatalf("first qb add calls = %d, want 1", got)
}
refreshed, err := repos.DownloadClient.FindByID(t.Context(), firstClient.ID)
if err != nil {
t.Fatal(err)
}
if refreshed == nil || !refreshed.IsDefault {
t.Fatalf("first enabled qB should be persisted as default, got %#v", refreshed)
}
}
func TestReloadConfigManagedModeDoesNotFallbackToLegacyWithoutRows(t *testing.T) {
var addCalls int32
qb := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/add":
atomic.AddInt32(&addCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "qbittorrent.url", qb.URL); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.username", "admin"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "qbittorrent.password", "admin"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
_, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
Title: "Movie 2026 1080p",
})
if err == nil {
t.Fatal("expected managed mode to reject missing default downloader")
}
if got := atomic.LoadInt32(&addCalls); got != 0 {
t.Fatalf("qb add calls = %d, want 0", got)
}
}
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)
}
client := &model.DownloadClient{
Name: "aria2",
Type: "aria2",
Host: aria2.URL,
Enabled: true,
}
if err := repos.DownloadClient.Create(t.Context(), client); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
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(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)
}
}
-75
View File
@@ -1,75 +0,0 @@
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)
}
}
-105
View File
@@ -1,105 +0,0 @@
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
}
-195
View File
@@ -1,195 +0,0 @@
package service
import (
"encoding/json"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestDeleteMarksMatchingDownloadTaskDeleted(t *testing.T) {
const hash = "abc123"
const title = "Delete Marker Show S01E01 1080p"
var deleteCalls int32
qb := 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/info":
_, _ = w.Write([]byte(`[{"hash":"abc123","name":"Delete Marker Show S01E01 1080p","state":"downloading","progress":0.5}]`))
case "/api/v2/torrents/delete":
atomic.AddInt32(&deleteCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
task := &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "https://pt.example/download?id=1",
Title: title,
SavePath: "/downloads/tv",
Status: "downloading",
Progress: 0.5,
}
if err := repos.Download.Create(t.Context(), task); 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 err := svc.Delete(t.Context(), hash, false); err != nil {
t.Fatal(err)
}
if got := atomic.LoadInt32(&deleteCalls); got != 1 {
t.Fatalf("delete calls = %d, want 1", got)
}
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, want deleted", updated.Status)
}
}
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
qb := 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/info":
_, _ = w.Write([]byte(`[]`))
case "/api/v2/torrents/delete":
atomic.AddInt32(&deleteCalls, 1)
_, _ = w.Write([]byte("Ok."))
default:
http.NotFound(w, r)
}
}))
defer qb.Close()
db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
task := &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "magnet:?xt=urn:btih:" + hash + "&dn=Codex.Path.Verify.S01E01.2026",
Title: "Codex Path Verify S01E01 2026",
SavePath: "/downloads/tv",
Status: "queued",
}
if err := repos.Download.Create(t.Context(), task); 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 err := svc.Delete(t.Context(), hash, false); err != nil {
t.Fatal(err)
}
if got := atomic.LoadInt32(&deleteCalls); got != 1 {
t.Fatalf("delete calls = %d, want 1", got)
}
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, want deleted", updated.Status)
}
}
+75
View File
@@ -0,0 +1,75 @@
package service
import (
"strings"
"unicode"
)
// 整理入库与外部搜索依赖的通用标题归一化辅助函数。
// 这些函数原本与订阅/下载逻辑共存,移除对应功能时保留为通用帮助函数。
func normalizeAvailabilityComparable(value string) string {
var b strings.Builder
for _, r := range strings.ToLower(value) {
if unicode.IsLetter(r) || unicode.IsDigit(r) {
b.WriteRune(r)
}
}
return b.String()
}
func firstNonEmpty(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
func compactUniqueStrings(values ...string) []string {
seen := map[string]struct{}{}
out := make([]string, 0, len(values))
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
key := normalizeAvailabilityComparable(value)
if key == "" {
continue
}
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
out = append(out, value)
}
return out
}
func titleMatchesResolution(titleFold, resolution string) bool {
switch strings.ToLower(strings.TrimSpace(resolution)) {
case "2160p", "4k", "uhd":
return strings.Contains(titleFold, "2160p") || strings.Contains(titleFold, "4k") || strings.Contains(titleFold, "uhd")
case "1080p":
return strings.Contains(titleFold, "1080p") || strings.Contains(titleFold, "fhd")
case "720p":
return strings.Contains(titleFold, "720p")
default:
return strings.Contains(titleFold, strings.ToLower(strings.TrimSpace(resolution)))
}
}
func detectResolutionScore(titleFold string) int {
switch {
case titleMatchesResolution(titleFold, "2160p"):
return 4
case titleMatchesResolution(titleFold, "1080p"):
return 3
case titleMatchesResolution(titleFold, "720p"):
return 2
default:
return 1
}
}
-207
View File
@@ -1,207 +0,0 @@
package service
import (
"fmt"
"net/url"
"path"
"regexp"
"strings"
"unicode"
)
var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`)
var downloadPackTitleToken = regexp.MustCompile(`(?i)(?:^|[\s._-])(?:complete|batch|pack|合集|全集|整季|全季)(?:[\s._-]|$)`)
func localAvailabilityTitleCandidates(title string) []string {
seen := map[string]struct{}{}
out := make([]string, 0, 6)
add := func(value string) {
value = strings.TrimSpace(value)
if value == "" {
return
}
if _, ok := seen[value]; ok {
return
}
seen[value] = struct{}{}
out = append(out, value)
}
add(availabilityQuery(title, ""))
if cleaned, _ := CleanQuery(title); cleaned != "" {
for _, candidate := range titleCandidates(cleaned) {
add(candidate)
fields := strings.Fields(candidate)
for i := len(fields) - 1; i >= 1; i-- {
prefix := strings.Join(fields[:i], " ")
if containsCJK(prefix) {
add(prefix)
}
}
}
}
return out
}
func downloadTaskBlocksDuplicate(status string) bool {
switch strings.ToLower(strings.TrimSpace(status)) {
case "failed", "error", "removed", "cancelled", "canceled":
return false
default:
return true
}
}
func downloadTaskBlocksReadd(status string) bool {
switch strings.ToLower(strings.TrimSpace(status)) {
case "failed", "error", "deleted", "removed", "cancelled", "canceled":
return false
default:
return true
}
}
func downloadTaskIdentityKey(name string) string {
if key := downloadMediaIdentityKey(name); key != "" {
return key
}
return normalizedDownloadTitleKey(name)
}
type downloadMediaIdentity struct {
TitleKey string
Year int
Episodes []episodeRef
Pack bool
}
func parseDownloadMediaIdentity(name string) downloadMediaIdentity {
title, year := CleanQuery(name)
if isSeriesPackTitle(name) {
title = downloadPackTitleToken.ReplaceAllString(title, " ")
}
titleKey := normalizeAvailabilityComparable(title)
if titleKey == "" {
titleKey = normalizeAvailabilityComparable(availabilityQuery(name, ""))
}
return downloadMediaIdentity{
TitleKey: titleKey,
Year: year,
Episodes: episodeRefsFromTitle(name),
Pack: isSeriesPackTitle(name),
}
}
func downloadTitleCoversRequest(existing, requested string) bool {
current := parseDownloadMediaIdentity(existing)
want := parseDownloadMediaIdentity(requested)
if current.TitleKey == "" || want.TitleKey == "" {
currentKey := normalizedDownloadTitleKey(existing)
wantKey := normalizedDownloadTitleKey(requested)
return currentKey != "" && wantKey != "" && (currentKey == wantKey || strings.Contains(currentKey, wantKey) || strings.Contains(wantKey, currentKey))
}
if current.TitleKey != want.TitleKey {
return false
}
if current.Year > 0 && want.Year > 0 && current.Year != want.Year {
return false
}
if downloadIdentityCoversWholeSeason(existing, current) {
return true
}
if len(current.Episodes) == 0 || len(want.Episodes) == 0 {
return len(current.Episodes) == len(want.Episodes)
}
currentEpisodes := map[string]struct{}{}
for _, ref := range current.Episodes {
currentEpisodes[episodeKey(ref.Season, ref.Episode)] = struct{}{}
}
for _, ref := range want.Episodes {
if _, ok := currentEpisodes[episodeKey(ref.Season, ref.Episode)]; !ok {
return false
}
}
return true
}
func downloadIdentityCoversWholeSeason(title string, identity downloadMediaIdentity) bool {
if !identity.Pack || len(identity.Episodes) > 0 {
return false
}
return seriesPackRE.MatchString(title)
}
func downloadMediaIdentityKey(name string) string {
name = strings.ToLower(strings.TrimSpace(name))
if name == "" {
return ""
}
identity := parseDownloadMediaIdentity(name)
titleKey := identity.TitleKey
if titleKey == "" {
return ""
}
parts := []string{titleKey}
if identity.Year > 0 {
parts = append(parts, fmt.Sprintf("y%d", identity.Year))
}
if len(identity.Episodes) > 0 {
first := identity.Episodes[0]
last := identity.Episodes[len(identity.Episodes)-1]
parts = append(parts, fmt.Sprintf("s%02de%03d", first.Season, first.Episode))
if len(identity.Episodes) > 1 {
parts = append(parts, fmt.Sprintf("to%03d", last.Episode))
}
}
return strings.Join(parts, "|")
}
func normalizedDownloadTitleKey(name string) string {
name = strings.ToLower(strings.TrimSpace(name))
var b strings.Builder
for _, r := range name {
if unicode.IsLetter(r) || unicode.IsDigit(r) {
b.WriteRune(r)
}
}
return b.String()
}
func publicDownloadTitle(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return "下载任务"
}
if u, err := url.Parse(raw); err == nil {
if dn := strings.TrimSpace(u.Query().Get("dn")); dn != "" {
if decoded, err := url.QueryUnescape(dn); err == nil && strings.TrimSpace(decoded) != "" {
return strings.TrimSpace(decoded)
}
return dn
}
if u.Host != "" {
base := path.Base(u.Path)
if base != "." && base != "/" && base != "" {
base = strings.TrimSuffix(base, path.Ext(base))
if base != "" {
return base
}
}
return u.Host
}
}
if strings.HasPrefix(strings.ToLower(raw), "magnet:") {
return "磁力下载"
}
return "下载任务"
}
func normalizeTorrentName(name string) string {
name = torrentEpisodeToken.ReplaceAllString(strings.ToLower(name), "")
var b strings.Builder
for _, r := range name {
if unicode.IsLetter(r) || unicode.IsDigit(r) {
b.WriteRune(r)
}
}
return b.String()
}
@@ -1,52 +0,0 @@
package service
import "time"
func (d *DownloadService) currentTime() time.Time {
if d != nil && d.now != nil {
return d.now()
}
return time.Now()
}
func (d *DownloadService) recordLiveTorrentSnapshot(live []QBitTorrent) {
if d == nil {
return
}
snapshot := cloneQBitTorrentSlice(live)
d.mu.Lock()
d.liveTorrents = snapshot
d.liveTorrentsAt = d.currentTime()
d.mu.Unlock()
}
func (d *DownloadService) LiveTorrentSnapshot(maxAge time.Duration) []QBitTorrent {
snapshot, ok := d.liveTorrentSnapshot(maxAge)
if !ok {
return nil
}
return snapshot
}
func (d *DownloadService) liveTorrentSnapshot(maxAge time.Duration) ([]QBitTorrent, bool) {
if d == nil {
return nil, false
}
now := d.currentTime()
d.mu.Lock()
defer d.mu.Unlock()
if d.liveTorrentsAt.IsZero() {
return nil, false
}
if maxAge > 0 && now.Sub(d.liveTorrentsAt) > maxAge {
return nil, false
}
return cloneQBitTorrentSlice(d.liveTorrents), true
}
func cloneQBitTorrentSlice(in []QBitTorrent) []QBitTorrent {
if len(in) == 0 {
return nil
}
return append([]QBitTorrent(nil), in...)
}
-346
View File
@@ -1,346 +0,0 @@
// Package service — 下载管理器,管理多个下载客户端适配器。
//
// DownloadManager 提供多客户端分发能力,支持运行时热插拔。
// 调用方通过 GetDefault() 或 GetClient(id) 获取适配器来执行下载操作。
package service
import (
"context"
"encoding/json"
"errors"
"fmt"
"sync"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// DownloadManager 管理多个下载客户端适配器实例。
type DownloadManager struct {
log *zap.Logger
repo *repository.Container
crypto *CryptoService
mu sync.RWMutex
clients map[string]DownloadAdapter // clientID -> adapter
configs map[string]DownloadClientConfig
models map[string]model.DownloadClient
order []string
}
// NewDownloadManager 创建新的下载管理器。
func NewDownloadManager(log *zap.Logger, repo *repository.Container, crypto *CryptoService) *DownloadManager {
return &DownloadManager{
log: log,
repo: repo,
crypto: crypto,
clients: make(map[string]DownloadAdapter),
configs: make(map[string]DownloadClientConfig),
models: make(map[string]model.DownloadClient),
}
}
// LoadAll 从数据库加载所有已启用的客户端并初始化适配器。
func (m *DownloadManager) LoadAll(ctx context.Context) error {
dbClients, err := m.repo.DownloadClient.ListEnabled(ctx)
if err != nil {
return err
}
if err := m.ensureEnabledDefault(ctx, dbClients); err != nil {
return err
}
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 {
adapter, cfg, ok := m.initializeClient(ctx, dc)
if !ok {
continue
}
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(_ context.Context) (*model.DownloadClient, DownloadAdapter, error) {
m.mu.RLock()
defer m.mu.RUnlock()
for _, id := range m.order {
client := m.models[id]
if client.IsDefault {
if adapter, ok := m.clients[id]; ok {
copy := client
return &copy, adapter, nil
}
}
}
for _, id := range m.order {
if adapter, ok := m.clients[id]; ok {
client := m.models[id]
copy := client
return &copy, adapter, nil
}
}
return nil, nil, errors.New("no download client available")
}
// GetClient 返回指定 ID 的下载客户端适配器。
func (m *DownloadManager) GetClient(id string) (DownloadAdapter, error) {
m.mu.RLock()
defer m.mu.RUnlock()
adapter, ok := m.clients[id]
if !ok {
return nil, errors.New("download client not found or not initialized")
}
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)
if err != nil {
return err
}
adapter := AdapterFactory(dc.Type)
if adapter == nil {
return errors.New("unknown download client type: " + dc.Type)
}
if err := adapter.Initialize(ctx, cfg); err != nil {
return err
}
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
}
// RemoveClient 移除一个下载客户端(停止适配器,不删除数据库记录)。
func (m *DownloadManager) RemoveClient(id string) {
m.mu.Lock()
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 更新已有客户端的配置并重新初始化。
func (m *DownloadManager) UpdateClient(ctx context.Context, dc *model.DownloadClient) error {
m.RemoveClient(dc.ID)
return m.AddClient(ctx, dc)
}
// TestConnection 测试客户端连接。
func (m *DownloadManager) TestConnection(ctx context.Context, dc *model.DownloadClient) error {
cfg, err := m.buildConfig(dc)
if err != nil {
return err
}
adapter := AdapterFactory(dc.Type)
if adapter == nil {
return errors.New("unknown download client type: " + dc.Type)
}
return adapter.Initialize(ctx, cfg)
}
// ListAll 获取所有已加载客户端的种子列表。
func (m *DownloadManager) ListAll(ctx context.Context, filter string) (map[string][]TorrentInfo, error) {
result := make(map[string][]TorrentInfo)
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", target.client.ID),
zap.Error(err),
)
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
}
}
return result, errors.Join(listErrs...)
}
// GetAdapterTypes 返回支持的下载客户端类型列表。
func (m *DownloadManager) GetAdapterTypes() []AdapterTypeInfo {
return []AdapterTypeInfo{
{Type: "qbittorrent", Name: "qBittorrent", Description: "qBittorrent WebUI API (v2)"},
{Type: "transmission", Name: "Transmission", Description: "Transmission RPC API"},
{Type: "aria2", Name: "Aria2", Description: "Aria2 JSON-RPC API"},
}
}
// AdapterTypeInfo 描述下载客户端类型信息。
type AdapterTypeInfo struct {
Type string `json:"type"`
Name string `json:"name"`
Description string `json:"description"`
}
// buildConfig 从数据库模型构建适配器配置。
func (m *DownloadManager) buildConfig(dc *model.DownloadClient) (DownloadClientConfig, error) {
password := dc.Password
if m.crypto != nil && password != "" {
password = m.crypto.Decrypt(password)
}
host, err := normalizeDownloadClientEndpoint(dc.Type, dc.Host)
if err != nil {
return DownloadClientConfig{}, err
}
cfg := DownloadClientConfig{
Host: host,
Username: dc.Username,
Password: password,
}
// 解析 Extra JSON 配置
if dc.Extra != "" {
extraStr := dc.Extra
if m.crypto != nil {
extraStr = m.crypto.Decrypt(extraStr)
}
var extra map[string]string
if err := json.Unmarshal([]byte(extraStr), &extra); err == nil {
cfg.Extra = extra
}
}
return cfg, nil
}
@@ -1,417 +0,0 @@
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])
}
}
-84
View File
@@ -1,84 +0,0 @@
package service
import (
"context"
"path/filepath"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func (d *DownloadService) notifyDownloadComplete(ctx context.Context, torrent QBitTorrent, task *model.DownloadTask) {
if d == nil || d.notify == nil {
return
}
if d.completedTorrentNotified(ctx, torrent) {
return
}
d.markCompletedTorrentNotified(ctx, torrent)
body, data := downloadCompleteNotificationPayload(torrent, task)
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
d.notify.BroadcastEvent(ctx, NotifyEvent{
Type: EventDownloadComplete,
Title: "MediaStationGo 下载完成",
Message: body,
Data: data,
})
}()
}
func downloadCompleteNotificationPayload(torrent QBitTorrent, task *model.DownloadTask) (string, map[string]interface{}) {
name := downloadCompleteNotificationName(torrent, task)
body := "任务:" + name
data := downloadCompleteNotificationData(torrent, task)
return body, data
}
func downloadCompleteNotificationName(torrent QBitTorrent, task *model.DownloadTask) string {
name := strings.TrimSpace(torrent.Name)
if name == "" {
name = strings.TrimSpace(filepath.Base(torrent.ContentPath))
}
if task != nil && strings.TrimSpace(task.Title) != "" {
name = strings.TrimSpace(task.Title)
}
if name == "" {
name = "下载任务"
}
return name
}
func downloadCompleteNotificationData(torrent QBitTorrent, task *model.DownloadTask) map[string]interface{} {
data := map[string]interface{}{}
if rt := strings.TrimSpace(torrent.Name); rt != "" {
data["resource_title"] = rt
}
if task == nil {
return data
}
addTrimmedString(data, "poster_url", task.PosterURL)
addTrimmedString(data, "backdrop_url", task.BackdropURL)
addTrimmedString(data, "media_type", task.MediaType)
addTrimmedString(data, "media_category", task.MediaCategory)
addTrimmedString(data, "title", task.Title)
addTrimmedString(data, "overview", task.Overview)
addTrimmedString(data, "original_title", task.OriginalName)
addTrimmedString(data, "original_language", task.OriginalLanguage)
if task.Year > 0 {
data["year"] = task.Year
}
if task.Rating > 0 {
data["rating"] = task.Rating
}
addTrimmedString(data, "genres", task.Genres)
return data
}
func addTrimmedString(data map[string]interface{}, key, value string) {
if strings.TrimSpace(value) != "" {
data[key] = value
}
}
-201
View File
@@ -1,201 +0,0 @@
package service
import (
"context"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
const completedTorrentOrganizeQueueSize = 64
var completedTorrentOrganizeCooldown = 3 * time.Second
// 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)
defer t.Stop()
// prevStates tracks previous completion states to detect "just finished"
if d.prevStates == nil {
d.prevStates = make(map[string]bool)
}
for {
select {
case <-ctx.Done():
return
case <-d.stopCh:
return
case <-t.C:
}
live, err := d.listLiveTorrents(ctx, "")
if err != nil && len(live) == 0 {
continue
}
rows, _ := d.repo.Download.List(ctx)
taskByKey := tasksByTorrentIdentity(rows)
d.processDownloadSnapshot(ctx, live, taskByKey)
d.hub.Publish("download", map[string]any{"torrents": live})
}
}
func (d *DownloadService) processDownloadSnapshot(ctx context.Context, live []QBitTorrent, taskByKey map[string]model.DownloadTask) {
d.recordLiveTorrentSnapshot(live)
firstSnapshot := d.beginDownloadSnapshot()
for _, torrent := range live {
d.processTorrentSnapshot(ctx, torrent, taskByKey, firstSnapshot)
}
}
func (d *DownloadService) beginDownloadSnapshot() bool {
d.mu.Lock()
defer d.mu.Unlock()
if d.prevStates == nil {
d.prevStates = make(map[string]bool)
}
firstSnapshot := !d.pollInitialized
if firstSnapshot {
d.pollInitialized = true
}
return firstSnapshot
}
func (d *DownloadService) processTorrentSnapshot(ctx context.Context, torrent QBitTorrent, taskByKey map[string]model.DownloadTask, firstSnapshot bool) {
stateKey := completedTorrentQueueKey(torrent)
taskNeedsOrganize := d.downloadSnapshotTaskNeedsOrganize(ctx, torrent, taskByKey)
d.syncDownloadTaskProgress(ctx, torrent, taskByKey)
if stateKey == "" {
return
}
if d.completedTorrentShouldQueue(stateKey, qbitTorrentCompleted(torrent), firstSnapshot, taskNeedsOrganize) &&
d.enqueueCompletedTorrent(torrent) {
d.markCompletedTorrentState(stateKey)
}
}
func (d *DownloadService) downloadSnapshotTaskNeedsOrganize(ctx context.Context, torrent QBitTorrent, taskByKey map[string]model.DownloadTask) bool {
matchedTask, hasTask := findMatchingTaskForTorrent(torrent, taskByKey)
if !hasTask || d.completedTorrentCatchupRecorded(ctx, torrent) || !d.downloadAutoOrganizeEnabled(ctx) {
return false
}
return downloadTaskNeedsCompletion(matchedTask) || recentlyCompletedTorrent(torrent, time.Now())
}
func (d *DownloadService) completedTorrentShouldQueue(stateKey string, complete, firstSnapshot, taskNeedsOrganize bool) bool {
d.mu.Lock()
defer d.mu.Unlock()
if d.prevStates == nil {
d.prevStates = make(map[string]bool)
}
wasComplete, wasSeen := d.prevStates[stateKey]
switch {
case complete && (firstSnapshot || !wasSeen):
// 首次快照里已完成的种子:此前一律标记「已见过」并跳过整理,
// 导致「下载完成时应用恰好不在线/正在重启」的种子永远不会被
// 自动整理入库。现在对最近完成的种子补一次整理
// (onTorrentComplete 内部仍受 organize.auto 开关约束,且
// 整理对已存在的目标文件幂等跳过)。
d.prevStates[stateKey] = true
return taskNeedsOrganize
case complete && !wasComplete:
return true
case complete && taskNeedsOrganize:
return true
case complete:
d.prevStates[stateKey] = true
default:
d.prevStates[stateKey] = false
}
return false
}
func (d *DownloadService) markCompletedTorrentState(stateKey string) {
d.mu.Lock()
d.prevStates[stateKey] = true
d.mu.Unlock()
}
func (d *DownloadService) startAutoOrganizeWorker(ctx context.Context) {
d.mu.Lock()
if d.organizeQueue == nil {
d.organizeQueue = make(chan QBitTorrent, completedTorrentOrganizeQueueSize)
}
if d.organizeQueued == nil {
d.organizeQueued = make(map[string]struct{})
}
d.mu.Unlock()
d.organizeOnce.Do(func() {
go d.autoOrganizeWorker(ctx)
})
}
func (d *DownloadService) enqueueCompletedTorrent(torrent QBitTorrent) bool {
key := completedTorrentQueueKey(torrent)
if key == "" {
return false
}
d.mu.Lock()
if d.organizeQueue == nil {
d.organizeQueue = make(chan QBitTorrent, completedTorrentOrganizeQueueSize)
}
if d.organizeQueued == nil {
d.organizeQueued = make(map[string]struct{})
}
if _, ok := d.organizeQueued[key]; ok {
d.mu.Unlock()
return true
}
select {
case d.organizeQueue <- torrent:
d.organizeQueued[key] = struct{}{}
d.mu.Unlock()
return true
default:
d.mu.Unlock()
if d.log != nil {
d.log.Warn("auto organize queue full; will retry completed torrent later",
zap.String("hash", torrent.Hash),
zap.String("name", torrent.Name))
}
return false
}
}
func (d *DownloadService) autoOrganizeWorker(ctx context.Context) {
for {
select {
case <-ctx.Done():
return
case <-d.stopCh:
return
case torrent := <-d.organizeQueue:
d.onTorrentComplete(ctx, torrent)
d.markCompletedTorrentOrganizeDone(torrent)
if completedTorrentOrganizeCooldown <= 0 {
continue
}
timer := time.NewTimer(completedTorrentOrganizeCooldown)
select {
case <-ctx.Done():
timer.Stop()
return
case <-d.stopCh:
timer.Stop()
return
case <-timer.C:
}
}
}
}
func (d *DownloadService) markCompletedTorrentOrganizeDone(torrent QBitTorrent) {
key := completedTorrentQueueKey(torrent)
if key == "" {
return
}
d.mu.Lock()
delete(d.organizeQueued, key)
d.mu.Unlock()
}
-191
View File
@@ -1,191 +0,0 @@
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
}
@@ -1,80 +0,0 @@
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
}
-217
View File
@@ -1,217 +0,0 @@
package service
import (
"math"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
type DownloadTaskView struct {
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"`
BackdropURL string `json:"backdrop_url,omitempty"`
Overview string `json:"overview,omitempty"`
MediaType string `json:"media_type,omitempty"`
MediaCategory string `json:"media_category,omitempty"`
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"`
Downloaded int64 `json:"downloaded"`
SavePath string `json:"save_path"`
}
func DownloadViews(rows []model.DownloadTask, live []QBitTorrent) ([]DownloadTaskView, []DownloadTorrentView) {
liveByKey := map[string]QBitTorrent{}
for _, torrent := range live {
key := normalizeTorrentName(torrent.Name)
if key != "" {
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 := findMatchingTorrentForTask(row, liveByKey); ok {
view = downloadTaskView(row, torrent)
}
taskViews = append(taskViews, view)
}
torrentViews := make([]DownloadTorrentView, 0, len(live))
for _, torrent := range live {
var row model.DownloadTask
if matched, ok := findMatchingTaskForTorrent(torrent, taskByKey); ok {
row = matched
}
torrentViews = append(torrentViews, downloadTorrentView(torrent, row))
}
return taskViews, torrentViews
}
func downloadTaskView(row model.DownloadTask, torrent QBitTorrent) DownloadTaskView {
progress := row.Progress
state := row.Status
if torrent.Name != "" {
progress = torrent.Progress
state = torrent.State
}
size := torrent.Size
return DownloadTaskView{
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,
}
}
func downloadTorrentView(torrent QBitTorrent, row model.DownloadTask) DownloadTorrentView {
title := torrent.Name
if row.Title != "" {
title = row.Title
}
return DownloadTorrentView{
Hash: torrent.Hash,
ClientID: torrent.ClientID,
Source: torrent.Source,
Name: torrent.Name,
Title: firstNonEmpty(title, "下载任务"),
PosterURL: row.PosterURL,
BackdropURL: row.BackdropURL,
Overview: row.Overview,
MediaType: row.MediaType,
MediaCategory: firstNonEmpty(row.MediaCategory, torrent.Category),
State: torrent.State,
Progress: torrent.Progress,
DLSpeed: torrent.DLSpeed,
UpSpeed: torrent.UpSpeed,
NumSeeds: torrent.NumSeeds,
NumLeechs: torrent.NumLeech,
Size: torrent.Size,
Downloaded: downloadedBytes(torrent.Size, torrent.Progress),
SavePath: torrent.SavePath,
}
}
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 == "" {
return QBitTorrent{}, false
}
if torrent, ok := liveByKey[key]; ok {
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
}
}
return QBitTorrent{}, false
}
func downloadedBytes(size int64, progress float32) int64 {
if size <= 0 || progress <= 0 {
return 0
}
if progress > 1 {
progress = 1
}
return int64(math.Round(float64(size) * float64(progress)))
}
func firstNonEmpty(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
-308
View File
@@ -1,308 +0,0 @@
// Package service — download manager.
//
// DownloadService persists user-initiated downloads, dispatches them to
// 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):
//
// qbittorrent.url e.g. http://127.0.0.1:8080
// qbittorrent.username qBittorrent WebUI user
// qbittorrent.password qBittorrent WebUI password
// qbittorrent.savepath optional default save dir
//
// Settings can be updated at runtime via the admin UI; ReloadConfig()
// re-reads them and re-authenticates.
package service
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// DownloadService is the single download orchestrator.
type DownloadService struct {
log *zap.Logger
repo *repository.Container
hub *Hub
qb *QBitClient
manager *DownloadManager
organizer *OrganizerService
organizePipeline *OrganizePipelineService
scanner *ScannerService
site *SiteService
tasks *TaskTrackerService
notify *NotifyChannelService
mu sync.Mutex
stopCh chan struct{}
pollOnce sync.Once
organizeOnce sync.Once
prevStates map[string]bool // client/task identity -> wasCompleted
pollInitialized bool
liveTorrents []QBitTorrent
liveTorrentsAt time.Time
now func() time.Time
organizeQueue chan QBitTorrent
organizeQueued map[string]struct{}
}
func (d *DownloadService) SetScanner(scanner *ScannerService) {
d.scanner = scanner
}
func (d *DownloadService) SetOrganizePipeline(pipeline *OrganizePipelineService) {
d.organizePipeline = pipeline
}
func (d *DownloadService) SetTaskTracker(tasks *TaskTrackerService) {
d.tasks = tasks
}
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 a downloader. Subscriptions treat this as a
// successful dedup hit, not as a retryable enqueue failure.
var ErrDownloadAlreadyExists = errors.New("download already exists")
// ErrMediaAlreadyInLibrary tells callers that the requested movie/episode is
// already present in the scanned media library and must not be sent to the
// 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)
}
// NewDownloadService is the constructor.
func NewDownloadService(log *zap.Logger, repo *repository.Container, hub *Hub, organizer *OrganizerService, site ...*SiteService) *DownloadService {
var siteSvc *SiteService
if len(site) > 0 {
siteSvc = site[0]
}
return &DownloadService{
log: log,
repo: repo,
hub: hub,
qb: NewQBitClient(log, QBitConfig{}),
organizer: organizer,
site: siteSvc,
prevStates: make(map[string]bool),
now: time.Now,
organizeQueue: make(chan QBitTorrent, completedTorrentOrganizeQueueSize),
organizeQueued: make(map[string]struct{}),
stopCh: make(chan struct{}),
}
}
// Start kicks off the background poller (idempotent).
func (d *DownloadService) Start(ctx context.Context) {
d.pollOnce.Do(func() {
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)
})
}
// Stop terminates the poller.
func (d *DownloadService) Stop() {
close(d.stopCh)
}
func (d *DownloadService) TorrentExistsByName(ctx context.Context, name string) bool {
query := normalizeTorrentName(name)
if query == "" {
return false
}
live, err := d.listLiveTorrents(ctx, "")
if err != nil && len(live) == 0 {
return false
}
for _, torrent := range live {
if downloadTitleCoversRequest(torrent.Name, name) {
return true
}
current := normalizeTorrentName(torrent.Name)
if current == "" {
continue
}
if current == query {
return true
}
}
return false
}
// 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.listLiveTorrents(ctx, "")
if err != 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 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")
}
requestedClientID := ""
if len(clientID) > 0 {
requestedClientID = clientID[0]
}
resolvedClientID, torrentName, err := d.resolveOperationClientID(ctx, hash, requestedClientID)
if err != nil {
return err
}
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)
d.mu.Unlock()
return nil
}
func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, torrentName string, clientID ...string) {
if d == nil || d.repo == nil || d.repo.DB == nil {
return
}
rows, err := d.repo.Download.List(ctx)
if err != nil {
return
}
if matched, ok := findDownloadTaskByHash(rows, hash, clientID...); ok {
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).
Where("id = ?", matched.ID).
Updates(map[string]any{
"status": "deleted",
"progress": matched.Progress,
}).Error
return
}
if strings.TrimSpace(torrentName) == "" {
return
}
taskByKey := tasksByTorrentIdentity(rows)
matched, ok := findMatchingTaskForTorrent(QBitTorrent{Name: torrentName, Hash: hash, ClientID: firstString(clientID)}, taskByKey)
if !ok {
return
}
_ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).
Where("id = ?", matched.ID).
Updates(map[string]any{
"status": "deleted",
"progress": matched.Progress,
}).Error
}
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
}
}
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, clientID ...string) error {
if strings.TrimSpace(hash) == "" {
return errors.New("hash is required")
}
if strings.TrimSpace(location) == "" {
return errors.New("location is required")
}
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))
}
@@ -1,157 +0,0 @@
package service
import (
"os"
"path/filepath"
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestCompletedTorrentSourceDoesNotFallbackToSavePath(t *testing.T) {
root := t.TempDir()
savePath := filepath.Join(root, "downloads", "日番")
if err := os.MkdirAll(savePath, 0o755); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), newOrganizerTestRepo(t), NewHub(zap.NewNop()), nil)
got := svc.completedTorrentSource(t.Context(), QBitTorrent{
Hash: "done123",
Name: "Missing.Payload.S01",
SavePath: savePath,
ContentPath: filepath.Join(savePath, "Missing.Payload.S01", "Missing.Payload.S01E01.mkv"),
})
if got != "" {
t.Fatalf("completedTorrentSource fell back to whole save_path %q; want empty", got)
}
}
func TestDownloadCompleteRecordsUnsupportedVideoAsHandled(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "downloads", "Toy.Story.4.2019.iso")
dest := filepath.Join(root, "media")
writeOrgFile(t, src, "iso")
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
for key, value := range map[string]string{
"organizer.auto_after_download": "true",
"organize.target_dir": dest,
"organize.transfer_mode": "copy",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
torrent := QBitTorrent{
Hash: "unsupported-iso",
Name: "Toy.Story.4.2019",
Progress: 1,
SavePath: filepath.Dir(src),
ContentPath: src,
CompletionOn: time.Now().Add(-time.Hour).Unix(),
}
org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), org)
svc.onTorrentComplete(t.Context(), torrent)
if !svc.completedTorrentCatchupRecorded(t.Context(), torrent) {
t.Fatalf("unsupported completed torrent should be marked handled to avoid repeated auto-organize retries")
}
}
func TestAutoOrganizeSyncsVisibilityWhenTargetAlreadyExists(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "downloads", "国产剧", "狂飙.S01E01.2023.1080p.mkv")
dest := filepath.Join(root, "media")
writeOrgFile(t, src, "episode")
repos := newOrganizerTestRepo(t)
for key, value := range map[string]string{
"organizer.auto_after_download": "true",
"organize.target_dir": dest,
"organize.transfer_mode": "copy",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
lib := model.Library{Name: "国产剧", Path: filepath.Join(dest, "电视剧", "国产剧"), Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
if _, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
SourcePath: src,
DestPath: dest,
TransferMode: TransferCopy,
}); err != nil {
t.Fatalf("seed organized destination: %v", err)
}
scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), org)
svc.SetScanner(scanner)
svc.onTorrentComplete(t.Context(), QBitTorrent{
Hash: "done123",
Name: "狂飙.S01E01.2023.1080p",
Progress: 1,
SavePath: filepath.Dir(src),
ContentPath: src,
})
var count int64
if err := repos.DB.Model(&model.Media{}).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("target already exists should still be scanned into DB, count=%d want 1", count)
}
}
func TestCompletedTorrentSourceUsesConfiguredMapping(t *testing.T) {
root := t.TempDir()
localRoot := filepath.Join(root, "localdl")
payload := filepath.Join(localRoot, "Show.S01")
if err := os.MkdirAll(payload, 0o755); err != nil {
t.Fatal(err)
}
repos := newOrganizerTestRepo(t)
if err := repos.Setting.Set(t.Context(), DownloadPathMappingsSettingKey, "/qb/downloads="+localRoot); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
got := svc.completedTorrentSource(t.Context(), QBitTorrent{ContentPath: "/qb/downloads/Show.S01"})
if got != payload {
t.Fatalf("completedTorrentSource = %q, want %q", got, payload)
}
}
func TestUserPathMappingsParsing(t *testing.T) {
repos := newOrganizerTestRepo(t)
raw := "# comment\n/a=/b\n/c => /d\n/e:/f\nbad-line\n"
if err := repos.Setting.Set(t.Context(), DownloadPathMappingsSettingKey, raw); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
got := svc.userPathMappings(t.Context())
want := map[string]string{"/a": "/b", "/c": "/d", "/e": "/f"}
if len(got) != len(want) {
t.Fatalf("userPathMappings = %v, want %v", got, want)
}
for k, v := range want {
if got[k] != v {
t.Fatalf("mapping %q = %q, want %q", k, got[k], v)
}
}
}
@@ -1,134 +0,0 @@
package service
import (
"context"
"errors"
"fmt"
"strings"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
const settingDownloadClientsManaged = "download_clients.managed"
// ReloadConfig reloads managed adapters and preserves the legacy qBittorrent
// settings fallback for deployments that never used download_clients.
//
// 配置来源优先级:
//
// 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 表。此分支主要服务未注入 DownloadManager 的
// 单元/兼容调用;生产容器已在上方通过原生适配器返回。
if d.repo.DownloadClient != nil {
hasConfiguredClients, _ = d.repo.DownloadClient.HasAnyIncludingDeleted(ctx)
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.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
}
}
}
if d.repo.Setting != nil {
managedRaw, _ := d.repo.Setting.Get(ctx, settingDownloadClientsManaged)
managedByDownloadClients = strings.EqualFold(strings.TrimSpace(managedRaw), "true")
}
// Path 2: legacy Setting 表。
// 仅在旧部署“从未使用过 download_clients 表”时回退。只要操作员曾经
// 配置过下载器,删除/禁用全部下载器就表示应停止投递,不能再偷偷用
// qbittorrent.* 旧设置继续往下载器添加任务。
if cfg.BaseURL == "" && !hasConfiguredClients && !managedByDownloadClients {
get := func(k string) string {
v, _ := d.repo.Setting.Get(ctx, k)
return v
}
cfg.BaseURL = get("qbittorrent.url")
cfg.Username = get("qbittorrent.username")
cfg.Password = get("qbittorrent.password")
}
d.qb.Configure(cfg)
return nil
}
func (d *DownloadService) preferredEnabledClient(ctx context.Context) (*model.DownloadClient, error) {
if d == nil || d.repo == nil || d.repo.DownloadClient == nil {
return nil, nil
}
rows, err := d.repo.DownloadClient.ListEnabled(ctx)
if err != nil {
return nil, err
}
if len(rows) == 0 {
return nil, 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 + ": 请在下载客户端中配置并启用下载器")
}
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 + ": 请在下载客户端中启用下载器;当前没有已启用的下载器")
}
var enabled []string
for _, row := range rows {
label := strings.TrimSpace(row.Name)
if label == "" {
label = row.Type
} else if row.Type != "" {
label += "(" + row.Type + ")"
}
enabled = append(enabled, label)
}
return fmt.Errorf("%s: 请检查已启用下载器的连接和默认设置;当前启用的下载器为 %s", prefix, strings.Join(enabled, ", "))
}
-195
View File
@@ -1,195 +0,0 @@
package service
import (
"os"
"path/filepath"
"testing"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestDownloadCompleteAutoOrganizesContentPath(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "downloads", "国产剧", "狂飙.S01E01.2023.1080p.mkv")
dest := filepath.Join(root, "media")
if err := os.MkdirAll(filepath.Dir(src), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(src, []byte("episode"), 0o644); err != nil {
t.Fatal(err)
}
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
for key, value := range map[string]string{
"organizer.auto_after_download": "true",
"organize.target_dir": dest,
"organize.transfer_mode": "copy",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), org)
svc.onTorrentComplete(t.Context(), QBitTorrent{
Hash: "done123",
Name: "狂飙.S01E01.2023.1080p",
Progress: 1,
SavePath: filepath.Join(root, "downloads", "国产剧"),
ContentPath: src,
})
want := filepath.Join(dest, "电视剧", "国产剧", "狂飙", "Season 01", "狂飙 - S01E01.mkv")
if _, err := os.Stat(want); err != nil {
t.Fatalf("auto organized file missing at %q: %v", want, err)
}
if _, err := os.Stat(src); err != nil {
t.Fatalf("copy mode should keep source: %v", err)
}
}
func TestDownloadCompleteAutoOrganizeUsesTaskMediaCategory(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "downloads", "Motherhood.of.Taihang.S01E01.2026.1080p.mkv")
dest := filepath.Join(root, "media")
writeOrgFile(t, src, "episode")
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
for key, value := range map[string]string{
"organizer.auto_after_download": "true",
"organize.target_dir": dest,
"organize.transfer_mode": "copy",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
task := &model.DownloadTask{
UserID: "u1",
Source: "qbittorrent",
URL: "magnet:?xt=urn:btih:motherhood",
Title: "Motherhood.of.Taihang.S01E01.2026.1080p",
SavePath: filepath.Join(root, "downloads"),
MediaType: "tv",
MediaCategory: "国产剧",
Status: "completed",
Progress: 1,
}
if err := repos.Download.Create(t.Context(), task); err != nil {
t.Fatal(err)
}
org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), org)
svc.onTorrentComplete(t.Context(), QBitTorrent{
Hash: "done-category",
Name: "Motherhood.of.Taihang.S01E01.2026.1080p",
Progress: 1,
SavePath: filepath.Join(root, "downloads"),
ContentPath: src,
})
categoryRoot := filepath.Join(dest, "电视剧", "国产剧")
var organized string
err := filepath.WalkDir(categoryRoot, func(path string, entry os.DirEntry, err error) error {
if err != nil || entry.IsDir() {
return err
}
if filepath.Ext(path) == ".mkv" {
organized = path
}
return nil
})
if err != nil {
t.Fatalf("walk organized category: %v", err)
}
if organized == "" {
t.Fatalf("expected organized file under %q", categoryRoot)
}
wrongRoot := filepath.Join(dest, "电视剧", "Motherhood Of Taihang")
if _, err := os.Stat(wrongRoot); !os.IsNotExist(err) {
t.Fatalf("unexpected uncategorized organize root %q, err=%v", wrongRoot, err)
}
}
func TestDownloadCompleteOnlyReplacesExistingWhenTaskAllowsWash(t *testing.T) {
for _, tc := range []struct {
name string
allowWash bool
wantContent string
}{
{name: "wash disabled keeps existing", allowWash: false, wantContent: "inception-1080p"},
{name: "wash enabled replaces existing", allowWash: true, wantContent: "inception-2160p"},
} {
t.Run(tc.name, func(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "downloads", "Inception 2010 2160p BluRay.mkv")
dest := filepath.Join(root, "media")
existing := filepath.Join(dest, "电影", "Inception (2010)", "Inception (2010).mkv")
writeOrgFile(t, src, "inception-2160p")
writeOrgFile(t, existing, "inception-1080p")
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
for key, value := range map[string]string{
"organizer.auto_after_download": "true",
"organize.target_dir": dest,
"organize.transfer_mode": "copy",
} {
if err := repos.Setting.Set(t.Context(), key, value); err != nil {
t.Fatal(err)
}
}
if err := repos.Media.Upsert(t.Context(), &model.Media{
Title: "Inception",
Path: existing,
Year: 2010,
Container: "mkv",
Width: 1920,
Height: 1080,
}); err != nil {
t.Fatal(err)
}
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
Source: "qbittorrent",
URL: "magnet:?xt=urn:btih:inception",
Title: "Inception 2010 2160p BluRay",
SavePath: filepath.Dir(src),
Status: "completed",
Progress: 1,
AllowExistingLibrary: tc.allowWash,
}); err != nil {
t.Fatal(err)
}
org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), org)
svc.onTorrentComplete(t.Context(), QBitTorrent{
Hash: "inception",
Name: "Inception 2010 2160p BluRay",
Progress: 1,
SavePath: filepath.Dir(src),
ContentPath: src,
})
got, err := os.ReadFile(existing)
if err != nil {
t.Fatal(err)
}
if string(got) != tc.wantContent {
t.Fatalf("existing content = %q, want %q", string(got), tc.wantContent)
}
})
}
}
-142
View File
@@ -1,142 +0,0 @@
package service
import (
"testing"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestDownloadPollBaselinesAlreadyCompletedTorrents(t *testing.T) {
repos := newOrganizerTestRepo(t)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{{
Hash: "already-complete",
Name: "Already Complete S01E01",
Progress: 1,
State: "stalledUP",
}}, nil)
if got := len(svc.organizeQueue); got != 0 {
t.Fatalf("first poll queued %d organize jobs, want 0", got)
}
if !svc.prevStates["already-complete"] {
t.Fatal("first poll should remember completed baseline state")
}
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{{
Hash: "late-complete",
Name: "Late Complete S01E01",
Progress: 1,
State: "stalledUP",
}}, nil)
if got := len(svc.organizeQueue); got != 0 {
t.Fatalf("newly discovered completed torrent queued %d organize jobs, want 0", got)
}
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{{
Hash: "new-download",
Name: "New Download S01E01",
Progress: 0.5,
}}, nil)
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{{
Hash: "new-download",
Name: "New Download S01E01",
Progress: 1,
State: "stalledUP",
}}, nil)
if got := len(svc.organizeQueue); got != 1 {
t.Fatalf("completion transition queued %d organize jobs, want 1", got)
}
}
func TestDownloadPollCatchesUpRecentlyCompletedTorrents(t *testing.T) {
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
task := &model.DownloadTask{
Source: "qbittorrent",
URL: "magnet:?xt=urn:btih:fresh",
Title: "Fresh Complete S01E01",
SavePath: "/downloads",
Status: "queued",
Progress: 0,
}
if err := repos.Download.Create(t.Context(), task); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(t.Context(), "organizer.auto_after_download", "true"); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{
{Hash: "fresh-complete", Name: "Fresh Complete S01E01", Progress: 1, State: "stalledUP", CompletionOn: time.Now().Add(-time.Hour).Unix()},
{Hash: "stale-complete", Name: "Stale Complete S01E01", Progress: 1, State: "stalledUP", CompletionOn: time.Now().Add(-48 * time.Hour).Unix()},
{Hash: "no-timestamp", Name: "No Timestamp S01E01", Progress: 1, State: "stalledUP"},
}, tasksByTorrentIdentity([]model.DownloadTask{*task}))
// 只有补整理时间窗内、且存在本地追踪任务的种子会被补整理;无 completion_on 的保守跳过。
if got := len(svc.organizeQueue); got != 1 {
t.Fatalf("first poll queued %d organize jobs, want 1 (recent tracked completion only)", got)
}
}
func TestDownloadPollDoesNotCatchUpWhenAutoOrganizeDisabled(t *testing.T) {
repos := newOrganizerTestRepo(t)
if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil {
t.Fatal(err)
}
task := &model.DownloadTask{
Source: "qbittorrent",
URL: "magnet:?xt=urn:btih:fresh",
Title: "Fresh Complete S01E01",
SavePath: "/downloads",
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: "fresh-complete",
Name: "Fresh Complete S01E01",
Progress: 1,
State: "stalledUP",
CompletionOn: time.Now().Add(-time.Hour).Unix(),
}
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{torrent}, tasksByTorrentIdentity([]model.DownloadTask{*task}))
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{torrent}, tasksByTorrentIdentity([]model.DownloadTask{*task}))
if got := len(svc.organizeQueue); got != 0 {
t.Fatalf("auto-organize disabled queued %d completed jobs, want 0", got)
}
}
func TestDownloadPollSkipsRecordedCompletedTorrentCatchup(t *testing.T) {
repos := newOrganizerTestRepo(t)
torrent := QBitTorrent{
Hash: "fresh-complete",
Name: "Fresh Complete S01E01",
Progress: 1,
State: "stalledUP",
CompletionOn: time.Now().Add(-time.Hour).Unix(),
}
if err := repos.Setting.Set(t.Context(), completedTorrentCatchupSettingKey(torrent), "true"); err != nil {
t.Fatal(err)
}
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
svc.processDownloadSnapshot(t.Context(), []QBitTorrent{torrent}, nil)
if got := len(svc.organizeQueue); got != 0 {
t.Fatalf("recorded completed torrent queued %d organize jobs, want 0", got)
}
}

Some files were not shown because too many files have changed in this diff Show More