mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-06 21:36:37 +08:00
初始化
初始化项目
This commit is contained in:
@@ -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 后重试。"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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() }
|
||||
@@ -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]
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 = ©
|
||||
continue
|
||||
}
|
||||
if lib.ID == rootLibraryID || CloudLibraryAutoCategory(lib) || !lib.Enabled || targetKey == "" {
|
||||
continue
|
||||
}
|
||||
key, ok := CloudLibraryMergeKey(lib)
|
||||
if target == nil && ok && key == targetKey {
|
||||
copy := lib
|
||||
target = ©
|
||||
}
|
||||
}
|
||||
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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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:")
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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", "")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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 ©, adapter, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, id := range m.order {
|
||||
if adapter, ok := m.clients[id]; ok {
|
||||
client := m.models[id]
|
||||
copy := client
|
||||
return ©, 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])
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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, ", "))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user