diff --git a/.gitattributes b/.gitattributes
index 4da7fb3..1e6e344 100644
--- a/.gitattributes
+++ b/.gitattributes
@@ -7,3 +7,14 @@
Dockerfile text eol=lf
*.ps1 text eol=crlf
+
+# GitHub Linguist: keep repository language stats focused on product code
+# (Go backend + React/TypeScript frontend + Docker packaging). Deployment
+# helpers, generated lock files, and static brand assets are still tracked but
+# should not appear as primary project languages.
+scripts/** linguist-vendored
+docker-entrypoint.sh linguist-vendored
+web/package-lock.json linguist-generated
+web/*.config.js linguist-vendored
+web/public/** linguist-vendored
+web/src/**/*.css linguist-vendored
diff --git a/Dockerfile b/Dockerfile
index cae4285..c0888a6 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -85,24 +85,9 @@ EXPOSE 8080
HEALTHCHECK --interval=30s --timeout=5s --start-period=15s --retries=3 \
CMD busybox wget -q --spider http://127.0.0.1:8080/api/health || exit 1
-# Tiny entrypoint that lets us swap to a different UID/GID via PUID/PGID
-# (handy on NAS deployments where bind-mounted volumes belong to a non-root
-# user). When PUID == 0 we skip su-exec entirely and run as root.
-RUN printf '#!/bin/sh\n\
-PUID=${PUID:-$(id -u mediastation)}\n\
-PGID=${PGID:-$(id -g mediastation)}\n\
-if [ "$PUID" != "$(id -u mediastation)" ] || [ "$PGID" != "$(id -g mediastation)" ]; then\n\
- deluser mediastation 2>/dev/null || true\n\
- delgroup mediastation 2>/dev/null || true\n\
- addgroup -g "$PGID" -S mediastation\n\
- adduser -u "$PUID" -G mediastation -S mediastation\n\
-fi\n\
-chown -R mediastation:mediastation /data /cache 2>/dev/null || true\n\
-chown mediastation:mediastation /media 2>/dev/null || true\n\
-if [ "$PUID" = "0" ]; then\n\
- exec mediastation-go\n\
-fi\n\
-exec su-exec mediastation mediastation-go\n' > /entrypoint.sh \
- && chmod +x /entrypoint.sh
+# Tiny entrypoint that lets us run as a NAS host UID/GID via PUID/PGID without
+# rewriting /etc/passwd or /etc/group on every container start.
+COPY docker-entrypoint.sh /entrypoint.sh
+RUN chmod +x /entrypoint.sh
CMD ["/entrypoint.sh"]
diff --git a/cmd/server/main.go b/cmd/server/main.go
index 2998e10..a65755d 100644
--- a/cmd/server/main.go
+++ b/cmd/server/main.go
@@ -218,14 +218,14 @@ func serveSPA(r *gin.Engine, webDir string) {
assets.Static("/", filepath.Join(webDir, "assets"))
brand := r.Group("/brand")
brand.Use(func(c *gin.Context) {
- c.Header("Cache-Control", "public, max-age=86400")
+ setNoCacheHeaders(c)
c.Next()
})
brand.Static("/", filepath.Join(webDir, "brand"))
- for _, icon := range []string{"/favicon.ico", "/favicon.svg"} {
- iconPath := filepath.Join(webDir, strings.TrimPrefix(icon, "/"))
- r.GET(icon, serveNoCacheFile(iconPath))
- r.HEAD(icon, serveNoCacheFile(iconPath))
+ for _, rootFile := range []string{"/favicon.ico", "/favicon.svg", "/artwork-cache-sw.js"} {
+ filePath := filepath.Join(webDir, strings.TrimPrefix(rootFile, "/"))
+ r.GET(rootFile, serveNoCacheFile(filePath))
+ r.HEAD(rootFile, serveNoCacheFile(filePath))
}
r.NoRoute(func(c *gin.Context) {
path := c.Request.URL.Path
diff --git a/cmd/server/main_test.go b/cmd/server/main_test.go
index 93ab70b..0f2ce25 100644
--- a/cmd/server/main_test.go
+++ b/cmd/server/main_test.go
@@ -64,6 +64,9 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
if err := os.WriteFile(filepath.Join(webDir, "brand", "mediastationgo-logo.svg"), []byte(""), 0o644); err != nil {
t.Fatal(err)
}
+ if err := os.WriteFile(filepath.Join(webDir, "artwork-cache-sw.js"), []byte("self.addEventListener('fetch', () => {})"), 0o644); err != nil {
+ t.Fatal(err)
+ }
router := gin.New()
serveSPA(router, webDir)
@@ -84,10 +87,26 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
if brandResp.Code != http.StatusOK {
t.Fatalf("brand asset status = %d, want 200", brandResp.Code)
}
+ if got := brandResp.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
+ t.Fatalf("brand asset Cache-Control = %q, want no-store", got)
+ }
if strings.Contains(brandResp.Body.String(), "index") {
t.Fatalf("brand asset should not serve SPA index: %q", brandResp.Body.String())
}
+ swReq := httptest.NewRequest(http.MethodGet, "/artwork-cache-sw.js", nil)
+ swResp := httptest.NewRecorder()
+ router.ServeHTTP(swResp, swReq)
+ if swResp.Code != http.StatusOK {
+ t.Fatalf("service worker status = %d, want 200", swResp.Code)
+ }
+ if got := swResp.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
+ t.Fatalf("service worker Cache-Control = %q, want no-store", got)
+ }
+ if strings.Contains(swResp.Body.String(), "index") {
+ t.Fatalf("service worker should not serve SPA index: %q", swResp.Body.String())
+ }
+
for _, path := range []string{
"/api/missing",
"/emby",
diff --git a/docker-entrypoint.sh b/docker-entrypoint.sh
new file mode 100644
index 0000000..ed4000e
--- /dev/null
+++ b/docker-entrypoint.sh
@@ -0,0 +1,28 @@
+#!/bin/sh
+set -eu
+
+run_uid="${PUID:-$(id -u mediastation 2>/dev/null || echo 1000)}"
+run_gid="${PGID:-$(id -g mediastation 2>/dev/null || echo 1000)}"
+
+case "$run_uid" in
+ ''|*[!0-9]*)
+ echo "PUID must be a numeric uid, got: $run_uid" >&2
+ exit 1
+ ;;
+esac
+
+case "$run_gid" in
+ ''|*[!0-9]*)
+ echo "PGID must be a numeric gid, got: $run_gid" >&2
+ exit 1
+ ;;
+esac
+
+if [ "$run_uid" = "0" ]; then
+ exec mediastation-go
+fi
+
+chown -R "$run_uid:$run_gid" /data /cache 2>/dev/null || true
+chown "$run_uid:$run_gid" /media 2>/dev/null || true
+
+exec su-exec "$run_uid:$run_gid" mediastation-go
diff --git a/internal/database/database.go b/internal/database/database.go
index ca85fdd..e56b661 100644
--- a/internal/database/database.go
+++ b/internal/database/database.go
@@ -3,42 +3,25 @@
package database
import (
- "context"
"errors"
"fmt"
- "os"
- "path/filepath"
- "reflect"
- "sort"
"strings"
- "time"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/driver/postgres"
"gorm.io/gorm"
- "gorm.io/gorm/clause"
"gorm.io/gorm/logger"
- "gorm.io/gorm/schema"
"github.com/ShukeBta/MediaStationGo/internal/config"
- "github.com/ShukeBta/MediaStationGo/internal/model"
)
// Open initialises the configured GORM database. database.type=auto chooses
-// PostgreSQL when database.dsn is present (the Docker Compose default) and
-// otherwise falls back to SQLite for old/bare-metal installs.
+// PostgreSQL when database.dsn is present and otherwise falls back to SQLite.
func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
- gormLogger := logger.New(
- zapStdLogger{log: log},
- logger.Config{
- SlowThreshold: 0,
- LogLevel: logger.Warn,
- IgnoreRecordNotFoundError: true,
- Colorful: false,
- },
- )
-
+ if cfg == nil {
+ return nil, errors.New("database config is required")
+ }
dialect := normalizeDatabaseType(cfg.Database.Type)
if dialect == "auto" {
dialect = effectiveAutoDatabaseType(cfg)
@@ -48,7 +31,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
return nil, err
}
db, err := gorm.Open(dialector, &gorm.Config{
- Logger: gormLogger,
+ Logger: newGormLogger(log),
PrepareStmt: true,
DisableForeignKeyConstraintWhenMigrating: false,
})
@@ -58,9 +41,31 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
if dialect == "sqlite" {
installSQLiteWriteGate(db)
}
+ if err := configureConnectionPool(db, cfg); err != nil {
+ return nil, err
+ }
+ return db, nil
+}
+
+func newGormLogger(log *zap.Logger) logger.Interface {
+ if log == nil {
+ log = zap.NewNop()
+ }
+ return logger.New(
+ zapStdLogger{log: log},
+ logger.Config{
+ SlowThreshold: 0,
+ LogLevel: logger.Warn,
+ IgnoreRecordNotFoundError: true,
+ Colorful: false,
+ },
+ )
+}
+
+func configureConnectionPool(db *gorm.DB, cfg *config.Config) error {
sqlDB, err := db.DB()
if err != nil {
- return nil, fmt.Errorf("gorm sqldb: %w", err)
+ return fmt.Errorf("gorm sqldb: %w", err)
}
if cfg.Database.MaxOpenConns > 0 {
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
@@ -68,7 +73,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
if cfg.Database.MaxIdleConns > 0 {
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
}
- return db, nil
+ return nil
}
func normalizeDatabaseType(value string) string {
@@ -106,718 +111,12 @@ func databaseDialector(cfg *config.Config, dialect string) (gorm.Dialector, erro
}
}
-// MigrateSQLiteToCurrentIfNeeded copies an existing SQLite database into
-// PostgreSQL. Redis is not migrated because it is a rebuildable cache, not a
-// source of truth.
-const sqliteMigrationCompleteSettingKey = "database.sqlite_migration_complete"
-
-func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *zap.Logger) error {
- if cfg == nil || target == nil || target.Dialector == nil || target.Dialector.Name() != "postgres" {
- return nil
- }
- sqlitePath, err := sqliteMigrationSourcePath(cfg, log)
- if err != nil {
- return err
- }
- if sqlitePath == "" {
- return nil
- }
- if complete, err := sqliteMigrationMarkedComplete(target); err != nil {
- return err
- } else if complete {
- if log != nil {
- log.Info("skip sqlite to postgres migration: already completed")
- }
- return nil
- }
-
- src, err := openSQLiteMigrationSource(cfg, sqlitePath)
- if err != nil {
- return fmt.Errorf("open sqlite migration source: %w", err)
- }
- sqlDB, err := src.DB()
- if err == nil {
- defer sqlDB.Close()
- }
-
- started := time.Now()
- if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil {
- return err
- }
- copied, err := copyModelTables(src, target, 500)
- if err != nil {
- return err
- }
- if err := markSQLiteMigrationComplete(target); err != nil {
- return err
- }
- if log != nil {
- log.Info("sqlite data migrated to postgres",
- zap.String("source", sqlitePath),
- zap.Int64("rows", copied),
- zap.Duration("duration", time.Since(started)))
- }
- return nil
-}
-
-func openSQLiteMigrationSource(cfg *config.Config, sqlitePath string) (*gorm.DB, error) {
- srcCfg := *cfg
- srcCfg.Database.Type = "sqlite"
- srcCfg.Database.DBPath = sqlitePath
- return gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{
- Logger: logger.Default.LogMode(logger.Silent),
- })
-}
-
-func sqliteMigrationSourcePath(cfg *config.Config, log *zap.Logger) (string, error) {
- configured := strings.TrimSpace(cfg.Database.DBPath)
- if configured != "" {
- exists, err := regularFileExists(configured)
- if err != nil {
- return "", fmt.Errorf("stat sqlite migration source: %w", err)
- }
- if exists {
- return configured, nil
- }
- }
-
- fallback := filepath.Join(strings.TrimSpace(cfg.App.DataDir), "mediastation.db")
- if fallback == "" || sameCleanPath(configured, fallback) {
- return "", nil
- }
- exists, err := regularFileExists(fallback)
- if err != nil {
- return "", fmt.Errorf("stat default sqlite migration source: %w", err)
- }
- if !exists {
- return "", nil
- }
- if log != nil && configured != "" {
- log.Warn("configured sqlite migration source not found; using data-dir default",
- zap.String("configured", configured),
- zap.String("fallback", fallback))
- }
- return fallback, nil
-}
-
-func regularFileExists(path string) (bool, error) {
- if strings.TrimSpace(path) == "" {
- return false, nil
- }
- info, err := os.Stat(path)
- if err != nil {
- if errors.Is(err, os.ErrNotExist) {
- return false, nil
- }
- return false, err
- }
- return !info.IsDir(), nil
-}
-
-func sameCleanPath(a, b string) bool {
- if a == "" || b == "" {
- return false
- }
- return filepath.Clean(a) == filepath.Clean(b)
-}
-
-func resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target *gorm.DB, log *zap.Logger) error {
- hasRows, err := sqliteSourceHasMigratableRows(src)
- if err != nil {
- return err
- }
- if !hasRows {
- return nil
- }
- bootstrapOnly, err := targetLooksLikeBootstrapOnly(target)
- if err != nil || !bootstrapOnly {
- return err
- }
- for i := len(model.AllModels()) - 1; i >= 0; i-- {
- m := model.AllModels()[i]
- if !target.Migrator().HasTable(m) {
- continue
- }
- if err := target.Session(&gorm.Session{AllowGlobalUpdate: true}).Unscoped().Delete(m).Error; err != nil {
- return fmt.Errorf("clear bootstrap target table %T: %w", m, err)
- }
- }
- if log != nil {
- log.Warn("cleared bootstrap postgres rows before sqlite migration")
- }
- return nil
-}
-
-func sqliteSourceHasMigratableRows(src *gorm.DB) (bool, error) {
- for _, table := range []string{"users", "libraries", "media", "settings"} {
- exists, err := sqliteTableExists(src, table)
- if err != nil {
- return false, err
- }
- if !exists {
- continue
- }
- var count int64
- if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&count).Error; err != nil {
- return false, fmt.Errorf("count sqlite table %s: %w", table, err)
- }
- if count > 0 {
- return true, nil
- }
- }
- return false, nil
-}
-
-func targetLooksLikeBootstrapOnly(target *gorm.DB) (bool, error) {
- for _, m := range []any{
- &model.Library{},
- &model.Series{},
- &model.Media{},
- &model.PlaybackHistory{},
- &model.Favorite{},
- &model.Playlist{},
- &model.PlaylistItem{},
- &model.DownloadTask{},
- &model.Subscription{},
- } {
- if !target.Migrator().HasTable(m) {
- continue
- }
- var count int64
- if err := target.Unscoped().Model(m).Count(&count).Error; err != nil {
- return false, err
- }
- if count > 0 {
- return false, nil
- }
- }
-
- var userCount int64
- if !target.Migrator().HasTable(&model.User{}) {
- return true, nil
- }
- if err := target.Model(&model.User{}).Count(&userCount).Error; err != nil {
- return false, err
- }
- if userCount == 0 {
- return true, nil
- }
- if userCount != 1 {
- return false, nil
- }
- var user model.User
- if err := target.Unscoped().Where("username = ?", "admin").First(&user).Error; err != nil {
- return false, nil
- }
- return user.Role == "admin", nil
-}
-
-func sqliteMigrationMarkedComplete(db *gorm.DB) (bool, error) {
- var value string
- err := db.Raw("SELECT value FROM "+quoteIdent("settings")+" WHERE "+quoteIdent("key")+" = ?", sqliteMigrationCompleteSettingKey).Scan(&value).Error
- if err != nil {
- return false, fmt.Errorf("check sqlite migration marker: %w", err)
- }
- return strings.EqualFold(strings.TrimSpace(value), "true"), nil
-}
-
-func markSQLiteMigrationComplete(db *gorm.DB) error {
- now := time.Now()
- if err := db.Clauses(clause.OnConflict{
- Columns: []clause.Column{{Name: "key"}},
- DoUpdates: clause.AssignmentColumns([]string{"value", "updated_at"}),
- }).Create(&model.Setting{
- Key: sqliteMigrationCompleteSettingKey,
- Value: "true",
- UpdatedAt: now,
- }).Error; err != nil {
- return fmt.Errorf("mark sqlite migration complete: %w", err)
- }
- return nil
-}
-
-func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
- if batchSize <= 0 {
- batchSize = 500
- }
- var copied int64
- for _, m := range model.AllModels() {
- table, err := modelTableName(src, m)
- if err != nil {
- return copied, err
- }
- primaryColumns, err := modelPrimaryColumns(src, m)
- if err != nil {
- return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
- }
- exists, err := sqliteTableExists(src, table)
- if err != nil {
- return copied, err
- }
- if !exists {
- continue
- }
- var sourceCount int64
- if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil {
- return copied, fmt.Errorf("count sqlite table %s: %w", table, err)
- }
- if sourceCount == 0 {
- continue
- }
- var targetCount int64
- if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil {
- return copied, fmt.Errorf("count target table %s: %w", table, err)
- }
- modelType := reflect.TypeOf(m)
- if modelType.Kind() != reflect.Ptr {
- return copied, fmt.Errorf("model %T is not a pointer", m)
- }
- sliceType := reflect.SliceOf(modelType.Elem())
- slicePtr := reflect.New(sliceType)
- if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
- return copied, fmt.Errorf("read sqlite table %s: %w", table, err)
- }
- filtered := slicePtr.Elem()
- if targetCount > 0 {
- primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
- if err != nil {
- return copied, err
- }
- filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
- }
- if filtered.Len() == 0 {
- continue
- }
- filteredPtr := reflect.New(filtered.Type())
- filteredPtr.Elem().Set(filtered)
- if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
- return copied, fmt.Errorf("copy sqlite table %s: %w", table, err)
- }
- copied += int64(filtered.Len())
- }
- return copied, nil
-}
-
-func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) {
- stmt := &gorm.Statement{DB: db}
- if err := stmt.Parse(m); err != nil {
- return nil, err
- }
- var cols []string
- for _, field := range stmt.Schema.PrimaryFields {
- cols = append(cols, field.DBName)
- }
- if len(cols) == 0 {
- return nil, fmt.Errorf("no primary key columns")
- }
- return cols, nil
-}
-
-func targetPrimaryKeySet(target *gorm.DB, table string, primaryColumns []string) (map[string]struct{}, error) {
- if len(primaryColumns) != 1 {
- return nil, nil
- }
- var values []string
- if err := target.Raw("SELECT " + quoteIdent(primaryColumns[0]) + " FROM " + quoteIdent(table)).Scan(&values).Error; err != nil {
- return nil, fmt.Errorf("read target primary keys for table %s: %w", table, err)
- }
- set := make(map[string]struct{}, len(values))
- for _, value := range values {
- set[value] = struct{}{}
- }
- return set, nil
-}
-
-func filterRowsMissingInTarget(target *gorm.DB, table string, primaryColumns []string, rows reflect.Value, primaryKeySet map[string]struct{}) reflect.Value {
- if rows.Kind() != reflect.Slice || rows.Len() == 0 || len(primaryColumns) == 0 {
- return rows
- }
- out := reflect.MakeSlice(rows.Type(), 0, rows.Len())
- for i := 0; i < rows.Len(); i++ {
- row := rows.Index(i)
- keys, ok := rowPrimaryKeys(row, primaryColumns)
- if !ok {
- out = reflect.Append(out, row)
- continue
- }
- if primaryKeySet != nil {
- if _, exists := primaryKeySet[fmt.Sprint(keys[primaryColumns[0]])]; !exists {
- out = reflect.Append(out, row)
- }
- continue
- }
- if !targetHasPrimaryKey(target, table, keys) {
- out = reflect.Append(out, row)
- }
- }
- return out
-}
-
-func rowPrimaryKeys(row reflect.Value, primaryColumns []string) (map[string]any, bool) {
- if row.Kind() == reflect.Pointer {
- if row.IsNil() {
- return nil, false
- }
- row = row.Elem()
- }
- if row.Kind() != reflect.Struct {
- return nil, false
- }
- keys := make(map[string]any, len(primaryColumns))
- for _, column := range primaryColumns {
- value, ok := fieldByDBName(row, column)
- if !ok || value.IsZero() {
- return nil, false
- }
- keys[column] = value.Interface()
- }
- return keys, true
-}
-
-func fieldByDBName(row reflect.Value, column string) (reflect.Value, bool) {
- rowType := row.Type()
- for i := 0; i < row.NumField(); i++ {
- fieldType := rowType.Field(i)
- field := row.Field(i)
- if fieldType.Anonymous {
- if value, ok := fieldByDBName(field, column); ok {
- return value, true
- }
- }
- if columnNameForStructField(fieldType) == column {
- if field.Kind() == reflect.Pointer && field.IsNil() {
- return reflect.Value{}, false
- }
- return field, field.CanInterface()
- }
- }
- return reflect.Value{}, false
-}
-
-func columnNameForStructField(field reflect.StructField) string {
- if field.PkgPath != "" && !field.Anonymous {
- return ""
- }
- tag := field.Tag.Get("gorm")
- settings := schema.ParseTagSetting(tag, ";")
- if column := settings["COLUMN"]; column != "" {
- return column
- }
- return schema.NamingStrategy{}.ColumnName("", field.Name)
-}
-
-func targetHasPrimaryKey(target *gorm.DB, table string, keys map[string]any) bool {
- where := make([]string, 0, len(keys))
- args := make([]any, 0, len(keys))
- for _, column := range sortedMapKeys(keys) {
- where = append(where, quoteIdent(column)+" = ?")
- args = append(args, keys[column])
- }
- var count int64
- err := target.Raw("SELECT COUNT(1) FROM "+quoteIdent(table)+" WHERE "+strings.Join(where, " AND "), args...).Scan(&count).Error
- return err == nil && count > 0
-}
-
-func sortedMapKeys(m map[string]any) []string {
- keys := make([]string, 0, len(m))
- for key := range m {
- keys = append(keys, key)
- }
- sort.Strings(keys)
- return keys
-}
-
-func sqliteTableExists(db *gorm.DB, table string) (bool, error) {
- var count int64
- if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
- return false, fmt.Errorf("inspect sqlite table %s: %w", table, err)
- }
- return count > 0, nil
-}
-
-func modelTableName(db *gorm.DB, m any) (string, error) {
- stmt := &gorm.Statement{DB: db}
- if err := stmt.Parse(m); err != nil {
- return "", err
- }
- return stmt.Schema.Table, nil
-}
-
-func quoteIdent(value string) string {
- return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
-}
-
-func installSQLiteWriteGate(db *gorm.DB) {
- if db == nil {
- return
- }
- const lockedKey = "mediastation:sqlite_write_locked"
- gate := newSQLiteWriteGate()
- lock := func(tx *gorm.DB) {
- ctx := context.Background()
- if tx.Statement != nil && tx.Statement.Context != nil {
- ctx = tx.Statement.Context
- }
- if err := gate.Lock(ctx); err != nil {
- _ = tx.AddError(err)
- return
- }
- tx.InstanceSet(lockedKey, struct{}{})
- }
- unlock := func(tx *gorm.DB) {
- if _, ok := tx.InstanceGet(lockedKey); ok {
- gate.Unlock()
- }
- }
- _ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock)
- _ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock)
- _ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock)
- _ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock)
- _ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock)
- _ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock)
- _ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock)
- _ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock)
-}
-
-// sqliteWriteGate 串行化进程内的 SQLite 写操作,避免多连接写竞争触发
-// SQLITE_BUSY。Lock 尊重语句自身的 context:此前用 sync.Mutex 时,一条
-// 长写语句(如 FTS 回填批次)会让登录等关键写操作无限期排队——客户端
-// 早已超时断开,goroutine 还挂在互斥锁上。现在等待方可随 context 取消
-// 及时失败,不再把整个进程的写路径拖死。
-type sqliteWriteGate struct {
- ch chan struct{}
-}
-
-func newSQLiteWriteGate() *sqliteWriteGate {
- return &sqliteWriteGate{ch: make(chan struct{}, 1)}
-}
-
-func (g *sqliteWriteGate) Lock(ctx context.Context) error {
- select {
- case g.ch <- struct{}{}:
- return nil
- default:
- }
- if ctx == nil {
- ctx = context.Background()
- }
- select {
- case g.ch <- struct{}{}:
- return nil
- case <-ctx.Done():
- return ctx.Err()
- }
-}
-
-func (g *sqliteWriteGate) Unlock() {
- select {
- case <-g.ch:
- default:
- }
-}
-
-func buildSQLiteDSN(cfg *config.Config) string {
- dbPath := cfg.Database.DBPath
- if !filepath.IsAbs(dbPath) {
- // keep as-is to respect user-provided relative paths.
- dbPath = filepath.Clean(dbPath)
- }
- dsn := dbPath + "?_pragma=foreign_keys(1)"
- if cfg.Database.WALMode {
- dsn += "&_pragma=journal_mode(WAL)"
- }
- if cfg.Database.BusyTimeout > 0 {
- dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout)
- }
- if cfg.Database.CacheSize != 0 {
- dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize)
- }
- return dsn
-}
-
-// AutoMigrate creates tables for every model registered in the model package.
-func AutoMigrate(db *gorm.DB) error {
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- return err
- }
- if err := ensurePostgresColumnCompatibility(db); err != nil {
- return err
- }
- if err := enforceTelegramBindingOneToOne(db); err != nil {
- return err
- }
- if err := ensurePerformanceIndexes(db); err != nil {
- return err
- }
- if isSQLite(db) {
- return ensureMediaSearchIndex(db)
- }
- return nil
-}
-
-func ensurePostgresColumnCompatibility(db *gorm.DB) error {
- if !isPostgres(db) {
- return nil
- }
- statements := []string{
- `ALTER TABLE media ALTER COLUMN container TYPE varchar(128)`,
- `ALTER TABLE media ALTER COLUMN genres TYPE text`,
- `ALTER TABLE media ALTER COLUMN series_id TYPE varchar(128)`,
- `ALTER TABLE media ALTER COLUMN duplicate_of TYPE varchar(128)`,
- `ALTER TABLE playback_histories ALTER COLUMN media_id TYPE varchar(128)`,
- `ALTER TABLE favorites ALTER COLUMN media_id TYPE varchar(128)`,
- `ALTER TABLE playlist_items ALTER COLUMN media_id TYPE varchar(128)`,
- `ALTER TABLE strm_records ALTER COLUMN media_id TYPE varchar(128)`,
- }
- for _, stmt := range statements {
- if err := db.Exec(stmt).Error; err != nil {
- return err
- }
- }
- return nil
-}
-
-func ensurePerformanceIndexes(db *gorm.DB) error {
- statements := []string{
- `CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_favorites_user_media_active ON favorites(user_id, media_id) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_playback_histories_user_media_active ON playback_histories(user_id, media_id, watched_at DESC) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_playback_histories_resume_active ON playback_histories(user_id, completed, watched_at DESC) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_play_profiles_user_created_active ON play_profiles(user_id, created_at DESC) WHERE deleted_at IS NULL`,
- }
- if isSQLite(db) {
- statements = append(statements,
- `CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title COLLATE NOCASE) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name COLLATE NOCASE) WHERE deleted_at IS NULL`,
- )
- } else {
- statements = append(statements,
- `CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title) WHERE deleted_at IS NULL`,
- `CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name) WHERE deleted_at IS NULL`,
- )
- }
- for _, stmt := range statements {
- if err := db.Exec(stmt).Error; err != nil {
- return err
- }
- }
- return nil
-}
-
-func isSQLite(db *gorm.DB) bool {
- return db != nil && db.Dialector != nil && db.Dialector.Name() == "sqlite"
-}
-
-func isPostgres(db *gorm.DB) bool {
- return db != nil && db.Dialector != nil && db.Dialector.Name() == "postgres"
-}
-
-// mediaSearchIndexSchemaVersion 标识 FTS 索引的物理布局版本。
-// v2:FTS 行的 rowid 与 media.rowid 对齐,并由触发器实时维护。
-const mediaSearchIndexSchemaVersion = 2
-
-func ensureMediaSearchIndex(db *gorm.DB) error {
- if err := db.Exec(`CREATE TABLE IF NOT EXISTS media_search_meta (id INTEGER PRIMARY KEY CHECK (id = 1), version INTEGER NOT NULL)`).Error; err != nil {
- return nil
- }
- var version int
- _ = db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error
- if version != mediaSearchIndexSchemaVersion {
- // 旧版(v1)FTS 表按 UNINDEXED 的 media_id 寻址。FTS5 的普通列
- // 不支持索引查找,按 media_id 的 DELETE / NOT EXISTS 都是整表
- // 扫描:十几万行的库每次启动回填要做上百亿次行访问,纯 Go
- // sqlite 直接把 CPU 钉满数小时,并隔着全局写锁拖死登录。
- // v2 起 FTS 行的 rowid 与 media.rowid 对齐,所有寻址走 rowid
- // 点查,索引一致性交给下方触发器维护。
- for _, stmt := range []string{
- `DROP TRIGGER IF EXISTS media_search_fts_ai`,
- `DROP TRIGGER IF EXISTS media_search_fts_au`,
- `DROP TRIGGER IF EXISTS media_search_fts_ad`,
- `DROP TABLE IF EXISTS media_search_fts`,
- } {
- _ = db.Exec(stmt).Error
- }
- }
- if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='trigram')`).Error; err != nil {
- if fallbackErr := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='unicode61')`).Error; fallbackErr != nil {
- // FTS is an acceleration path. Some embedded SQLite builds may omit
- // FTS5; keep startup working and let repository queries fall back to
- // LIKE-based Chinese fuzzy search.
- return nil
- }
- }
- // 触发器让 FTS 与 media 行保持同步(新增/标题刮削改写/软删/恢复/
- // 硬删全覆盖),应用层不再需要按 media_id 手工刷新索引——也顺带
- // 修复了刮削直写 Updates() 后新标题搜不到的问题。
- for _, stmt := range []string{
- `CREATE TRIGGER IF NOT EXISTS media_search_fts_ai AFTER INSERT ON media WHEN new.deleted_at IS NULL BEGIN
- DELETE FROM media_search_fts WHERE rowid = new.rowid;
- INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
- VALUES (new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, ''));
- END`,
- `CREATE TRIGGER IF NOT EXISTS media_search_fts_au AFTER UPDATE OF title, original_name, path, genres, deleted_at ON media BEGIN
- DELETE FROM media_search_fts WHERE rowid = old.rowid;
- INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
- SELECT new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, '')
- WHERE new.deleted_at IS NULL;
- END`,
- `CREATE TRIGGER IF NOT EXISTS media_search_fts_ad AFTER DELETE ON media BEGIN
- DELETE FROM media_search_fts WHERE rowid = old.rowid;
- END`,
- } {
- if err := db.Exec(stmt).Error; err != nil {
- return err
- }
- }
- if version != mediaSearchIndexSchemaVersion {
- if err := db.Exec(`INSERT INTO media_search_meta(id, version) VALUES (1, ?) ON CONFLICT(id) DO UPDATE SET version = excluded.version`, mediaSearchIndexSchemaVersion).Error; err != nil {
- return err
- }
- }
- return nil
-}
-
-func enforceTelegramBindingOneToOne(db *gorm.DB) error {
- if !db.Migrator().HasTable(&model.TelegramBinding{}) {
- return nil
- }
- return db.Transaction(func(tx *gorm.DB) error {
- if err := tx.Exec(`
-DELETE FROM telegram_bindings
-WHERE deleted_at IS NULL
- AND user_id IN (
- SELECT user_id
- FROM telegram_bindings
- WHERE deleted_at IS NULL
- GROUP BY user_id
- HAVING COUNT(*) > 1
- )
- AND id NOT IN (
- SELECT id
- FROM (
- SELECT id,
- ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY updated_at DESC, created_at DESC, id DESC) AS rn
- FROM telegram_bindings
- WHERE deleted_at IS NULL
- ) AS ranked_bindings
- WHERE rn = 1
- )
-`).Error; err != nil {
- return err
- }
- return tx.Exec(`
-CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_bindings_user_id_active
-ON telegram_bindings(user_id)
-WHERE deleted_at IS NULL
-`).Error
- })
-}
-
// zapStdLogger adapts a *zap.Logger to GORM's tiny logger interface.
type zapStdLogger struct{ log *zap.Logger }
func (z zapStdLogger) Printf(format string, args ...interface{}) {
+ if z.log == nil {
+ return
+ }
z.log.Sugar().Infof(format, args...)
}
diff --git a/internal/database/database_test.go b/internal/database/database_test.go
index c744ef0..c874cc7 100644
--- a/internal/database/database_test.go
+++ b/internal/database/database_test.go
@@ -2,6 +2,7 @@ package database
import (
"path/filepath"
+ "strings"
"testing"
"time"
@@ -12,6 +13,47 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/model"
)
+func TestOpenRequiresConfig(t *testing.T) {
+ db, err := Open(nil, nil)
+ if err == nil {
+ t.Fatal("expected nil config to return an error")
+ }
+ if db != nil {
+ t.Fatal("db should be nil when config is missing")
+ }
+ if !strings.Contains(err.Error(), "database config") {
+ t.Fatalf("error = %v, want database config message", err)
+ }
+}
+
+func TestOpenSQLiteWithNilLoggerConfiguresPool(t *testing.T) {
+ cfg := &config.Config{}
+ cfg.Database.Type = "sqlite"
+ cfg.Database.DBPath = filepath.Join(t.TempDir(), "mediastation.db")
+ cfg.Database.WALMode = true
+ cfg.Database.BusyTimeout = 5000
+ cfg.Database.CacheSize = -2000
+ cfg.Database.MaxOpenConns = 3
+ cfg.Database.MaxIdleConns = 2
+
+ db, err := Open(cfg, nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ sqlDB, err := db.DB()
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer sqlDB.Close()
+ if err := db.Exec("SELECT 1").Error; err != nil {
+ t.Fatal(err)
+ }
+ stats := sqlDB.Stats()
+ if stats.MaxOpenConnections != 3 {
+ t.Fatalf("MaxOpenConnections = %d, want 3", stats.MaxOpenConnections)
+ }
+}
+
func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
@@ -83,6 +125,58 @@ func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) {
}
}
+func TestEnsureMediaSearchIndexCreatesVersionedTriggers(t *testing.T) {
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ if err := ensureMediaSearchIndex(db); err != nil {
+ t.Fatal(err)
+ }
+ if !sqliteFTSTableExists(t, db, "media_search_fts") {
+ t.Skip("SQLite FTS5 is unavailable in this build")
+ }
+ var version int
+ if err := db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error; err != nil {
+ t.Fatal(err)
+ }
+ if version != mediaSearchIndexSchemaVersion {
+ t.Fatalf("media search schema version = %d, want %d", version, mediaSearchIndexSchemaVersion)
+ }
+ for _, trigger := range []string{"media_search_fts_ai", "media_search_fts_au", "media_search_fts_ad"} {
+ var count int
+ if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'trigger' AND name = ?`, trigger).Scan(&count).Error; err != nil {
+ t.Fatal(err)
+ }
+ if count != 1 {
+ t.Fatalf("trigger %s count = %d, want 1", trigger, count)
+ }
+ }
+ media := model.Media{LibraryID: "lib-1", Title: "中文搜索电影", Path: "/media/movie.mkv", Genres: "动画,冒险"}
+ if err := db.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ var indexed int
+ if err := db.Raw(`SELECT COUNT(1) FROM media_search_fts WHERE media_id = ?`, media.ID).Scan(&indexed).Error; err != nil {
+ t.Fatal(err)
+ }
+ if indexed != 1 {
+ t.Fatalf("indexed rows = %d, want inserted media indexed", indexed)
+ }
+}
+
+func sqliteFTSTableExists(t *testing.T, db *gorm.DB, table string) bool {
+ t.Helper()
+ var count int
+ if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
+ t.Fatal(err)
+ }
+ return count == 1
+}
+
func TestCopyModelTablesMigratesExistingSQLiteRows(t *testing.T) {
src, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
diff --git a/internal/database/schema_migration.go b/internal/database/schema_migration.go
new file mode 100644
index 0000000..cbd9845
--- /dev/null
+++ b/internal/database/schema_migration.go
@@ -0,0 +1,202 @@
+package database
+
+import (
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// AutoMigrate creates tables for every model registered in the model package.
+func AutoMigrate(db *gorm.DB) error {
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ return err
+ }
+ if err := ensurePostgresColumnCompatibility(db); err != nil {
+ return err
+ }
+ if err := enforceTelegramBindingOneToOne(db); err != nil {
+ return err
+ }
+ if err := ensurePerformanceIndexes(db); err != nil {
+ return err
+ }
+ if isSQLite(db) {
+ return ensureMediaSearchIndex(db)
+ }
+ return nil
+}
+
+func ensurePostgresColumnCompatibility(db *gorm.DB) error {
+ if !isPostgres(db) {
+ return nil
+ }
+ statements := []string{
+ `ALTER TABLE media ALTER COLUMN container TYPE varchar(128)`,
+ `ALTER TABLE media ALTER COLUMN genres TYPE text`,
+ `ALTER TABLE media ALTER COLUMN series_id TYPE varchar(128)`,
+ `ALTER TABLE media ALTER COLUMN duplicate_of TYPE varchar(128)`,
+ `ALTER TABLE playback_histories ALTER COLUMN media_id TYPE varchar(128)`,
+ `ALTER TABLE favorites ALTER COLUMN media_id TYPE varchar(128)`,
+ `ALTER TABLE playlist_items ALTER COLUMN media_id TYPE varchar(128)`,
+ `ALTER TABLE strm_records ALTER COLUMN media_id TYPE varchar(128)`,
+ }
+ for _, stmt := range statements {
+ if err := db.Exec(stmt).Error; err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func ensurePerformanceIndexes(db *gorm.DB) error {
+ statements := []string{
+ `CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_favorites_user_media_active ON favorites(user_id, media_id) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_playback_histories_user_media_active ON playback_histories(user_id, media_id, watched_at DESC) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_playback_histories_resume_active ON playback_histories(user_id, completed, watched_at DESC) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_play_profiles_user_created_active ON play_profiles(user_id, created_at DESC) WHERE deleted_at IS NULL`,
+ }
+ if isSQLite(db) {
+ statements = append(statements,
+ `CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title COLLATE NOCASE) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name COLLATE NOCASE) WHERE deleted_at IS NULL`,
+ )
+ } else {
+ statements = append(statements,
+ `CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title) WHERE deleted_at IS NULL`,
+ `CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name) WHERE deleted_at IS NULL`,
+ )
+ }
+ for _, stmt := range statements {
+ if err := db.Exec(stmt).Error; err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+// mediaSearchIndexSchemaVersion identifies the physical FTS index layout.
+// v2 aligns FTS rowids with media rowids and keeps the index current with
+// triggers.
+const mediaSearchIndexSchemaVersion = 2
+
+func ensureMediaSearchIndex(db *gorm.DB) error {
+ if err := ensureMediaSearchMetaTable(db); err != nil {
+ return nil
+ }
+ version := currentMediaSearchIndexVersion(db)
+ if version != mediaSearchIndexSchemaVersion {
+ resetMediaSearchIndex(db)
+ }
+ if !createMediaSearchFTSTable(db) {
+ return nil
+ }
+ if err := createMediaSearchTriggers(db); err != nil {
+ return err
+ }
+ if version != mediaSearchIndexSchemaVersion {
+ return markMediaSearchIndexVersion(db)
+ }
+ return nil
+}
+
+func ensureMediaSearchMetaTable(db *gorm.DB) error {
+ return db.Exec(`CREATE TABLE IF NOT EXISTS media_search_meta (id INTEGER PRIMARY KEY CHECK (id = 1), version INTEGER NOT NULL)`).Error
+}
+
+func currentMediaSearchIndexVersion(db *gorm.DB) int {
+ var version int
+ _ = db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error
+ return version
+}
+
+func resetMediaSearchIndex(db *gorm.DB) {
+ for _, stmt := range []string{
+ `DROP TRIGGER IF EXISTS media_search_fts_ai`,
+ `DROP TRIGGER IF EXISTS media_search_fts_au`,
+ `DROP TRIGGER IF EXISTS media_search_fts_ad`,
+ `DROP TABLE IF EXISTS media_search_fts`,
+ } {
+ _ = db.Exec(stmt).Error
+ }
+}
+
+func createMediaSearchFTSTable(db *gorm.DB) bool {
+ if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='trigram')`).Error; err == nil {
+ return true
+ }
+ if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='unicode61')`).Error; err == nil {
+ return true
+ }
+ // FTS is an acceleration path. Some embedded SQLite builds may omit FTS5;
+ // keep startup working and let repository queries fall back to LIKE search.
+ return false
+}
+
+func createMediaSearchTriggers(db *gorm.DB) error {
+ for _, stmt := range mediaSearchTriggerStatements {
+ if err := db.Exec(stmt).Error; err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+var mediaSearchTriggerStatements = []string{
+ `CREATE TRIGGER IF NOT EXISTS media_search_fts_ai AFTER INSERT ON media WHEN new.deleted_at IS NULL BEGIN
+ DELETE FROM media_search_fts WHERE rowid = new.rowid;
+ INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
+ VALUES (new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, ''));
+ END`,
+ `CREATE TRIGGER IF NOT EXISTS media_search_fts_au AFTER UPDATE OF title, original_name, path, genres, deleted_at ON media BEGIN
+ DELETE FROM media_search_fts WHERE rowid = old.rowid;
+ INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
+ SELECT new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, '')
+ WHERE new.deleted_at IS NULL;
+ END`,
+ `CREATE TRIGGER IF NOT EXISTS media_search_fts_ad AFTER DELETE ON media BEGIN
+ DELETE FROM media_search_fts WHERE rowid = old.rowid;
+ END`,
+}
+
+func markMediaSearchIndexVersion(db *gorm.DB) error {
+ return db.Exec(`INSERT INTO media_search_meta(id, version) VALUES (1, ?) ON CONFLICT(id) DO UPDATE SET version = excluded.version`, mediaSearchIndexSchemaVersion).Error
+}
+
+func enforceTelegramBindingOneToOne(db *gorm.DB) error {
+ if !db.Migrator().HasTable(&model.TelegramBinding{}) {
+ return nil
+ }
+ return db.Transaction(func(tx *gorm.DB) error {
+ if err := tx.Exec(`
+DELETE FROM telegram_bindings
+WHERE deleted_at IS NULL
+ AND user_id IN (
+ SELECT user_id
+ FROM telegram_bindings
+ WHERE deleted_at IS NULL
+ GROUP BY user_id
+ HAVING COUNT(*) > 1
+ )
+ AND id NOT IN (
+ SELECT id
+ FROM (
+ SELECT id,
+ ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY updated_at DESC, created_at DESC, id DESC) AS rn
+ FROM telegram_bindings
+ WHERE deleted_at IS NULL
+ ) AS ranked_bindings
+ WHERE rn = 1
+ )
+`).Error; err != nil {
+ return err
+ }
+ return tx.Exec(`
+CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_bindings_user_id_active
+ON telegram_bindings(user_id)
+WHERE deleted_at IS NULL
+`).Error
+ })
+}
diff --git a/internal/database/sqlite_migration.go b/internal/database/sqlite_migration.go
new file mode 100644
index 0000000..8ae6b14
--- /dev/null
+++ b/internal/database/sqlite_migration.go
@@ -0,0 +1,463 @@
+package database
+
+import (
+ "errors"
+ "fmt"
+ "os"
+ "path/filepath"
+ "reflect"
+ "sort"
+ "strings"
+ "time"
+
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+ "gorm.io/gorm/clause"
+ "gorm.io/gorm/logger"
+ "gorm.io/gorm/schema"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// MigrateSQLiteToCurrentIfNeeded copies an existing SQLite database into
+// PostgreSQL. Redis is not migrated because it is a rebuildable cache, not a
+// source of truth.
+const sqliteMigrationCompleteSettingKey = "database.sqlite_migration_complete"
+
+func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *zap.Logger) error {
+ if cfg == nil || target == nil || target.Dialector == nil || target.Dialector.Name() != "postgres" {
+ return nil
+ }
+ sqlitePath, err := sqliteMigrationSourcePath(cfg, log)
+ if err != nil {
+ return err
+ }
+ if sqlitePath == "" {
+ return nil
+ }
+ if complete, err := sqliteMigrationMarkedComplete(target); err != nil {
+ return err
+ } else if complete {
+ if log != nil {
+ log.Info("skip sqlite to postgres migration: already completed")
+ }
+ return nil
+ }
+
+ src, err := openSQLiteMigrationSource(cfg, sqlitePath)
+ if err != nil {
+ return fmt.Errorf("open sqlite migration source: %w", err)
+ }
+ sqlDB, err := src.DB()
+ if err == nil {
+ defer sqlDB.Close()
+ }
+
+ started := time.Now()
+ if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil {
+ return err
+ }
+ copied, err := copyModelTables(src, target, 500)
+ if err != nil {
+ return err
+ }
+ if err := markSQLiteMigrationComplete(target); err != nil {
+ return err
+ }
+ if log != nil {
+ log.Info("sqlite data migrated to postgres",
+ zap.String("source", sqlitePath),
+ zap.Int64("rows", copied),
+ zap.Duration("duration", time.Since(started)))
+ }
+ return nil
+}
+
+func openSQLiteMigrationSource(cfg *config.Config, sqlitePath string) (*gorm.DB, error) {
+ srcCfg := *cfg
+ srcCfg.Database.Type = "sqlite"
+ srcCfg.Database.DBPath = sqlitePath
+ return gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{
+ Logger: logger.Default.LogMode(logger.Silent),
+ })
+}
+
+func sqliteMigrationSourcePath(cfg *config.Config, log *zap.Logger) (string, error) {
+ configured := strings.TrimSpace(cfg.Database.DBPath)
+ if configured != "" {
+ exists, err := regularFileExists(configured)
+ if err != nil {
+ return "", fmt.Errorf("stat sqlite migration source: %w", err)
+ }
+ if exists {
+ return configured, nil
+ }
+ }
+
+ fallback := filepath.Join(strings.TrimSpace(cfg.App.DataDir), "mediastation.db")
+ if fallback == "" || sameCleanPath(configured, fallback) {
+ return "", nil
+ }
+ exists, err := regularFileExists(fallback)
+ if err != nil {
+ return "", fmt.Errorf("stat default sqlite migration source: %w", err)
+ }
+ if !exists {
+ return "", nil
+ }
+ if log != nil && configured != "" {
+ log.Warn("configured sqlite migration source not found; using data-dir default",
+ zap.String("configured", configured),
+ zap.String("fallback", fallback))
+ }
+ return fallback, nil
+}
+
+func regularFileExists(path string) (bool, error) {
+ if strings.TrimSpace(path) == "" {
+ return false, nil
+ }
+ info, err := os.Stat(path)
+ if err != nil {
+ if errors.Is(err, os.ErrNotExist) {
+ return false, nil
+ }
+ return false, err
+ }
+ return !info.IsDir(), nil
+}
+
+func sameCleanPath(a, b string) bool {
+ if a == "" || b == "" {
+ return false
+ }
+ return filepath.Clean(a) == filepath.Clean(b)
+}
+
+func resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target *gorm.DB, log *zap.Logger) error {
+ hasRows, err := sqliteSourceHasMigratableRows(src)
+ if err != nil {
+ return err
+ }
+ if !hasRows {
+ return nil
+ }
+ bootstrapOnly, err := targetLooksLikeBootstrapOnly(target)
+ if err != nil || !bootstrapOnly {
+ return err
+ }
+ for i := len(model.AllModels()) - 1; i >= 0; i-- {
+ m := model.AllModels()[i]
+ if !target.Migrator().HasTable(m) {
+ continue
+ }
+ if err := target.Session(&gorm.Session{AllowGlobalUpdate: true}).Unscoped().Delete(m).Error; err != nil {
+ return fmt.Errorf("clear bootstrap target table %T: %w", m, err)
+ }
+ }
+ if log != nil {
+ log.Warn("cleared bootstrap postgres rows before sqlite migration")
+ }
+ return nil
+}
+
+func sqliteSourceHasMigratableRows(src *gorm.DB) (bool, error) {
+ for _, table := range []string{"users", "libraries", "media", "settings"} {
+ exists, err := sqliteTableExists(src, table)
+ if err != nil {
+ return false, err
+ }
+ if !exists {
+ continue
+ }
+ var count int64
+ if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&count).Error; err != nil {
+ return false, fmt.Errorf("count sqlite table %s: %w", table, err)
+ }
+ if count > 0 {
+ return true, nil
+ }
+ }
+ return false, nil
+}
+
+func targetLooksLikeBootstrapOnly(target *gorm.DB) (bool, error) {
+ for _, m := range []any{
+ &model.Library{},
+ &model.Series{},
+ &model.Media{},
+ &model.PlaybackHistory{},
+ &model.Favorite{},
+ &model.Playlist{},
+ &model.PlaylistItem{},
+ &model.DownloadTask{},
+ &model.Subscription{},
+ } {
+ if !target.Migrator().HasTable(m) {
+ continue
+ }
+ var count int64
+ if err := target.Unscoped().Model(m).Count(&count).Error; err != nil {
+ return false, err
+ }
+ if count > 0 {
+ return false, nil
+ }
+ }
+
+ var userCount int64
+ if !target.Migrator().HasTable(&model.User{}) {
+ return true, nil
+ }
+ if err := target.Model(&model.User{}).Count(&userCount).Error; err != nil {
+ return false, err
+ }
+ if userCount == 0 {
+ return true, nil
+ }
+ if userCount != 1 {
+ return false, nil
+ }
+ var user model.User
+ if err := target.Unscoped().Where("username = ?", "admin").First(&user).Error; err != nil {
+ return false, nil
+ }
+ return user.Role == "admin", nil
+}
+
+func sqliteMigrationMarkedComplete(db *gorm.DB) (bool, error) {
+ var value string
+ err := db.Raw("SELECT value FROM "+quoteIdent("settings")+" WHERE "+quoteIdent("key")+" = ?", sqliteMigrationCompleteSettingKey).Scan(&value).Error
+ if err != nil {
+ return false, fmt.Errorf("check sqlite migration marker: %w", err)
+ }
+ return strings.EqualFold(strings.TrimSpace(value), "true"), nil
+}
+
+func markSQLiteMigrationComplete(db *gorm.DB) error {
+ now := time.Now()
+ if err := db.Clauses(clause.OnConflict{
+ Columns: []clause.Column{{Name: "key"}},
+ DoUpdates: clause.AssignmentColumns([]string{"value", "updated_at"}),
+ }).Create(&model.Setting{
+ Key: sqliteMigrationCompleteSettingKey,
+ Value: "true",
+ UpdatedAt: now,
+ }).Error; err != nil {
+ return fmt.Errorf("mark sqlite migration complete: %w", err)
+ }
+ return nil
+}
+
+func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
+ if batchSize <= 0 {
+ batchSize = 500
+ }
+ var copied int64
+ for _, m := range model.AllModels() {
+ table, err := modelTableName(src, m)
+ if err != nil {
+ return copied, err
+ }
+ primaryColumns, err := modelPrimaryColumns(src, m)
+ if err != nil {
+ return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
+ }
+ exists, err := sqliteTableExists(src, table)
+ if err != nil {
+ return copied, err
+ }
+ if !exists {
+ continue
+ }
+ var sourceCount int64
+ if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil {
+ return copied, fmt.Errorf("count sqlite table %s: %w", table, err)
+ }
+ if sourceCount == 0 {
+ continue
+ }
+ var targetCount int64
+ if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil {
+ return copied, fmt.Errorf("count target table %s: %w", table, err)
+ }
+ modelType := reflect.TypeOf(m)
+ if modelType.Kind() != reflect.Ptr {
+ return copied, fmt.Errorf("model %T is not a pointer", m)
+ }
+ sliceType := reflect.SliceOf(modelType.Elem())
+ slicePtr := reflect.New(sliceType)
+ if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
+ return copied, fmt.Errorf("read sqlite table %s: %w", table, err)
+ }
+ filtered := slicePtr.Elem()
+ if targetCount > 0 {
+ primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
+ if err != nil {
+ return copied, err
+ }
+ filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
+ }
+ if filtered.Len() == 0 {
+ continue
+ }
+ filteredPtr := reflect.New(filtered.Type())
+ filteredPtr.Elem().Set(filtered)
+ if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
+ return copied, fmt.Errorf("copy sqlite table %s: %w", table, err)
+ }
+ copied += int64(filtered.Len())
+ }
+ return copied, nil
+}
+
+func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) {
+ stmt := &gorm.Statement{DB: db}
+ if err := stmt.Parse(m); err != nil {
+ return nil, err
+ }
+ var cols []string
+ for _, field := range stmt.Schema.PrimaryFields {
+ cols = append(cols, field.DBName)
+ }
+ if len(cols) == 0 {
+ return nil, fmt.Errorf("no primary key columns")
+ }
+ return cols, nil
+}
+
+func targetPrimaryKeySet(target *gorm.DB, table string, primaryColumns []string) (map[string]struct{}, error) {
+ if len(primaryColumns) != 1 {
+ return nil, nil
+ }
+ var values []string
+ if err := target.Raw("SELECT " + quoteIdent(primaryColumns[0]) + " FROM " + quoteIdent(table)).Scan(&values).Error; err != nil {
+ return nil, fmt.Errorf("read target primary keys for table %s: %w", table, err)
+ }
+ set := make(map[string]struct{}, len(values))
+ for _, value := range values {
+ set[value] = struct{}{}
+ }
+ return set, nil
+}
+
+func filterRowsMissingInTarget(target *gorm.DB, table string, primaryColumns []string, rows reflect.Value, primaryKeySet map[string]struct{}) reflect.Value {
+ if rows.Kind() != reflect.Slice || rows.Len() == 0 || len(primaryColumns) == 0 {
+ return rows
+ }
+ out := reflect.MakeSlice(rows.Type(), 0, rows.Len())
+ for i := 0; i < rows.Len(); i++ {
+ row := rows.Index(i)
+ keys, ok := rowPrimaryKeys(row, primaryColumns)
+ if !ok {
+ out = reflect.Append(out, row)
+ continue
+ }
+ if primaryKeySet != nil {
+ if _, exists := primaryKeySet[fmt.Sprint(keys[primaryColumns[0]])]; !exists {
+ out = reflect.Append(out, row)
+ }
+ continue
+ }
+ if !targetHasPrimaryKey(target, table, keys) {
+ out = reflect.Append(out, row)
+ }
+ }
+ return out
+}
+
+func rowPrimaryKeys(row reflect.Value, primaryColumns []string) (map[string]any, bool) {
+ if row.Kind() == reflect.Pointer {
+ if row.IsNil() {
+ return nil, false
+ }
+ row = row.Elem()
+ }
+ if row.Kind() != reflect.Struct {
+ return nil, false
+ }
+ keys := make(map[string]any, len(primaryColumns))
+ for _, column := range primaryColumns {
+ value, ok := fieldByDBName(row, column)
+ if !ok || value.IsZero() {
+ return nil, false
+ }
+ keys[column] = value.Interface()
+ }
+ return keys, true
+}
+
+func fieldByDBName(row reflect.Value, column string) (reflect.Value, bool) {
+ rowType := row.Type()
+ for i := 0; i < row.NumField(); i++ {
+ fieldType := rowType.Field(i)
+ field := row.Field(i)
+ if fieldType.Anonymous {
+ if value, ok := fieldByDBName(field, column); ok {
+ return value, true
+ }
+ }
+ if columnNameForStructField(fieldType) == column {
+ if field.Kind() == reflect.Pointer && field.IsNil() {
+ return reflect.Value{}, false
+ }
+ return field, field.CanInterface()
+ }
+ }
+ return reflect.Value{}, false
+}
+
+func columnNameForStructField(field reflect.StructField) string {
+ if field.PkgPath != "" && !field.Anonymous {
+ return ""
+ }
+ tag := field.Tag.Get("gorm")
+ settings := schema.ParseTagSetting(tag, ";")
+ if column := settings["COLUMN"]; column != "" {
+ return column
+ }
+ return schema.NamingStrategy{}.ColumnName("", field.Name)
+}
+
+func targetHasPrimaryKey(target *gorm.DB, table string, keys map[string]any) bool {
+ where := make([]string, 0, len(keys))
+ args := make([]any, 0, len(keys))
+ for _, column := range sortedMapKeys(keys) {
+ where = append(where, quoteIdent(column)+" = ?")
+ args = append(args, keys[column])
+ }
+ var count int64
+ err := target.Raw("SELECT COUNT(1) FROM "+quoteIdent(table)+" WHERE "+strings.Join(where, " AND "), args...).Scan(&count).Error
+ return err == nil && count > 0
+}
+
+func sortedMapKeys(m map[string]any) []string {
+ keys := make([]string, 0, len(m))
+ for key := range m {
+ keys = append(keys, key)
+ }
+ sort.Strings(keys)
+ return keys
+}
+
+func sqliteTableExists(db *gorm.DB, table string) (bool, error) {
+ var count int64
+ if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
+ return false, fmt.Errorf("inspect sqlite table %s: %w", table, err)
+ }
+ return count > 0, nil
+}
+
+func modelTableName(db *gorm.DB, m any) (string, error) {
+ stmt := &gorm.Statement{DB: db}
+ if err := stmt.Parse(m); err != nil {
+ return "", err
+ }
+ return stmt.Schema.Table, nil
+}
+
+func quoteIdent(value string) string {
+ return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
+}
diff --git a/internal/database/sqlite_runtime.go b/internal/database/sqlite_runtime.go
new file mode 100644
index 0000000..f18247e
--- /dev/null
+++ b/internal/database/sqlite_runtime.go
@@ -0,0 +1,104 @@
+package database
+
+import (
+ "context"
+ "fmt"
+ "path/filepath"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+)
+
+func installSQLiteWriteGate(db *gorm.DB) {
+ if db == nil {
+ return
+ }
+ const lockedKey = "mediastation:sqlite_write_locked"
+ gate := newSQLiteWriteGate()
+ lock := func(tx *gorm.DB) {
+ ctx := context.Background()
+ if tx.Statement != nil && tx.Statement.Context != nil {
+ ctx = tx.Statement.Context
+ }
+ if err := gate.Lock(ctx); err != nil {
+ _ = tx.AddError(err)
+ return
+ }
+ tx.InstanceSet(lockedKey, struct{}{})
+ }
+ unlock := func(tx *gorm.DB) {
+ if _, ok := tx.InstanceGet(lockedKey); ok {
+ gate.Unlock()
+ }
+ }
+ _ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock)
+ _ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock)
+ _ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock)
+ _ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock)
+ _ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock)
+ _ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock)
+ _ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock)
+ _ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock)
+}
+
+// sqliteWriteGate serializes in-process SQLite writes while respecting the
+// statement context, so request cancellation can break out of a queued write.
+type sqliteWriteGate struct {
+ ch chan struct{}
+}
+
+func newSQLiteWriteGate() *sqliteWriteGate {
+ return &sqliteWriteGate{ch: make(chan struct{}, 1)}
+}
+
+func (g *sqliteWriteGate) Lock(ctx context.Context) error {
+ select {
+ case g.ch <- struct{}{}:
+ return nil
+ default:
+ }
+ if ctx == nil {
+ ctx = context.Background()
+ }
+ select {
+ case g.ch <- struct{}{}:
+ return nil
+ case <-ctx.Done():
+ return ctx.Err()
+ }
+}
+
+func (g *sqliteWriteGate) Unlock() {
+ select {
+ case <-g.ch:
+ default:
+ }
+}
+
+func buildSQLiteDSN(cfg *config.Config) string {
+ dbPath := cfg.Database.DBPath
+ if !filepath.IsAbs(dbPath) {
+ // keep as-is to respect user-provided relative paths.
+ dbPath = filepath.Clean(dbPath)
+ }
+ dsn := dbPath + "?_pragma=foreign_keys(1)"
+ if cfg.Database.WALMode {
+ dsn += "&_pragma=journal_mode(WAL)"
+ }
+ if cfg.Database.BusyTimeout > 0 {
+ dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout)
+ }
+ if cfg.Database.CacheSize != 0 {
+ dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize)
+ }
+ return dsn
+}
+
+func isSQLite(db *gorm.DB) bool {
+ return db != nil && db.Dialector != nil && db.Dialector.Name() == "sqlite"
+}
+
+func isPostgres(db *gorm.DB) bool {
+ return db != nil && db.Dialector != nil && db.Dialector.Name() == "postgres"
+}
diff --git a/internal/handler/admin.go b/internal/handler/admin.go
index dd6b47b..3b35390 100644
--- a/internal/handler/admin.go
+++ b/internal/handler/admin.go
@@ -24,6 +24,9 @@ func listUsersHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
+ if svc.Sessions != nil {
+ svc.Sessions.ApplyToUsers(c.Request.Context(), users)
+ }
c.JSON(http.StatusOK, users)
}
}
@@ -132,6 +135,10 @@ func deleteUserHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusForbidden, gin.H{"error": "default admin cannot be deleted"})
return
}
+ if svc.Sessions != nil && svc.Sessions.UserRecentlyActive(c.Request.Context(), c.Param("id"), service.RealtimeDeletionGuardWindow()) {
+ c.JSON(http.StatusConflict, gin.H{"error": "user has a recent realtime session; confirm the user is offline before deletion"})
+ return
+ }
if err := svc.Repo.User.Delete(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
diff --git a/internal/handler/admin_test.go b/internal/handler/admin_test.go
new file mode 100644
index 0000000..d3a984e
--- /dev/null
+++ b/internal/handler/admin_test.go
@@ -0,0 +1,49 @@
+package handler
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestDeleteUserRefusesRecentRealtimeSession(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.User{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ admin := model.User{Base: model.Base{ID: "admin"}, Username: "admin", PasswordHash: "x", Role: "admin", IsActive: true}
+ viewer := model.User{Base: model.Base{ID: "viewer"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
+ if err := repos.DB.Create(&[]model.User{admin, viewer}).Error; err != nil {
+ t.Fatal(err)
+ }
+ tracker := service.NewSessionTrackerService(zap.NewNop())
+ tracker.RecordLogin(t.Context(), viewer.ID, viewer.Username, "dev-1", "Apple TV", "Yamby", "10.0.0.8")
+ svc := &service.Container{Repo: repos, Sessions: tracker}
+ router := gin.New()
+ router.DELETE("/admin/users/:id", deleteUserHandler(svc))
+
+ req := httptest.NewRequest(http.MethodDelete, "/admin/users/viewer", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusConflict {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ if found, _ := repos.User.FindByID(t.Context(), viewer.ID); found == nil {
+ t.Fatal("recent realtime user should not be deleted")
+ }
+}
diff --git a/internal/handler/auth.go b/internal/handler/auth.go
index 658c7f7..41547c5 100644
--- a/internal/handler/auth.go
+++ b/internal/handler/auth.go
@@ -41,6 +41,12 @@ func loginHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
+ if svc.Sessions != nil {
+ svc.Sessions.RecordLogin(c.Request.Context(), resp.User.ID, resp.User.Username, "", "Web", "Web", c.ClientIP())
+ }
+ if resp.Tokens != nil {
+ setAccessTokenCookie(c, resp.Tokens.AccessToken, int(resp.Tokens.ExpiresIn))
+ }
c.JSON(http.StatusOK, gin.H{
"user": resp.User,
"tokens": resp.Tokens,
@@ -65,6 +71,9 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
+ if tokens != nil {
+ setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
+ }
c.JSON(http.StatusCreated, gin.H{
"user": u,
"tokens": tokens,
diff --git a/internal/handler/auth_cookie.go b/internal/handler/auth_cookie.go
new file mode 100644
index 0000000..0fe3c88
--- /dev/null
+++ b/internal/handler/auth_cookie.go
@@ -0,0 +1,55 @@
+package handler
+
+import (
+ "net/http"
+ "strings"
+ "time"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func setAccessTokenCookie(c *gin.Context, token string, maxAgeSeconds int) {
+ token = strings.TrimSpace(token)
+ if token == "" {
+ return
+ }
+ if maxAgeSeconds <= 0 {
+ maxAgeSeconds = int(service.AccessTokenDuration.Seconds())
+ }
+ writeAccessTokenCookie(c, token, maxAgeSeconds)
+}
+
+func clearAccessTokenCookie(c *gin.Context) {
+ writeAccessTokenCookie(c, "", -1)
+}
+
+func writeAccessTokenCookie(c *gin.Context, value string, maxAgeSeconds int) {
+ cookie := &http.Cookie{
+ Name: middleware.AccessTokenCookieName,
+ Value: value,
+ Path: middleware.AccessTokenCookiePath,
+ MaxAge: maxAgeSeconds,
+ HttpOnly: true,
+ SameSite: http.SameSiteLaxMode,
+ Secure: requestIsHTTPS(c),
+ }
+ if maxAgeSeconds > 0 {
+ cookie.Expires = time.Now().Add(time.Duration(maxAgeSeconds) * time.Second)
+ } else if maxAgeSeconds < 0 {
+ cookie.Expires = time.Unix(0, 0)
+ }
+ http.SetCookie(c.Writer, cookie)
+}
+
+func requestIsHTTPS(c *gin.Context) bool {
+ if c == nil || c.Request == nil {
+ return false
+ }
+ if c.Request.TLS != nil {
+ return true
+ }
+ return strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
+}
diff --git a/internal/handler/auth_cookie_test.go b/internal/handler/auth_cookie_test.go
new file mode 100644
index 0000000..3440bd4
--- /dev/null
+++ b/internal/handler/auth_cookie_test.go
@@ -0,0 +1,71 @@
+package handler
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+)
+
+func TestSetAccessTokenCookie(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "https://media.local/api/auth/login", nil)
+
+ setAccessTokenCookie(c, "access-token", 3600)
+
+ cookie := findResponseCookie(t, w, middleware.AccessTokenCookieName)
+ if cookie.Value != "access-token" {
+ t.Fatalf("cookie value = %q", cookie.Value)
+ }
+ if cookie.Path != middleware.AccessTokenCookiePath {
+ t.Fatalf("cookie path = %q, want %q", cookie.Path, middleware.AccessTokenCookiePath)
+ }
+ if cookie.MaxAge != 3600 {
+ t.Fatalf("cookie max age = %d, want 3600", cookie.MaxAge)
+ }
+ if !cookie.HttpOnly {
+ t.Fatal("cookie should be HttpOnly")
+ }
+ if !cookie.Secure {
+ t.Fatal("https request should set Secure cookie")
+ }
+ if cookie.SameSite != http.SameSiteLaxMode {
+ t.Fatalf("cookie SameSite = %v, want Lax", cookie.SameSite)
+ }
+}
+
+func TestClearAccessTokenCookie(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "http://127.0.0.1:8080/api/me/logout", nil)
+
+ clearAccessTokenCookie(c)
+
+ cookie := findResponseCookie(t, w, middleware.AccessTokenCookieName)
+ if cookie.MaxAge >= 0 {
+ t.Fatalf("clear cookie max age = %d, want negative", cookie.MaxAge)
+ }
+ if cookie.Path != middleware.AccessTokenCookiePath {
+ t.Fatalf("cookie path = %q, want %q", cookie.Path, middleware.AccessTokenCookiePath)
+ }
+ if cookie.Secure {
+ t.Fatal("plain http request should not set Secure cookie")
+ }
+}
+
+func findResponseCookie(t *testing.T, w *httptest.ResponseRecorder, name string) *http.Cookie {
+ t.Helper()
+ for _, cookie := range w.Result().Cookies() {
+ if cookie.Name == name {
+ return cookie
+ }
+ }
+ t.Fatalf("missing response cookie %q", name)
+ return nil
+}
diff --git a/internal/handler/auth_extra.go b/internal/handler/auth_extra.go
index ec70e5e..7d7195d 100644
--- a/internal/handler/auth_extra.go
+++ b/internal/handler/auth_extra.go
@@ -39,6 +39,7 @@ func refreshHandler(svc *service.Container) gin.HandlerFunc {
// so the Vue frontend's logout button gets a 200 instead of 404.
func logoutHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ clearAccessTokenCookie(c)
c.Status(http.StatusNoContent)
}
}
diff --git a/internal/handler/cloud.go b/internal/handler/cloud.go
index d005b5c..9156e89 100644
--- a/internal/handler/cloud.go
+++ b/internal/handler/cloud.go
@@ -3,18 +3,10 @@
package handler
import (
- "crypto/sha256"
- "encoding/hex"
- "io"
"net/http"
- "net/url"
- "path"
- "sort"
"strings"
- "time"
"github.com/gin-gonic/gin"
- "go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
@@ -25,6 +17,10 @@ import (
func cloudListHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider", "items": []any{}})
+ return
+ }
dir := c.Query("dir")
entries, err := svc.StorageCfg.CloudList(c.Request.Context(), typ, dir)
if err != nil {
@@ -35,10 +31,62 @@ func cloudListHandler(svc *service.Container) gin.HandlerFunc {
}
}
+func cloudMkdirHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ typ := c.Param("type")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
+ var in struct {
+ Dir string `json:"dir"`
+ Name string `json:"name" binding:"required"`
+ }
+ if err := c.ShouldBindJSON(&in); err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+ return
+ }
+ entry, err := svc.StorageCfg.CloudMkdir(c.Request.Context(), typ, in.Dir, in.Name)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"entry": entry})
+ }
+}
+
+func cloudRenameHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ typ := c.Param("type")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
+ var in struct {
+ Ref string `json:"ref" binding:"required"`
+ Name string `json:"name" binding:"required"`
+ }
+ if err := c.ShouldBindJSON(&in); err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+ return
+ }
+ entry, err := svc.StorageCfg.CloudRename(c.Request.Context(), typ, in.Ref, in.Name)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"entry": entry})
+ }
+}
+
// cloudImportHandler turns a cloud file into a playable 302-backed media item.
func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
var in struct {
Ref string `json:"ref" binding:"required"`
Name string `json:"name"`
@@ -63,6 +111,10 @@ func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
func cloudMountHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
var in struct {
Dir string `json:"dir"`
DirPath string `json:"dir_path"`
@@ -284,209 +336,3 @@ func cloud115QRPollHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusOK, st)
}
}
-
-// cloudPlayHandler resolves a cloud file to its direct link and either issues a
-// 302 redirect (true offload — host does not stream the bytes) or, when the
-// provider requires authenticated headers, reverse-proxies the response.
-func cloudPlayHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- typ := c.Param("type")
- ref := c.Query("ref")
- if ref == "" {
- c.JSON(http.StatusBadRequest, gin.H{"error": "ref required"})
- return
- }
- if !enforceScopedCloudPlaybackToken(c, svc, typ, ref) {
- return
- }
- serveCloudResolvedLink(svc, c, typ, ref)
- }
-}
-
-func serveCloudResolvedLink(svc *service.Container, c *gin.Context, typ, ref string) {
- if isCloudImageRef(ref) && svc != nil && svc.ImageProxy != nil {
- if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, typ+":"+ref) {
- return
- }
- }
- if svc == nil || svc.StorageCfg == nil {
- c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
- return
- }
- resolveStart := time.Now()
- link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
- resolveDur := time.Since(resolveStart)
- if err != nil {
- logCloudPlayback(svc, "cloud playback resolve failed",
- append(cloudPlaybackLogFields(typ, ref, nil, resolveDur), zap.Error(err))...)
- c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
- return
- }
- if isCloudImageRef(ref) && svc.ImageProxy != nil {
- if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, typ+":"+ref, link); err != nil {
- c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
- }
- return
- }
- if isCloudImageRef(ref) {
- c.Header("Cache-Control", "public, max-age=2592000, immutable")
- }
- if !link.Proxy {
- // Pure offload: send the client straight to the cloud CDN.
- setRedirectNoStoreHeaders(c)
- logCloudPlayback(svc, "cloud playback redirect",
- append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
- zap.String("mode", "redirect"),
- zap.Int("status", http.StatusFound),
- zap.String("method", c.Request.Method),
- zap.String("range", c.GetHeader("Range")),
- )...)
- c.Redirect(http.StatusFound, link.URL)
- return
- }
- // Proxy mode: the direct link needs auth headers the browser cannot
- // carry. Stream through with Range forwarding.
- method := c.Request.Method
- if method == "" {
- method = http.MethodGet
- }
- req, err := http.NewRequestWithContext(c.Request.Context(), method, link.URL, nil)
- if err != nil {
- c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
- return
- }
- for k, v := range link.Headers {
- req.Header.Set(k, v)
- }
- if rng := c.GetHeader("Range"); rng != "" {
- req.Header.Set("Range", rng)
- }
- if accept := c.GetHeader("Accept"); accept != "" {
- req.Header.Set("Accept", accept)
- }
- if c.GetHeader("Accept-Encoding") == "" {
- req.Header.Set("Accept-Encoding", "identity")
- }
- upstreamStart := time.Now()
- resp, err := http.DefaultClient.Do(req)
- upstreamHeaderDur := time.Since(upstreamStart)
- if err != nil {
- logCloudPlayback(svc, "cloud playback proxy upstream failed",
- append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
- zap.String("mode", "proxy"),
- zap.String("method", method),
- zap.String("range", c.GetHeader("Range")),
- zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
- zap.Error(err),
- )...)
- c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
- return
- }
- defer resp.Body.Close()
- for _, h := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Last-Modified"} {
- if v := resp.Header.Get(h); v != "" {
- c.Header(h, v)
- }
- }
- if c.Writer.Header().Get("Accept-Ranges") == "" {
- c.Header("Accept-Ranges", "bytes")
- }
- if resp.StatusCode >= 400 {
- c.Header("Cache-Control", "no-store")
- }
- c.Status(resp.StatusCode)
- var copied int64
- var copyErr error
- streamStart := time.Now()
- if c.Request.Method != http.MethodHead {
- copied, copyErr = io.Copy(c.Writer, resp.Body)
- }
- fields := append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
- zap.String("mode", "proxy"),
- zap.String("method", method),
- zap.String("range", c.GetHeader("Range")),
- zap.Int("status", resp.StatusCode),
- zap.String("content_range", resp.Header.Get("Content-Range")),
- zap.String("content_length", resp.Header.Get("Content-Length")),
- zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
- zap.Int64("stream_ms", durationMilliseconds(time.Since(streamStart))),
- zap.Int64("total_ms", durationMilliseconds(time.Since(resolveStart))),
- zap.Int64("bytes", copied),
- )
- if copyErr != nil {
- logCloudPlayback(svc, "cloud playback proxy copy failed", append(fields, zap.Error(copyErr))...)
- return
- }
- logCloudPlayback(svc, "cloud playback proxy finished", fields...)
-}
-
-func isCloudImageRef(ref string) bool {
- ref = strings.ToLower(strings.TrimSpace(ref))
- for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp"} {
- if strings.HasSuffix(ref, suffix) {
- return true
- }
- }
- return false
-}
-
-func logCloudPlayback(svc *service.Container, msg string, fields ...zap.Field) {
- if svc == nil || svc.Log == nil {
- return
- }
- svc.Log.Info(msg, fields...)
-}
-
-func cloudPlaybackLogFields(typ, ref string, link *cloud.DirectLink, resolveDur time.Duration) []zap.Field {
- refHash, refExt := cloudPlaybackRefFingerprint(ref)
- fields := []zap.Field{
- zap.String("provider", strings.TrimSpace(typ)),
- zap.String("ref_hash", refHash),
- zap.String("ref_ext", refExt),
- zap.Int64("resolve_ms", durationMilliseconds(resolveDur)),
- }
- if link != nil {
- fields = append(fields,
- zap.String("target_host", cloudPlaybackLinkHost(link.URL)),
- zap.Bool("headers_required", len(link.Headers) > 0),
- zap.Strings("header_names", cloudPlaybackHeaderNames(link.Headers)),
- )
- }
- return fields
-}
-
-func cloudPlaybackRefFingerprint(ref string) (string, string) {
- ref = strings.TrimSpace(ref)
- sum := sha256.Sum256([]byte(ref))
- ext := strings.ToLower(path.Ext(strings.Trim(strings.ReplaceAll(ref, "\\", "/"), "/")))
- return hex.EncodeToString(sum[:])[:12], ext
-}
-
-func cloudPlaybackLinkHost(raw string) string {
- u, err := url.Parse(strings.TrimSpace(raw))
- if err != nil || u.Host == "" {
- return ""
- }
- return u.Host
-}
-
-func cloudPlaybackHeaderNames(headers map[string]string) []string {
- if len(headers) == 0 {
- return nil
- }
- out := make([]string, 0, len(headers))
- for key := range headers {
- if key = strings.TrimSpace(key); key != "" {
- out = append(out, key)
- }
- }
- sort.Strings(out)
- return out
-}
-
-func durationMilliseconds(d time.Duration) int64 {
- if d <= 0 {
- return 0
- }
- return d.Milliseconds()
-}
diff --git a/internal/handler/cloud_playback.go b/internal/handler/cloud_playback.go
new file mode 100644
index 0000000..6080743
--- /dev/null
+++ b/internal/handler/cloud_playback.go
@@ -0,0 +1,289 @@
+package handler
+
+import (
+ "crypto/sha256"
+ "encoding/hex"
+ "io"
+ "net/http"
+ "net/url"
+ "path"
+ "sort"
+ "strings"
+ "time"
+
+ "github.com/gin-gonic/gin"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+type cloudPlaybackRequest struct {
+ svc *service.Container
+ c *gin.Context
+ typ string
+ ref string
+ link *cloud.DirectLink
+ resolveStart time.Time
+ resolveDur time.Duration
+}
+
+// cloudPlayHandler resolves a cloud file to its direct link and either issues a
+// 302 redirect (true offload — host does not stream the bytes) or, when the
+// provider requires authenticated headers, reverse-proxies the response.
+func cloudPlayHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ typ := c.Param("type")
+ ref := c.Query("ref")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
+ if ref == "" {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "ref required"})
+ return
+ }
+ if !enforceScopedCloudPlaybackToken(c, svc, typ, ref) {
+ return
+ }
+ serveCloudResolvedLink(svc, c, typ, ref)
+ }
+}
+
+func serveCloudResolvedLink(svc *service.Container, c *gin.Context, typ, ref string) {
+ if isCloudImageRef(ref) && svc != nil && svc.ImageProxy != nil {
+ if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, typ+":"+ref) {
+ return
+ }
+ }
+ if svc == nil || svc.StorageCfg == nil {
+ c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
+ return
+ }
+ resolveStart := time.Now()
+ link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
+ resolveDur := time.Since(resolveStart)
+ if err != nil {
+ logCloudPlayback(svc, "cloud playback resolve failed",
+ append(cloudPlaybackLogFields(typ, ref, nil, resolveDur), zap.Error(err))...)
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ return
+ }
+ if isCloudImageRef(ref) && svc.ImageProxy != nil {
+ if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, typ+":"+ref, link); err != nil {
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ }
+ return
+ }
+ if isCloudImageRef(ref) {
+ c.Header("Cache-Control", "public, max-age=2592000, immutable")
+ }
+ if !link.Proxy {
+ // Pure offload: send the client straight to the cloud CDN.
+ setRedirectNoStoreHeaders(c)
+ logCloudPlayback(svc, "cloud playback redirect",
+ append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
+ zap.String("mode", "redirect"),
+ zap.Int("status", http.StatusFound),
+ zap.String("method", c.Request.Method),
+ zap.String("range", c.GetHeader("Range")),
+ )...)
+ c.Redirect(http.StatusFound, link.URL)
+ return
+ }
+ proxyCloudResolvedLink(cloudPlaybackRequest{
+ svc: svc,
+ c: c,
+ typ: typ,
+ ref: ref,
+ link: link,
+ resolveStart: resolveStart,
+ resolveDur: resolveDur,
+ })
+}
+
+func proxyCloudResolvedLink(playback cloudPlaybackRequest) {
+ c := playback.c
+ clientMethod := playback.c.Request.Method
+ if clientMethod == "" {
+ clientMethod = http.MethodGet
+ }
+ upstreamMethod := clientMethod
+ if upstreamMethod == http.MethodHead {
+ upstreamMethod = http.MethodGet
+ }
+ req, err := http.NewRequestWithContext(c.Request.Context(), upstreamMethod, playback.link.URL, nil)
+ if err != nil {
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ return
+ }
+ for k, v := range playback.link.Headers {
+ req.Header.Set(k, v)
+ }
+ if rng := c.GetHeader("Range"); rng != "" {
+ req.Header.Set("Range", rng)
+ } else if clientMethod == http.MethodHead {
+ req.Header.Set("Range", "bytes=0-0")
+ }
+ if accept := c.GetHeader("Accept"); accept != "" {
+ req.Header.Set("Accept", accept)
+ }
+ if c.GetHeader("Accept-Encoding") == "" {
+ req.Header.Set("Accept-Encoding", "identity")
+ }
+ upstreamStart := time.Now()
+ resp, err := http.DefaultClient.Do(req)
+ upstreamHeaderDur := time.Since(upstreamStart)
+ if err != nil {
+ logCloudPlayback(playback.svc, "cloud playback proxy upstream failed",
+ append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
+ zap.String("mode", "proxy"),
+ zap.String("method", clientMethod),
+ zap.String("upstream_method", upstreamMethod),
+ zap.String("range", c.GetHeader("Range")),
+ zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
+ zap.Error(err),
+ )...)
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ return
+ }
+ defer resp.Body.Close()
+ for _, h := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Last-Modified"} {
+ if v := resp.Header.Get(h); v != "" {
+ c.Header(h, v)
+ }
+ }
+ if c.Writer.Header().Get("Accept-Ranges") == "" {
+ c.Header("Accept-Ranges", "bytes")
+ }
+ if resp.StatusCode >= 400 {
+ handleCloudProxyError(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
+ return
+ }
+ streamCloudProxyResponse(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
+}
+
+func handleCloudProxyError(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
+ c := playback.c
+ c.Header("Cache-Control", "no-store")
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
+ fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
+ zap.String("mode", "proxy"),
+ zap.String("method", clientMethod),
+ zap.String("upstream_method", upstreamMethod),
+ zap.String("range", c.GetHeader("Range")),
+ zap.String("upstream_range", req.Header.Get("Range")),
+ zap.Int("status", resp.StatusCode),
+ zap.String("content_range", resp.Header.Get("Content-Range")),
+ zap.String("content_length", resp.Header.Get("Content-Length")),
+ zap.String("upstream_error_body", strings.TrimSpace(string(body))),
+ zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
+ zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
+ )
+ logCloudPlayback(playback.svc, "cloud playback proxy upstream returned error", fields...)
+ c.Status(resp.StatusCode)
+ if clientMethod != http.MethodHead && len(body) > 0 {
+ _, _ = c.Writer.Write(body)
+ }
+}
+
+func streamCloudProxyResponse(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
+ c := playback.c
+ c.Status(resp.StatusCode)
+ var copied int64
+ var copyErr error
+ streamStart := time.Now()
+ if c.Request.Method != http.MethodHead {
+ copied, copyErr = io.Copy(c.Writer, resp.Body)
+ }
+ fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
+ zap.String("mode", "proxy"),
+ zap.String("method", clientMethod),
+ zap.String("upstream_method", upstreamMethod),
+ zap.String("range", c.GetHeader("Range")),
+ zap.String("upstream_range", req.Header.Get("Range")),
+ zap.Int("status", resp.StatusCode),
+ zap.String("content_range", resp.Header.Get("Content-Range")),
+ zap.String("content_length", resp.Header.Get("Content-Length")),
+ zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
+ zap.Int64("stream_ms", durationMilliseconds(time.Since(streamStart))),
+ zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
+ zap.Int64("bytes", copied),
+ )
+ if copyErr != nil {
+ logCloudPlayback(playback.svc, "cloud playback proxy copy failed", append(fields, zap.Error(copyErr))...)
+ return
+ }
+ logCloudPlayback(playback.svc, "cloud playback proxy finished", fields...)
+}
+
+func isCloudImageRef(ref string) bool {
+ ref = strings.ToLower(strings.TrimSpace(ref))
+ for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
+ if strings.HasSuffix(ref, suffix) {
+ return true
+ }
+ }
+ return false
+}
+
+func logCloudPlayback(svc *service.Container, msg string, fields ...zap.Field) {
+ if svc == nil || svc.Log == nil {
+ return
+ }
+ svc.Log.Info(msg, fields...)
+}
+
+func cloudPlaybackLogFields(typ, ref string, link *cloud.DirectLink, resolveDur time.Duration) []zap.Field {
+ refHash, refExt := cloudPlaybackRefFingerprint(ref)
+ fields := []zap.Field{
+ zap.String("provider", strings.TrimSpace(typ)),
+ zap.String("ref_hash", refHash),
+ zap.String("ref_ext", refExt),
+ zap.Int64("resolve_ms", durationMilliseconds(resolveDur)),
+ }
+ if link != nil {
+ fields = append(fields,
+ zap.String("target_host", cloudPlaybackLinkHost(link.URL)),
+ zap.Bool("headers_required", len(link.Headers) > 0),
+ zap.Strings("header_names", cloudPlaybackHeaderNames(link.Headers)),
+ )
+ }
+ return fields
+}
+
+func cloudPlaybackRefFingerprint(ref string) (string, string) {
+ ref = strings.TrimSpace(ref)
+ sum := sha256.Sum256([]byte(ref))
+ ext := strings.ToLower(path.Ext(strings.Trim(strings.ReplaceAll(ref, "\\", "/"), "/")))
+ return hex.EncodeToString(sum[:])[:12], ext
+}
+
+func cloudPlaybackLinkHost(raw string) string {
+ u, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil || u.Host == "" {
+ return ""
+ }
+ return u.Host
+}
+
+func cloudPlaybackHeaderNames(headers map[string]string) []string {
+ if len(headers) == 0 {
+ return nil
+ }
+ out := make([]string, 0, len(headers))
+ for key := range headers {
+ if key = strings.TrimSpace(key); key != "" {
+ out = append(out, key)
+ }
+ }
+ sort.Strings(out)
+ return out
+}
+
+func durationMilliseconds(d time.Duration) int64 {
+ if d <= 0 {
+ return 0
+ }
+ return d.Milliseconds()
+}
diff --git a/internal/handler/cloud_test.go b/internal/handler/cloud_test.go
index 0a88bc8..b60b1db 100644
--- a/internal/handler/cloud_test.go
+++ b/internal/handler/cloud_test.go
@@ -1,8 +1,17 @@
package handler
import (
+ "net/http"
+ "net/http/httptest"
"strings"
"testing"
+
+ "github.com/gin-gonic/gin"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
func TestCloudMountLibraryNameDefaultsToDirectoryBaseName(t *testing.T) {
@@ -48,3 +57,98 @@ func TestCloudPlaybackDiagnosticsDoNotExposeRawRefOrURL(t *testing.T) {
t.Fatalf("header names = %q", got)
}
}
+
+func TestAdminCloudHandlersRejectQuarkBrowsing(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ router.GET("/admin/cloud/:type/list", cloudListHandler(nil))
+
+ req := httptest.NewRequest(http.MethodGet, "/admin/cloud/quark/list?dir=0", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
+ }
+ if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
+ t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
+ }
+}
+
+func TestCloudPlayRejectsQuarkProvider(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ router.GET("/api/cloud/play/:type", cloudPlayHandler(nil))
+
+ req := httptest.NewRequest(http.MethodGet, "/api/cloud/play/quark?ref=file-1", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
+ }
+ if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
+ t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
+ }
+}
+
+func TestCloudArtworkProxyServesCachedImageWithoutCloudResolve(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("cached-cloud-poster"))
+ }))
+ defer upstream.Close()
+
+ imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
+ stableKey := "openlist:/Anime/JianLai/poster.jpg"
+ if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
+ t.Fatal(err)
+ }
+
+ router := gin.New()
+ router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
+
+ req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
+ }
+ if got := w.Body.String(); got != "cached-cloud-poster" {
+ t.Fatalf("body = %q, want cached poster", got)
+ }
+ if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
+ t.Fatalf("cache-control = %q, want long static cache", got)
+ }
+}
+
+func TestCloudArtworkProxyAcceptsCachedTBNImage(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("cached-tbn-poster"))
+ }))
+ defer upstream.Close()
+
+ imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
+ stableKey := "openlist:/Movies/Movie.tbn"
+ if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/Movie.tbn"}); err != nil {
+ t.Fatal(err)
+ }
+
+ router := gin.New()
+ router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
+
+ req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FMovies%2FMovie.tbn", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
+ }
+ if got := w.Body.String(); got != "cached-tbn-poster" {
+ t.Fatalf("body = %q, want cached tbn poster", got)
+ }
+}
diff --git a/internal/handler/discover.go b/internal/handler/discover.go
index fe98d36..89b83b1 100644
--- a/internal/handler/discover.go
+++ b/internal/handler/discover.go
@@ -26,6 +26,7 @@ func trendingHandler(svc *service.Container) gin.HandlerFunc {
if items == nil {
items = []service.Match{}
}
+ svc.Discover.WarmMatchArtwork(items)
c.JSON(http.StatusOK, gin.H{"items": items})
}
}
@@ -41,6 +42,7 @@ func popularHandler(svc *service.Container) gin.HandlerFunc {
if items == nil {
items = []service.Match{}
}
+ svc.Discover.WarmMatchArtwork(items)
c.JSON(http.StatusOK, gin.H{"items": items})
}
}
diff --git a/internal/handler/discover_extra.go b/internal/handler/discover_extra.go
index 521f0de..68e4e22 100644
--- a/internal/handler/discover_extra.go
+++ b/internal/handler/discover_extra.go
@@ -16,24 +16,37 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/service"
)
+type discoverSectionDef struct {
+ Key string
+ Label string
+ Provider string
+}
+
+var discoverSectionCatalog = []discoverSectionDef{
+ {Key: "tmdb_trending_day", Label: "TMDb 今日趋势", Provider: "tmdb"},
+ {Key: "tmdb_trending_week", Label: "TMDb 本周热门", Provider: "tmdb"},
+ {Key: "tmdb_popular_movie", Label: "TMDb 热门电影", Provider: "tmdb"},
+ {Key: "tmdb_popular_tv", Label: "TMDb 热门剧集", Provider: "tmdb"},
+ {Key: "tmdb_top_rated_movie", Label: "TMDb 高分电影", Provider: "tmdb"},
+ {Key: "douban_hot_movie", Label: "豆瓣热门电影", Provider: "douban"},
+ {Key: "douban_hot_tv", Label: "豆瓣热门剧集", Provider: "douban"},
+ {Key: "douban_top_movie", Label: "豆瓣高分电影", Provider: "douban"},
+ {Key: "bangumi_calendar", Label: "Bangumi 每日放送", Provider: "bangumi"},
+}
+
// discoverSectionsHandler returns the catalog of sections the UI can
// pick from. The names match the upstream Vue UI so existing settings
// keep working.
-func discoverSectionsHandler(_ *service.Container) gin.HandlerFunc {
+func discoverSectionsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "sections": []gin.H{
- {"key": "tmdb_trending_day", "label": "TMDb 今日趋势", "provider": "tmdb"},
- {"key": "tmdb_trending_week", "label": "TMDb 本周热门", "provider": "tmdb"},
- {"key": "tmdb_popular_movie", "label": "TMDb 热门电影", "provider": "tmdb"},
- {"key": "tmdb_popular_tv", "label": "TMDb 热门剧集", "provider": "tmdb"},
- {"key": "tmdb_top_rated_movie", "label": "TMDb 高分电影", "provider": "tmdb"},
- {"key": "douban_hot_movie", "label": "豆瓣热门电影", "provider": "douban"},
- {"key": "douban_hot_tv", "label": "豆瓣热门剧集", "provider": "douban"},
- {"key": "douban_top_movie", "label": "豆瓣高分电影", "provider": "douban"},
- {"key": "bangumi_calendar", "label": "Bangumi 每日放送", "provider": "bangumi"},
- },
- })
+ sections := make([]gin.H, 0, len(discoverSectionCatalog))
+ for _, section := range discoverSectionCatalog {
+ if !discoverProviderEnabled(c.Request.Context(), svc, section.Provider) {
+ continue
+ }
+ sections = append(sections, gin.H{"key": section.Key, "label": section.Label, "provider": section.Provider})
+ }
+ c.JSON(http.StatusOK, gin.H{"sections": sections})
}
}
@@ -45,19 +58,51 @@ func discoverFeedHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
keys := strings.Split(c.DefaultQuery("sections", "tmdb_trending_day,tmdb_popular_movie,douban_hot_movie,bangumi_calendar"), ",")
out := gin.H{}
+ artworkItems := []service.ExternalMediaResult{}
for _, raw := range keys {
k := strings.TrimSpace(raw)
+ if provider := discoverSectionProvider(k); provider != "" && !discoverProviderEnabled(c.Request.Context(), svc, provider) {
+ out[k] = []service.ExternalMediaResult{}
+ continue
+ }
items, err := discoverSectionItems(c.Request.Context(), svc, k)
if err != nil {
svc.Log.Debug("discover fetch failed")
items = nil
}
+ artworkItems = append(artworkItems, items...)
out[k] = items
}
+ svc.Discover.WarmExternalArtwork(artworkItems)
c.JSON(http.StatusOK, out)
}
}
+func discoverSectionProvider(key string) string {
+ for _, section := range discoverSectionCatalog {
+ if section.Key == key {
+ return section.Provider
+ }
+ }
+ switch key {
+ case "trending_day", "trending_week", "popular_movie", "popular_tv", "top_rated_movie", "upcoming_movie":
+ return "tmdb"
+ default:
+ return ""
+ }
+}
+
+func discoverProviderEnabled(ctx context.Context, svc *service.Container, provider string) bool {
+ if svc == nil || svc.APIConfig == nil || strings.TrimSpace(provider) == "" {
+ return true
+ }
+ cfg, err := svc.APIConfig.Get(ctx, provider)
+ if err != nil || cfg == nil {
+ return true
+ }
+ return cfg.Enabled
+}
+
func discoverSectionItems(ctx context.Context, svc *service.Container, k string) ([]service.ExternalMediaResult, error) {
switch k {
case "tmdb_trending_day", "tmdb_trending_week", "tmdb_popular_movie", "tmdb_popular_tv", "tmdb_top_rated_movie",
diff --git a/internal/handler/discover_extra_test.go b/internal/handler/discover_extra_test.go
new file mode 100644
index 0000000..352d74f
--- /dev/null
+++ b/internal/handler/discover_extra_test.go
@@ -0,0 +1,37 @@
+package handler
+
+import (
+ "testing"
+
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestDiscoverProviderEnabledHonorsAPIConfigToggle(t *testing.T) {
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ apiConfig := service.NewAPIConfigService(zap.NewNop(), repos, service.NewCryptoService("", zap.NewNop()))
+ enabled := false
+ if _, err := apiConfig.Update(t.Context(), "douban", service.APIConfigPatch{Enabled: &enabled}); err != nil {
+ t.Fatal(err)
+ }
+ svc := &service.Container{APIConfig: apiConfig}
+
+ if discoverProviderEnabled(t.Context(), svc, "douban") {
+ t.Fatal("disabled API config should disable discover provider")
+ }
+ if !discoverProviderEnabled(t.Context(), svc, "missing-provider") {
+ t.Fatal("missing API config should keep discover provider available")
+ }
+}
diff --git a/internal/handler/emby.go b/internal/handler/emby.go
index edc839d..85ae3dc 100644
--- a/internal/handler/emby.go
+++ b/internal/handler/emby.go
@@ -1,1832 +1,3 @@
-// Package handler — Emby/Jellyfin compatibility shim.
-//
-// 路由挂在 /emby/* 和根路径下双前缀。Infuse / Yamby / Hills /
-// Senplayer / Kodi 这类客户端会自动尝试 /System/Info 与 /emby/System/Info
-// 两种 URL,我们都接住。
+// Package handler provides the HTTP API, including Emby/Jellyfin compatibility
+// routes mounted both at /emby/* and at the root path for client discovery.
package handler
-
-import (
- "bytes"
- "context"
- "encoding/json"
- "errors"
- "io"
- "net/http"
- "net/url"
- "strconv"
- "strings"
- "sync"
- "time"
-
- "github.com/gin-gonic/gin"
- "github.com/gorilla/websocket"
-
- "github.com/ShukeBta/MediaStationGo/internal/middleware"
- "github.com/ShukeBta/MediaStationGo/internal/service"
-)
-
-// embyError 返回 Emby 风格的错误(顶层 Code/Message)。
-func embyError(c *gin.Context, status int, msg string) {
- c.JSON(status, gin.H{"Code": status, "Message": msg})
-}
-
-// embyUserID 从中间件中获取 user id。Emby auth middleware 写入 CtxUserID。
-func embyUserID(c *gin.Context) string {
- if uid, ok := c.Get(middleware.CtxUserID); ok {
- if s, ok := uid.(string); ok {
- return s
- }
- }
- return ""
-}
-
-const embyCompatSessionTTL = 30 * time.Minute
-
-type embyCompatSession struct {
- token string
- expiresAt time.Time
-}
-
-var embyCompatSessions = struct {
- sync.RWMutex
- items map[string]embyCompatSession
-}{items: map[string]embyCompatSession{}}
-
-func embyAuthRequiredWithSessionFallback(secret string) gin.HandlerFunc {
- required := middleware.EmbyAuthRequired(secret)
- return func(c *gin.Context) {
- if embyRequestToken(c) == "" {
- if token := embyCompatSessionToken(c); token != "" {
- c.Request.Header.Set("X-Emby-Token", token)
- }
- }
- required(c)
- }
-}
-
-func embyRememberCompatSession(c *gin.Context, token string) {
- token = strings.TrimSpace(token)
- if token == "" {
- return
- }
- keys := embyCompatSessionKeys(c)
- if len(keys) == 0 {
- return
- }
- expiresAt := time.Now().Add(embyCompatSessionTTL)
- embyCompatSessions.Lock()
- defer embyCompatSessions.Unlock()
- if len(embyCompatSessions.items) > 1000 {
- now := time.Now()
- for key, session := range embyCompatSessions.items {
- if now.After(session.expiresAt) {
- delete(embyCompatSessions.items, key)
- }
- }
- if len(embyCompatSessions.items) > 1000 {
- embyCompatSessions.items = map[string]embyCompatSession{}
- }
- }
- for _, key := range keys {
- embyCompatSessions.items[key] = embyCompatSession{token: token, expiresAt: expiresAt}
- }
-}
-
-func embyCompatSessionToken(c *gin.Context) string {
- keys := embyCompatSessionKeys(c)
- if len(keys) == 0 {
- return ""
- }
- now := time.Now()
- embyCompatSessions.RLock()
- defer embyCompatSessions.RUnlock()
- for _, key := range keys {
- session, ok := embyCompatSessions.items[key]
- if ok && now.Before(session.expiresAt) {
- return session.token
- }
- }
- return ""
-}
-
-func embyCompatSessionKeys(c *gin.Context) []string {
- if c == nil {
- return nil
- }
- ip := strings.TrimSpace(c.ClientIP())
- if ip == "" {
- return nil
- }
- keys := []string{}
- add := func(kind, value string) {
- value = strings.TrimSpace(value)
- if value != "" {
- keys = append(keys, ip+"\x00"+kind+"\x00"+value)
- }
- }
- add("device", firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"))
- add("ua", c.GetHeader("User-Agent"))
- return keys
-}
-
-func firstHeaderValue(c *gin.Context, names ...string) string {
- for _, name := range names {
- if value := strings.TrimSpace(c.GetHeader(name)); value != "" {
- return value
- }
- }
- return ""
-}
-
-type embyClientInfo struct {
- DeviceID string
- DeviceName string
- Client string
-}
-
-func embyClientInfoFromRequest(c *gin.Context) embyClientInfo {
- auth := parseMediaBrowserAuthorization(firstHeaderValue(c,
- "X-Emby-Authorization",
- "X-MediaBrowser-Authorization",
- "Authorization",
- ))
- info := embyClientInfo{
- DeviceID: firstNonEmptyHeaderString(
- firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"),
- auth["DeviceId"],
- auth["DeviceID"],
- ),
- DeviceName: firstNonEmptyHeaderString(
- firstHeaderValue(c, "X-Emby-Device-Name", "X-Emby-DeviceName", "X-MediaBrowser-Device-Name", "X-MediaBrowser-DeviceName"),
- auth["Device"],
- ),
- Client: firstNonEmptyHeaderString(
- firstHeaderValue(c, "X-Emby-Client", "X-MediaBrowser-Client"),
- auth["Client"],
- ),
- }
- ua := strings.TrimSpace(c.GetHeader("User-Agent"))
- if info.Client == "" {
- info.Client = embyClientFromUserAgent(ua)
- }
- if info.DeviceName == "" {
- info.DeviceName = embyDeviceFromUserAgent(ua)
- }
- return info
-}
-
-func parseMediaBrowserAuthorization(raw string) map[string]string {
- out := map[string]string{}
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return out
- }
- for _, prefix := range []string{"MediaBrowser ", "Emby "} {
- if strings.HasPrefix(raw, prefix) {
- raw = strings.TrimSpace(strings.TrimPrefix(raw, prefix))
- break
- }
- }
- for _, part := range strings.Split(raw, ",") {
- key, value, ok := strings.Cut(strings.TrimSpace(part), "=")
- if !ok {
- continue
- }
- key = strings.TrimSpace(key)
- value = strings.Trim(strings.TrimSpace(value), `"`)
- if key != "" && value != "" {
- out[key] = value
- }
- }
- return out
-}
-
-func firstNonEmptyHeaderString(values ...string) string {
- for _, value := range values {
- if strings.TrimSpace(value) != "" {
- return strings.TrimSpace(value)
- }
- }
- return ""
-}
-
-func embyClientFromUserAgent(ua string) string {
- ua = strings.TrimSpace(ua)
- lower := strings.ToLower(ua)
- switch {
- case strings.Contains(lower, "infuse"):
- return "Infuse"
- case strings.Contains(lower, "emby"):
- return "Emby"
- case strings.Contains(lower, "jellyfin"):
- return "Jellyfin"
- case strings.Contains(lower, "yamby"):
- return "Yamby"
- case strings.Contains(lower, "vidhub"):
- return "VidHub"
- case strings.Contains(lower, "hills"):
- return "Hills"
- default:
- return ua
- }
-}
-
-func embyDeviceFromUserAgent(ua string) string {
- lower := strings.ToLower(strings.TrimSpace(ua))
- switch {
- case strings.Contains(lower, "android"):
- return "Android"
- case strings.Contains(lower, "iphone"):
- return "iPhone"
- case strings.Contains(lower, "ipad"):
- return "iPad"
- case strings.Contains(lower, "ios"):
- return "iOS"
- case strings.Contains(lower, "windows"):
- return "Windows PC"
- case strings.Contains(lower, "macintosh") || strings.Contains(lower, "mac os"):
- return "Mac"
- case strings.Contains(lower, "linux"):
- return "Linux PC"
- case strings.Contains(lower, "appletv") || strings.Contains(lower, "apple tv"):
- return "Apple TV"
- default:
- return ""
- }
-}
-
-// ─── System ──────────────────────────────────────────────────────────────────
-
-func embySystemInfoHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfo()))
- }
-}
-
-func embySystemInfoPublicHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfoPublic()))
- }
-}
-
-func embyRequestBaseURL(c *gin.Context) string {
- proto := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto"))
- if proto == "" {
- if c.Request != nil && c.Request.TLS != nil {
- proto = "https"
- } else {
- proto = "http"
- }
- }
- if comma := strings.Index(proto, ","); comma >= 0 {
- proto = strings.TrimSpace(proto[:comma])
- }
-
- host := strings.TrimSpace(c.GetHeader("X-Forwarded-Host"))
- if host == "" && c.Request != nil {
- host = strings.TrimSpace(c.Request.Host)
- }
- if host == "" {
- return ""
- }
- return strings.TrimRight(proto+"://"+host, "/")
-}
-
-func embyWithRequestAddress(c *gin.Context, payload map[string]any) map[string]any {
- out := make(map[string]any, len(payload)+2)
- for key, value := range payload {
- out[key] = value
- }
- if address := embyRequestBaseURL(c); address != "" {
- out["LocalAddress"] = address
- out["WanAddress"] = address
- out["PublishedServerUrl"] = address
- }
- return out
-}
-
-func embySystemEndpointHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "IsLocal": true,
- "IsInNetwork": true,
- })
- }
-}
-
-func embyPingHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- // Emby/Jellyfin 期望 plain text "Emby Server"
- c.String(http.StatusOK, "Emby Server")
- }
-}
-
-func embyRootHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, embyPublicSystemInfoPayload(c, svc))
- }
-}
-
-func embyPublicSystemInfoPayload(c *gin.Context, svc *service.Container) map[string]any {
- if svc != nil && svc.Emby != nil {
- return embyWithRequestAddress(c, svc.Emby.SystemInfoPublic())
- }
- return embyWithRequestAddress(c, map[string]any{
- "Id": "mediastation-go-001",
- "ServerId": "mediastation-go-001",
- "ServerName": "MediaStationGo",
- "Version": "4.8.10.0",
- "ServerVersion": "4.8.10.0",
- "ProductName": "Emby Server",
- "OperatingSystem": "Windows",
- "SupportsHttps": false,
- "SupportsAutoDiscovery": true,
- "StartupWizardCompleted": true,
- })
-}
-
-// ─── Users / Auth ────────────────────────────────────────────────────────────
-
-type embyAuthByNameReq struct {
- Username string `json:"Username"`
- Pw string `json:"Pw"`
- Password string `json:"Password"`
- PasswordMd5 string `json:"PasswordMd5"`
- PasswordSha1 string `json:"PasswordSha1"`
-}
-
-func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
- req := embyAuthByNameReq{}
- if strings.Contains(strings.ToLower(c.GetHeader("Content-Type")), "json") {
- var body map[string]any
- if err := c.ShouldBindJSON(&body); err != nil && !errors.Is(err, io.EOF) {
- return req, err
- }
- fillEmbyAuthFromMap(&req, body)
- }
-
- if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
- _ = c.Request.ParseForm()
- if req.Username == "" {
- req.Username = firstFormValue(c, "Username", "username", "Name", "name")
- }
- if req.Pw == "" {
- req.Pw = firstFormValue(c, "Pw", "pw")
- }
- if req.Password == "" {
- req.Password = firstFormValue(c, "Password", "password")
- }
- if req.PasswordMd5 == "" {
- req.PasswordMd5 = firstFormValue(c, "PasswordMd5", "passwordMd5", "password_md5")
- }
- if req.PasswordSha1 == "" {
- req.PasswordSha1 = firstFormValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
- }
- }
-
- if req.Username == "" {
- req.Username = firstQueryValue(c, "Username", "username", "Name", "name")
- }
- if req.Pw == "" {
- req.Pw = firstQueryValue(c, "Pw", "pw")
- }
- if req.Password == "" {
- req.Password = firstQueryValue(c, "Password", "password")
- }
- if req.PasswordMd5 == "" {
- req.PasswordMd5 = firstQueryValue(c, "PasswordMd5", "passwordMd5", "password_md5")
- }
- if req.PasswordSha1 == "" {
- req.PasswordSha1 = firstQueryValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
- }
- if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
- fillEmbyAuthFromRawBody(c, &req)
- }
- return req, nil
-}
-
-func fillEmbyAuthFromMap(req *embyAuthByNameReq, body map[string]any) {
- if req.Username == "" {
- req.Username = firstStringFromMap(body, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
- }
- if req.Pw == "" {
- req.Pw = firstStringFromMap(body, "Pw", "pw", "PW")
- }
- if req.Password == "" {
- req.Password = firstStringFromMap(body, "Password", "password", "Pass", "pass", "Pwd", "pwd")
- }
- if req.PasswordMd5 == "" {
- req.PasswordMd5 = firstStringFromMap(body, "PasswordMd5", "passwordMd5", "password_md5")
- }
- if req.PasswordSha1 == "" {
- req.PasswordSha1 = firstStringFromMap(body, "PasswordSha1", "passwordSha1", "password_sha1")
- }
-}
-
-func fillEmbyAuthFromRawBody(c *gin.Context, req *embyAuthByNameReq) {
- if c.Request == nil || c.Request.Body == nil {
- return
- }
- raw, err := io.ReadAll(io.LimitReader(c.Request.Body, 1<<20))
- if err != nil {
- return
- }
- c.Request.Body = io.NopCloser(bytes.NewReader(raw))
- raw = bytes.TrimSpace(raw)
- if len(raw) == 0 {
- return
- }
- if bytes.HasPrefix(raw, []byte("{")) {
- var body map[string]any
- if err := json.Unmarshal(raw, &body); err == nil {
- fillEmbyAuthFromMap(req, body)
- }
- return
- }
- if values, err := url.ParseQuery(string(raw)); err == nil {
- fillEmbyAuthFromValues(req, values)
- }
-}
-
-func fillEmbyAuthFromValues(req *embyAuthByNameReq, values url.Values) {
- if req.Username == "" {
- req.Username = firstValue(values, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
- }
- if req.Pw == "" {
- req.Pw = firstValue(values, "Pw", "pw", "PW")
- }
- if req.Password == "" {
- req.Password = firstValue(values, "Password", "password", "Pass", "pass", "Pwd", "pwd")
- }
- if req.PasswordMd5 == "" {
- req.PasswordMd5 = firstValue(values, "PasswordMd5", "passwordMd5", "password_md5")
- }
- if req.PasswordSha1 == "" {
- req.PasswordSha1 = firstValue(values, "PasswordSha1", "passwordSha1", "password_sha1")
- }
-}
-
-func firstValue(values url.Values, keys ...string) string {
- for _, key := range keys {
- if value := strings.TrimSpace(values.Get(key)); value != "" {
- return value
- }
- }
- return ""
-}
-
-func firstStringFromMap(body map[string]any, keys ...string) string {
- if len(body) == 0 {
- return ""
- }
- for _, key := range keys {
- if value, ok := body[key]; ok {
- if s, ok := value.(string); ok {
- return strings.TrimSpace(s)
- }
- }
- }
- return ""
-}
-
-func firstFormValue(c *gin.Context, keys ...string) string {
- for _, key := range keys {
- if values, ok := c.Request.PostForm[key]; ok && len(values) > 0 {
- if value := strings.TrimSpace(values[0]); value != "" {
- return value
- }
- }
- }
- return ""
-}
-
-func firstQueryValue(c *gin.Context, keys ...string) string {
- for _, key := range keys {
- if value := strings.TrimSpace(c.Query(key)); value != "" {
- return value
- }
- }
- return ""
-}
-
-// embyAuthByNameHandler 处理 POST /Users/AuthenticateByName。
-//
-// 这是 Emby 客户端登录的唯一入口(Infuse / Yamby / Hills 等都走这里)。
-// 用户名+密码 → 调用我们已有的 AuthService.Login → 返回 AccessToken + User。
-func embyAuthByNameHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- req, err := parseEmbyAuthByNameReq(c)
- if err != nil {
- embyError(c, http.StatusBadRequest, "invalid body")
- return
- }
- password := req.Pw
- if password == "" {
- password = req.Password
- }
- if strings.TrimSpace(req.Username) == "" || password == "" {
- if req.PasswordMd5 != "" || req.PasswordSha1 != "" {
- embyError(c, http.StatusBadRequest, "plain password required")
- return
- }
- embyError(c, http.StatusBadRequest, "missing username or password")
- return
- }
- resp, err := svc.Auth.Login(c.Request.Context(), req.Username, password)
- if err != nil {
- embyError(c, http.StatusUnauthorized, err.Error())
- return
- }
- // 记录登录设备会话并执行防共享检测(登录客户端数 / 设备指纹)。
- clientInfo := embyClientInfoFromRequest(c)
- if svc.Device != nil {
- svc.Device.RecordLogin(c.Request.Context(), resp.User.ID,
- clientInfo.DeviceID,
- clientInfo.DeviceName,
- clientInfo.Client,
- c.ClientIP())
- }
- userPayload, _ := svc.Emby.FindUser(c.Request.Context(), resp.User.ID)
- // Emby/Jellyfin 客户端没有 refresh token 机制:它们把这里返回的
- // AccessToken 长期保存并反复使用。若返回 60 分钟的普通 access
- // token,客户端每小时就会掉登录、无法播放、媒体库无法刷新。因此
- // 签发长期令牌(IssueEmbyToken)匹配 Emby 持久化令牌语义。
- accessToken := resp.Tokens.AccessToken
- if longLived, err := svc.Auth.IssueEmbyToken(resp.User); err == nil && longLived != "" {
- accessToken = longLived
- }
- embyRememberCompatSession(c, accessToken)
- c.JSON(http.StatusOK, gin.H{
- "AccessToken": accessToken,
- "ServerId": "mediastation-go-001",
- "User": userPayload,
- "SessionInfo": gin.H{
- "Id": resp.User.ID,
- "UserId": resp.User.ID,
- "UserName": resp.User.Username,
- "Client": clientInfo.Client,
- "DeviceId": clientInfo.DeviceID,
- "DeviceName": clientInfo.DeviceName,
- },
- })
- }
-}
-
-func embyPublicUsersHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- // 公开用户列表(Emby Web 客户端登录页拉这个,列出可见用户)。
- users, err := svc.Emby.ListUsers(c.Request.Context())
- if err != nil {
- c.JSON(http.StatusOK, []any{})
- return
- }
- // 公开版本只暴露 Id + Name,不包含 Policy。
- out := make([]map[string]any, 0, len(users))
- for _, u := range users {
- out = append(out, map[string]any{
- "Id": u["Id"],
- "Name": u["Name"],
- "ServerId": u["ServerId"],
- "HasPassword": true,
- })
- }
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyListUsersHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- users, err := svc.Emby.ListUsers(c.Request.Context())
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- c.JSON(http.StatusOK, users)
- }
-}
-
-func embyMeHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := embyUserID(c)
- if uid == "" {
- embyError(c, http.StatusUnauthorized, "not authenticated")
- return
- }
- u, err := svc.Emby.FindUser(c.Request.Context(), uid)
- if err != nil || u == nil {
- embyError(c, http.StatusNotFound, "user not found")
- return
- }
- c.JSON(http.StatusOK, u)
- }
-}
-
-func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- u, err := svc.Emby.FindUser(c.Request.Context(), c.Param("userId"))
- if err == nil && u != nil {
- c.JSON(http.StatusOK, u)
- return
- }
- if authUID := embyUserID(c); authUID != "" && authUID != c.Param("userId") {
- u, err = svc.Emby.FindUser(c.Request.Context(), authUID)
- if err == nil && u != nil {
- c.JSON(http.StatusOK, u)
- return
- }
- }
- c.JSON(http.StatusOK, embyFallbackUser(c.Param("userId")))
- }
-}
-
-func embyFallbackUser(id string) gin.H {
- if strings.TrimSpace(id) == "" {
- id = "mediastation-user"
- }
- return gin.H{
- "Id": id,
- "Name": "MediaStationGo",
- "ServerId": "mediastation-go-001",
- "HasPassword": true,
- "HasConfiguredPassword": true,
- "HasConfiguredEasyPassword": false,
- "EnableAutoLogin": false,
- "Policy": gin.H{
- "IsAdministrator": true,
- "EnableContentDeletion": true,
- "EnableRemoteControlOfOtherUsers": true,
- "EnableSharedDeviceControl": true,
- "EnableRemoteAccess": true,
- "EnableAllDevices": true,
- "EnableAllChannels": true,
- "EnableAllFolders": true,
- },
- }
-}
-
-// ─── Views / MediaFolders ────────────────────────────────────────────────────
-
-func embyViewsHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = embyUserID(c)
- }
- out, err := svc.Emby.Views(c.Request.Context(), uid)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyVirtualFoldersHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Header("Cache-Control", "no-store")
- libs, err := svc.Repo.Library.List(c.Request.Context())
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
- uid := embyUserID(c)
- visibility := service.UserDefaultMediaVisibility(c.Request.Context(), svc.Repo, uid)
- out := make([]gin.H, 0, len(libs))
- for _, lib := range libs {
- if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
- continue
- }
- collectionType := "movies"
- switch lib.Type {
- case "tv", "anime", "variety":
- collectionType = "tvshows"
- case "music":
- collectionType = "music"
- }
- out = append(out, gin.H{
- "Name": lib.Name,
- "Locations": []string{lib.Path},
- "CollectionType": collectionType,
- "ItemId": lib.ID,
- "Id": lib.ID,
- "PrimaryImageItemId": lib.ID,
- "RefreshStatus": "Idle",
- "LibraryOptions": gin.H{},
- })
- }
- c.JSON(http.StatusOK, out)
- }
-}
-
-// ─── Items ───────────────────────────────────────────────────────────────────
-
-func parseEmbyItemsParams(c *gin.Context) service.ItemsParams {
- limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "50"))
- offset, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "StartIndex", "startIndex", "startindex"), "0"))
- uid := c.Param("userId")
- if uid == "" {
- uid = firstQueryValue(c, "UserId", "userId", "userid")
- }
- if uid == "" {
- uid = embyUserID(c)
- }
- splitOpt := func(s string) []string {
- if s == "" {
- return nil
- }
- parts := strings.Split(s, ",")
- out := make([]string, 0, len(parts))
- for _, p := range parts {
- p = strings.TrimSpace(p)
- if p != "" {
- out = append(out, p)
- }
- }
- return out
- }
- return service.ItemsParams{
- UserID: uid,
- ParentID: firstQueryValue(c, "ParentId", "parentId", "parentid"),
- IDs: splitOpt(firstQueryValue(c, "Ids", "ids")),
- SearchTerm: firstQueryValue(c, "SearchTerm", "searchTerm", "searchterm"),
- IncludeItemTypes: splitOpt(firstQueryValue(c, "IncludeItemTypes", "includeItemTypes", "includeitemtypes")),
- Filters: splitOpt(firstQueryValue(c, "Filters", "filters")),
- Recursive: strings.EqualFold(firstQueryValue(c, "Recursive", "recursive"), "true"),
- SortBy: firstQueryValue(c, "SortBy", "sortBy", "sortby"),
- SortOrder: firstQueryValue(c, "SortOrder", "sortOrder", "sortorder"),
- Limit: limit,
- StartIndex: offset,
- }
-}
-
-func embyFirstNonEmptyString(values ...string) string {
- for _, value := range values {
- if strings.TrimSpace(value) != "" {
- return strings.TrimSpace(value)
- }
- }
- return ""
-}
-
-func embyItemsHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- out, err := svc.Emby.Items(c.Request.Context(), parseEmbyItemsParams(c))
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyItemByIDHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- id := c.Param("id")
- uid := c.Param("userId")
- if uid == "" {
- uid = embyUserID(c)
- }
- out, err := svc.Emby.Item(c.Request.Context(), id, uid)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- if out == nil {
- embyError(c, http.StatusNotFound, "item not found")
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyUserItemByIDHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- switch strings.ToLower(c.Param("id")) {
- case "latest":
- embyLatestItemsHandler(svc)(c)
- case "resume":
- embyResumeItemsHandler(svc)(c)
- default:
- embyItemByIDHandler(svc)(c)
- }
- }
-}
-
-func embyLatestItemsHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = firstQueryValue(c, "UserId", "userId", "userid")
- }
- if uid == "" {
- uid = embyUserID(c)
- }
- limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
- out, err := svc.Emby.LatestItems(c.Request.Context(), uid, firstQueryValue(c, "ParentId", "parentId", "parentid"), limit)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyResumeItemsHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = firstQueryValue(c, "UserId", "userId", "userid")
- }
- if uid == "" {
- uid = embyUserID(c)
- }
- limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
- out, err := svc.Emby.ResumeItems(c.Request.Context(), uid, limit)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyItemsCountsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "MovieCount": 0,
- "SeriesCount": 0,
- "EpisodeCount": 0,
- "ItemCount": 0,
- })
- }
-}
-
-func embyDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "Id": c.Param("id"),
- "ViewType": "Poster",
- "SortBy": "SortName",
- "SortOrder": "Ascending",
- "IndexBy": "SortName",
- "RememberIndexing": false,
- "PrimaryImageHeight": 250,
- "PrimaryImageWidth": 250,
- "ScrollDirection": "Horizontal",
- "ShowSidebar": true,
- "CustomPrefs": gin.H{
- "homeexploresection": "1",
- "homesection0": "smalllibrarytiles",
- "homesection1": "resume",
- "homesection2": "latestmedia",
- "homesection3": "nextup",
- "homesection4": "none",
- "homesection5": "none",
- "homesection6": "none",
- "latestItems": "true",
- "landing-livetv": "false",
- },
- })
- }
-}
-
-func embySaveDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Status(http.StatusNoContent)
- }
-}
-
-// ─── Images ──────────────────────────────────────────────────────────────────
-
-var embyPlaceholderPNG = []byte{
- 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
- 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
- 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
- 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
- 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
- 0x54, 0x78, 0x9c, 0x63, 0x50, 0xd1, 0x30, 0xf8,
- 0x0f, 0x00, 0x02, 0x6c, 0x01, 0x7c, 0x30, 0xed,
- 0x6e, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45,
- 0x4e, 0x44, 0xae, 0x42, 0x60, 0x82,
-}
-
-// embyItemImageHandler 把 /Items/{id}/Images/Primary 等请求直接输出为图片。
-// Emby 客户端缓存图片 URL 时经常不会继续携带 token;如果重定向到受保护的
-// /api/img 会变成 401,所以这里复用 ImageProxy 但不再走 /api 路由。
-func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- ctx, cancel := context.WithTimeout(c.Request.Context(), 8*time.Second)
- defer cancel()
- req := c.Request.WithContext(ctx)
- id := c.Param("id")
- imgType := strings.ToLower(c.Param("type"))
- raw, err := svc.Emby.ImageURL(ctx, id, imgType)
- if err != nil || raw == "" {
- embyServePlaceholderImage(c)
- return
- }
- if typ, ref, ok := parseCloudPlayImageURL(raw); ok {
- c.Request = req
- serveCloudResolvedLink(svc, c, typ, ref)
- return
- }
- if svc.ImageProxy == nil {
- embyServePlaceholderImage(c)
- return
- }
- if err := svc.ImageProxy.Serve(ctx, c.Writer, req, raw); err != nil {
- embyServePlaceholderImage(c)
- }
- }
-}
-
-func embyServePlaceholderImage(c *gin.Context) {
- c.Header("Content-Type", "image/png")
- c.Header("Cache-Control", "public, max-age=3600")
- c.Header("Content-Length", strconv.Itoa(len(embyPlaceholderPNG)))
- if c.Request.Method == http.MethodHead {
- c.Status(http.StatusOK)
- return
- }
- c.Data(http.StatusOK, "image/png", embyPlaceholderPNG)
-}
-
-func parseCloudPlayImageURL(raw string) (string, string, bool) {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return "", "", false
- }
- u, err := url.Parse(raw)
- if err != nil {
- return "", "", false
- }
- path := strings.Trim(u.Path, "/")
- const prefix = "api/cloud/play/"
- if !strings.HasPrefix(path, prefix) {
- return "", "", false
- }
- typ := strings.TrimSpace(strings.TrimPrefix(path, prefix))
- ref := strings.TrimSpace(u.Query().Get("ref"))
- if typ == "" || ref == "" {
- return "", "", false
- }
- return typ, ref, true
-}
-
-func embyShowSeasonsHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- params := service.ItemsParams{
- UserID: firstQueryValue(c, "UserId", "userId"),
- ParentID: c.Param("id"),
- Limit: 500,
- }
- out, err := svc.Emby.Items(c.Request.Context(), params)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyShowEpisodesHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- parentID := firstQueryValue(c, "SeasonId", "seasonId")
- if parentID == "" {
- parentID = c.Param("id")
- }
- params := service.ItemsParams{
- UserID: firstQueryValue(c, "UserId", "userId"),
- ParentID: parentID,
- IncludeItemTypes: []string{"Episode"},
- Recursive: true,
- Limit: 500,
- }
- out, err := svc.Emby.Items(c.Request.Context(), params)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-// ─── Playback ────────────────────────────────────────────────────────────────
-
-func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = embyUserID(c)
- }
- out, err := svc.Emby.PlaybackInfo(c.Request.Context(), c.Param("id"), uid)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- if out == nil {
- embyError(c, http.StatusNotFound, "not found")
- return
- }
- embyAttachRequestTokenToMediaSources(c, out)
- c.JSON(http.StatusOK, out)
- }
-}
-
-func embyAttachRequestTokenToMediaSources(c *gin.Context, out any) {
- token := embyRequestToken(c)
- if token == "" || out == nil {
- return
- }
- embyAttachTokenToMediaSourcesValue(out, token)
-}
-
-func embyAttachTokenToMediaSourcesValue(value any, token string) {
- switch typed := value.(type) {
- case map[string]any:
- embyAttachTokenToMediaSourcesMap(typed, token)
- case gin.H:
- embyAttachTokenToMediaSourcesMap(map[string]any(typed), token)
- case []map[string]any:
- for _, item := range typed {
- embyAttachTokenToMediaSourcesMap(item, token)
- }
- case []any:
- for _, item := range typed {
- embyAttachTokenToMediaSourcesValue(item, token)
- }
- }
-}
-
-func embyAttachTokenToMediaSourcesMap(out map[string]any, token string) {
- if out == nil {
- return
- }
- if sources, ok := out["MediaSources"].([]map[string]any); ok {
- embyAttachTokenToMediaSources(sources, token)
- } else if sources, ok := out["MediaSources"].([]any); ok {
- for _, source := range sources {
- if sourceMap, ok := source.(map[string]any); ok {
- embyAttachTokenToMediaSources([]map[string]any{sourceMap}, token)
- }
- }
- }
- if items, ok := out["Items"]; ok {
- embyAttachTokenToMediaSourcesValue(items, token)
- }
-}
-
-func embyAttachTokenToMediaSources(sources []map[string]any, token string) {
- for _, source := range sources {
- for _, key := range []string{"DirectStreamUrl", "TranscodingUrl"} {
- raw, ok := source[key].(string)
- if !ok {
- continue
- }
- source[key] = embyAppendAPIKey(raw, token)
- }
- }
-}
-
-func embyRequestToken(c *gin.Context) string {
- if c == nil {
- return ""
- }
- for _, key := range []string{"api_key", "apiKey", "ApiKey", "token", "X-Emby-Token", "X-MediaBrowser-Token"} {
- if value := strings.TrimSpace(c.Query(key)); value != "" {
- return value
- }
- }
- for _, header := range []string{"X-Emby-Token", "X-MediaBrowser-Token"} {
- if value := strings.TrimSpace(c.GetHeader(header)); value != "" {
- return value
- }
- }
- for _, header := range []string{"Authorization", "X-Emby-Authorization", "X-MediaBrowser-Authorization"} {
- if token := embyTokenFromAuthHeader(c.GetHeader(header)); token != "" {
- return token
- }
- }
- return ""
-}
-
-func embyTokenFromAuthHeader(value string) string {
- value = strings.TrimSpace(value)
- if value == "" {
- return ""
- }
- for _, prefix := range []string{"Bearer ", "Emby "} {
- if strings.HasPrefix(value, prefix) {
- return strings.TrimSpace(strings.TrimPrefix(value, prefix))
- }
- }
- for _, part := range strings.Split(value, ",") {
- part = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(part), "MediaBrowser "))
- if !strings.HasPrefix(part, "Token=") {
- continue
- }
- token := strings.TrimSpace(strings.TrimPrefix(part, "Token="))
- return strings.Trim(token, `"`)
- }
- if strings.Contains(value, "Token=") {
- return ""
- }
- return value
-}
-
-func embyAppendAPIKey(raw, token string) string {
- raw = strings.TrimSpace(raw)
- token = strings.TrimSpace(token)
- if raw == "" || token == "" {
- return raw
- }
- if strings.HasPrefix(raw, "//") {
- return raw
- }
- u, err := url.Parse(raw)
- if err != nil || u.IsAbs() {
- return raw
- }
- q := u.Query()
- if q.Get("api_key") == "" && q.Get("apiKey") == "" && q.Get("token") == "" {
- q.Set("api_key", token)
- u.RawQuery = q.Encode()
- }
- return u.String()
-}
-
-// embyVideoStreamHandler 是 GET /Videos/{id}/stream 的入口,
-// 直接代理到我们的 /api/stream/{id}(同一个 ServeFile)。
-func embyVideoStreamHandler(svc *service.Container, cloudMode string) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := embyUserID(c)
- item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
- if err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- if item == nil {
- c.Status(http.StatusNotFound)
- return
- }
- if embyShouldRedirectVideoStreamToSTRM(c, svc, c.Param("id"), cloudMode) {
- target := "/api/stream/" + url.PathEscape(strings.TrimSpace(c.Param("id")))
- if token := embyPlaybackRedirectToken(c, svc); token != "" {
- target = embyAppendAPIKey(target, token)
- }
- setRedirectNoStoreHeaders(c)
- c.Redirect(http.StatusFound, absoluteRequestURL(c, target))
- return
- }
- // 直接调用 Stream service 写入 response。
- // 此前这里把所有错误一律吞成 404:云盘 Cookie 过期、直链解析失败、
- // STRM 播放被关闭……在第三方播放器上全部表现为「404 不存在」,
- // 无法排查。现在区分:行不存在→404;云盘播放不可用/上游故障→502+原因。
- err = svc.Stream.ServeFileWithCloudMode(c.Writer, c.Request, c.Param("id"), cloudMode)
- switch {
- case err == nil:
- case errors.Is(err, service.ErrMediaNotFound):
- c.Status(http.StatusNotFound)
- case errors.Is(err, service.ErrCloudPlaybackDisabled):
- if !c.Writer.Written() {
- c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
- }
- default:
- if !c.Writer.Written() {
- c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
- }
- }
- }
-}
-
-func embyPlaybackRedirectToken(c *gin.Context, svc *service.Container) string {
- if token := embyRequestToken(c); token != "" {
- return token
- }
- if c == nil || svc == nil || svc.Auth == nil || svc.Repo == nil || svc.Repo.User == nil {
- return ""
- }
- uid := embyUserID(c)
- if uid == "" {
- return ""
- }
- u, err := svc.Repo.User.FindByID(c.Request.Context(), uid)
- if err != nil || u == nil {
- return ""
- }
- token, err := svc.Auth.IssueEmbyToken(u)
- if err != nil {
- return ""
- }
- return token
-}
-
-func embyShouldRedirectVideoStreamToSTRM(c *gin.Context, svc *service.Container, mediaID, cloudMode string) bool {
- if c == nil || svc == nil || svc.Repo == nil || svc.Repo.Media == nil || cloudMode != service.CloudPlaybackModeRedirectProxy {
- return false
- }
- settings := service.CloudPlaybackSettings(c.Request.Context(), svc.Repo)
- if settings.PreferredMode != service.CloudPlaybackModeSTRM || !settings.STRMEnabled {
- return false
- }
- m, err := svc.Repo.Media.FindByID(c.Request.Context(), mediaID)
- if err != nil || m == nil {
- return false
- }
- return strings.TrimSpace(m.STRMURL) != ""
-}
-
-func embyVideoHLSPlaylistHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := embyUserID(c)
- item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
- if err != nil || item == nil || svc.Stream == nil {
- c.Status(http.StatusNotFound)
- return
- }
- err = svc.Stream.ServeHLSPlaylist(c.Writer, c.Request, c.Param("id"))
- if errors.Is(err, service.ErrTranscodeDisabled) {
- c.JSON(http.StatusConflict, gin.H{"error": "transcode disabled"})
- return
- }
- if errors.Is(err, service.ErrTranscodeBusy) {
- c.JSON(http.StatusTooManyRequests, gin.H{"error": "transcode busy"})
- return
- }
- if err != nil {
- c.Status(http.StatusNotFound)
- }
- }
-}
-
-func embyVideoHLSSegmentHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := embyUserID(c)
- item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
- if err != nil || item == nil || svc.Stream == nil {
- c.Status(http.StatusNotFound)
- return
- }
- if err := svc.Stream.ServeHLSSegment(c.Writer, c.Request, c.Param("id"), c.Param("seg")); err != nil {
- c.Status(http.StatusNotFound)
- }
- }
-}
-
-// ─── 播放进度 / 收藏 / 已看 ────────────────────────────────────────────────
-
-type embyPlayingReq struct {
- ItemId string `json:"ItemId"`
- PositionTicks int64 `json:"PositionTicks"`
- RunTimeTicks int64 `json:"RunTimeTicks"`
-}
-
-func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := embyUserID(c)
- if uid == "" {
- c.Status(http.StatusUnauthorized)
- return
- }
- var req embyPlayingReq
- _ = c.ShouldBindJSON(&req)
- // 兼容 query 形式(一些客户端在 /Sessions/Playing/* 用 query)
- if req.ItemId == "" {
- req.ItemId = c.Query("ItemId")
- }
- if req.PositionTicks == 0 {
- req.PositionTicks, _ = strconv.ParseInt(c.Query("PositionTicks"), 10, 64)
- }
- if req.RunTimeTicks == 0 {
- req.RunTimeTicks, _ = strconv.ParseInt(c.Query("RunTimeTicks"), 10, 64)
- }
- if req.ItemId == "" {
- c.Status(http.StatusOK) // Emby 期望 2xx;不是关键操作
- return
- }
- // 被「一键踢下线」的设备拒绝继续播放,直到重新登录。
- clientInfo := embyClientInfoFromRequest(c)
- if svc.Device != nil && svc.Device.IsDeviceKicked(c.Request.Context(), uid, clientInfo.DeviceID) {
- c.Status(http.StatusUnauthorized)
- return
- }
- _ = svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks)
- // 标记该设备正在播放并执行并发播放防共享检测。
- if svc.Device != nil {
- svc.Device.RecordPlayback(c.Request.Context(), uid,
- clientInfo.DeviceID,
- clientInfo.DeviceName,
- clientInfo.Client)
- }
- c.Status(http.StatusNoContent)
- }
-}
-
-func embyFavoriteHandler(svc *service.Container, fav bool) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = embyUserID(c)
- }
- mid := c.Param("itemId")
- if uid == "" || mid == "" {
- c.Status(http.StatusBadRequest)
- return
- }
- if err := svc.Emby.SetFavorite(c.Request.Context(), uid, mid, fav); err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- // Emby 期望返回 UserItemDataDto;最小可工作版本:echo Item 即可。
- out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
- if out != nil {
- c.JSON(http.StatusOK, out["UserData"])
- return
- }
- c.JSON(http.StatusOK, gin.H{"IsFavorite": fav})
- }
-}
-
-func embyMarkPlayedHandler(svc *service.Container, played bool) gin.HandlerFunc {
- return func(c *gin.Context) {
- uid := c.Param("userId")
- if uid == "" {
- uid = embyUserID(c)
- }
- mid := c.Param("itemId")
- if uid == "" || mid == "" {
- c.Status(http.StatusBadRequest)
- return
- }
- if err := svc.Emby.MarkPlayed(c.Request.Context(), uid, mid, played); err != nil {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
- if played && svc.Device != nil {
- clientInfo := embyClientInfoFromRequest(c)
- svc.Device.RecordPlayback(c.Request.Context(), uid, clientInfo.DeviceID, clientInfo.DeviceName, clientInfo.Client)
- }
- out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
- if out != nil {
- c.JSON(http.StatusOK, out["UserData"])
- return
- }
- c.JSON(http.StatusOK, gin.H{"Played": played})
- }
-}
-
-// ─── Sessions / Branding 占位 ────────────────────────────────────────────────
-
-func embySessionsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, []any{})
- }
-}
-
-func embyNoContentHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Status(http.StatusNoContent)
- }
-}
-
-func embyWebSocketHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- if !websocket.IsWebSocketUpgrade(c.Request) {
- c.Status(http.StatusNoContent)
- return
- }
- conn, err := wsUpgrader.Upgrade(c.Writer, c.Request, nil)
- if err != nil {
- return
- }
- defer conn.Close()
-
- done := make(chan struct{})
- go func() {
- defer close(done)
- for {
- if _, _, err := conn.NextReader(); err != nil {
- return
- }
- }
- }()
-
- ticker := time.NewTicker(30 * time.Second)
- defer ticker.Stop()
- for {
- select {
- case <-done:
- return
- case <-ticker.C:
- _ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
- if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
- return
- }
- }
- }
- }
-}
-
-func embyServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "EnableFolderView": true,
- "EnableGroupingIntoCollections": true,
- "EnableExternalContentInSuggestions": false,
- "ImageSavingConvention": "Compatible",
- })
- }
-}
-
-func embyPublicServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "IsStartupWizardCompleted": true,
- "EnableRemoteAccess": true,
- "EnableUPnP": false,
- "EnableHttps": false,
- "RequireHttps": false,
- "LocalNetworkSubnets": []string{},
- "LocalNetworkAddresses": []string{},
- "RemoteClientBitrateLimit": 0,
- })
- }
-}
-
-func embyStartupConfigurationHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "IsStartupWizardCompleted": true,
- "StartupWizardCompleted": true,
- "EnableRemoteAccess": true,
- "UICulture": "zh-CN",
- "MetadataCountryCode": "CN",
- "PreferredMetadataLanguage": "zh-CN",
- })
- }
-}
-
-func embyQuickConnectEnabledHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, false)
- }
-}
-
-func embyEmptyItemsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{"Items": []any{}, "TotalRecordCount": 0})
- }
-}
-
-func embyEmptyArrayHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, []any{})
- }
-}
-
-func embyCustomCSSJSScriptsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Data(http.StatusOK, "application/javascript; charset=utf-8", nil)
- }
-}
-
-func embyLocalizationCulturesHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, []gin.H{
- {
- "DisplayName": "简体中文",
- "Name": "zh-CN",
- "ThreeLetterISOLanguageName": "zho",
- "TwoLetterISOLanguageName": "zh",
- "ThreeLetterISOLanguageNames": []string{"zho", "chi"},
- "IsRightToLeft": false,
- },
- {
- "DisplayName": "English",
- "Name": "en-US",
- "ThreeLetterISOLanguageName": "eng",
- "TwoLetterISOLanguageName": "en",
- "ThreeLetterISOLanguageNames": []string{"eng"},
- "IsRightToLeft": false,
- },
- })
- }
-}
-
-func embyThemeMediaHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- empty := gin.H{"Items": []any{}, "TotalRecordCount": 0}
- c.JSON(http.StatusOK, gin.H{
- "ThemeVideosResult": empty,
- "ThemeSongsResult": empty,
- "SoundtrackSongsResult": empty,
- })
- }
-}
-
-func embyServerDomainsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, []any{})
- }
-}
-
-func embyDanmuRawHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Data(http.StatusOK, "text/plain; charset=utf-8", nil)
- }
-}
-
-func embyBrandingConfigHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, gin.H{
- "LoginDisclaimer": "",
- "CustomCss": "",
- "SplashscreenEnabled": false,
- })
- }
-}
-
-func embyBrandingCSSHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.Data(http.StatusOK, "text/css; charset=utf-8", []byte(""))
- }
-}
-
-func embyLocalizationOptionsHandler(_ *service.Container) gin.HandlerFunc {
- return func(c *gin.Context) {
- c.JSON(http.StatusOK, []map[string]any{
- {"Name": "简体中文", "Value": "zh-CN"},
- {"Name": "English", "Value": "en-US"},
- })
- }
-}
-
-// registerEmbyRoutes 在 r 上挂双前缀("" + "/emby")的 Emby 兼容路由。
-func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) {
- for _, prefix := range []string{"/emby", ""} {
- grp := r.Group(prefix)
- grp.Use(func(c *gin.Context) {
- c.Header("Cache-Control", "no-store")
- c.Header("Pragma", "no-cache")
- c.Header("Expires", "0")
- c.Next()
- })
-
- if prefix == "/emby" {
- grp.GET("", embyRootHandler(svc))
- grp.HEAD("", embyRootHandler(svc))
- grp.GET("/", embyRootHandler(svc))
- grp.HEAD("/", embyRootHandler(svc))
- }
-
- // 公开端点
- for _, path := range []string{"/System/Info/Public", "/system/info/public"} {
- grp.GET(path, embySystemInfoPublicHandler(svc))
- grp.HEAD(path, embySystemInfoPublicHandler(svc))
- }
- for _, path := range []string{"/System/Info", "/system/info"} {
- grp.GET(path, embySystemInfoHandler(svc))
- grp.HEAD(path, embySystemInfoHandler(svc))
- }
- for _, path := range []string{"/System/Endpoint", "/system/endpoint"} {
- grp.GET(path, embySystemEndpointHandler(svc))
- }
- for _, path := range []string{"/System/Ext/ServerDomains", "/system/ext/serverdomains"} {
- grp.GET(path, embyServerDomainsHandler(svc))
- grp.HEAD(path, embyServerDomainsHandler(svc))
- }
- for _, path := range []string{"/System/Configuration/Public", "/system/configuration/public"} {
- grp.GET(path, embyPublicServerConfigurationHandler(svc))
- grp.HEAD(path, embyPublicServerConfigurationHandler(svc))
- }
- for _, path := range []string{"/Startup/Configuration", "/startup/configuration"} {
- grp.GET(path, embyStartupConfigurationHandler(svc))
- grp.HEAD(path, embyStartupConfigurationHandler(svc))
- }
- for _, path := range []string{"/Startup/Complete", "/startup/complete"} {
- grp.POST(path, embyNoContentHandler(svc))
- }
- for _, path := range []string{"/QuickConnect/Enabled", "/quickconnect/enabled"} {
- grp.GET(path, embyQuickConnectEnabledHandler(svc))
- grp.HEAD(path, embyQuickConnectEnabledHandler(svc))
- }
- for _, path := range []string{"/System/Ping", "/system/ping"} {
- grp.GET(path, embyPingHandler(svc))
- grp.HEAD(path, embyPingHandler(svc))
- grp.POST(path, embyPingHandler(svc))
- }
- for _, path := range []string{"/Sessions/Capabilities", "/Sessions/Capabilities/Full", "/sessions/capabilities", "/sessions/capabilities/full"} {
- grp.POST(path, embyNoContentHandler(svc))
- }
- // 30/min per IP: many Emby clients sit behind a single NAT/reverse-proxy
- // IP, so a low limit would throttle legitimate logins into 429s.
- embyLoginLimiter := middleware.NewRateLimiter(30, 1*time.Minute)
- for _, path := range []string{"/Users/AuthenticateByName", "/Users/authenticatebyname", "/users/AuthenticateByName", "/users/authenticatebyname"} {
- grp.POST(path, middleware.RateLimit(embyLoginLimiter), embyAuthByNameHandler(svc))
- }
- for _, path := range []string{"/Users/Public", "/users/public"} {
- grp.GET(path, embyPublicUsersHandler(svc))
- }
- for _, path := range []string{"/Branding/Configuration", "/branding/configuration"} {
- grp.GET(path, embyBrandingConfigHandler(svc))
- }
- for _, path := range []string{"/Branding/Css", "/branding/css"} {
- grp.GET(path, embyBrandingCSSHandler(svc))
- grp.HEAD(path, embyBrandingCSSHandler(svc))
- }
- for _, path := range []string{"/Localization/Options", "/localization/options"} {
- grp.GET(path, embyLocalizationOptionsHandler(svc))
- }
- for _, path := range []string{"/Localization/Cultures", "/Localization/cultures", "/localization/cultures"} {
- grp.GET(path, embyLocalizationCulturesHandler(svc))
- }
- for _, path := range []string{"/CustomCssJS/Scripts", "/customcssjs/scripts"} {
- grp.GET(path, embyCustomCSSJSScriptsHandler(svc))
- grp.HEAD(path, embyCustomCSSJSScriptsHandler(svc))
- }
- for _, path := range []string{"/embywebsocket", "/EmbyWebSocket"} {
- grp.GET(path, embyWebSocketHandler(svc))
- grp.HEAD(path, embyNoContentHandler(svc))
- }
- for _, path := range []string{"/Sessions/Logout", "/sessions/logout"} {
- grp.POST(path, embyNoContentHandler(svc))
- }
- grp.GET("/DisplayPreferences/:id", embyDisplayPreferencesHandler(svc))
- grp.POST("/DisplayPreferences/:id", embySaveDisplayPreferencesHandler(svc))
- grp.GET("/displaypreferences/:id", embyDisplayPreferencesHandler(svc))
- grp.POST("/displaypreferences/:id", embySaveDisplayPreferencesHandler(svc))
-
- // 图片公开(Infuse 缓存 URL 时会丢 token)
- grp.GET("/Items/:id/Images/:type", embyItemImageHandler(svc))
- grp.GET("/Items/:id/Images/:type/:index", embyItemImageHandler(svc))
- grp.HEAD("/Items/:id/Images/:type", embyItemImageHandler(svc))
- grp.GET("/items/:id/images/:type", embyItemImageHandler(svc))
- grp.GET("/items/:id/images/:type/:index", embyItemImageHandler(svc))
- grp.HEAD("/items/:id/images/:type", embyItemImageHandler(svc))
-
- // 鉴权后端点
- auth := grp.Group("", embyAuthRequiredWithSessionFallback(jwtSecret), activeEmbyUserRequired(svc))
- auth.GET("/Users/Me", embyMeHandler(svc))
- auth.GET("/Users", embyListUsersHandler(svc))
- auth.GET("/Users/:userId", embyGetUserByIDHandler(svc))
- auth.GET("/Users/:userId/Views", embyViewsHandler(svc))
- auth.GET("/Library/MediaFolders", embyViewsHandler(svc))
- auth.GET("/Library/VirtualFolders", embyVirtualFoldersHandler(svc))
- auth.GET("/Library/SelectableMediaFolders", embyVirtualFoldersHandler(svc))
-
- auth.GET("/Items", embyItemsHandler(svc))
- auth.GET("/Users/:userId/Items", embyItemsHandler(svc))
- auth.GET("/Items/Counts", embyItemsCountsHandler(svc))
- auth.GET("/Users/:userId/Items/Counts", embyItemsCountsHandler(svc))
- auth.GET("/Items/Latest", embyLatestItemsHandler(svc))
- auth.GET("/Items/Resume", embyResumeItemsHandler(svc))
- auth.GET("/Items/:id", embyItemByIDHandler(svc))
- auth.GET("/Users/:userId/Items/:id", embyUserItemByIDHandler(svc))
- auth.GET("/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
- auth.GET("/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
- auth.GET("/Users/:userId/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
- auth.GET("/Users/:userId/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
- auth.GET("/Shows/NextUp", embyEmptyItemsHandler(svc))
- auth.GET("/Users/:userId/Shows/NextUp", embyEmptyItemsHandler(svc))
- auth.GET("/MediaSegments/:id", embyEmptyItemsHandler(svc))
- auth.GET("/Artists", embyEmptyItemsHandler(svc))
- auth.GET("/Persons", embyEmptyItemsHandler(svc))
- auth.GET("/Genres", embyEmptyItemsHandler(svc))
- auth.GET("/Shows/Upcoming", embyEmptyItemsHandler(svc))
- auth.GET("/Users/:userId/Shows/Upcoming", embyEmptyItemsHandler(svc))
- auth.GET("/Items/:id/Similar", embyEmptyItemsHandler(svc))
- auth.GET("/Items/:id/ThumbnailSet", embyEmptyItemsHandler(svc))
- auth.GET("/Items/:id/ThemeMedia", embyThemeMediaHandler(svc))
- auth.GET("/Users/:userId/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
- auth.GET("/Users/:userId/Items/:id/Intros", embyEmptyItemsHandler(svc))
- auth.GET("/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
- auth.GET("/Items/:id/Intros", embyEmptyItemsHandler(svc))
- auth.GET("/api/danmu/:id/raw", embyDanmuRawHandler(svc))
-
- auth.GET("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
- auth.POST("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
- auth.GET("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
- auth.POST("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
-
- auth.GET("/Videos/:id/stream", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/Videos/:id/stream", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/Videos/:id/stream.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/Videos/:id/stream.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/Videos/:id/original", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/Videos/:id/original", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/Videos/:id/original.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/Videos/:id/original.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- if prefix == "/emby" {
- auth.GET("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
- auth.HEAD("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
- }
- auth.GET("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.HEAD("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.GET("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.HEAD("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.GET("/Videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
-
- auth.POST("/Sessions/Playing", embyPlayingProgressHandler(svc))
- auth.POST("/Sessions/Playing/Progress", embyPlayingProgressHandler(svc))
- auth.POST("/Sessions/Playing/Stopped", embyPlayingProgressHandler(svc))
-
- auth.POST("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, true))
- auth.DELETE("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, false))
- auth.POST("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, true))
- auth.DELETE("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, false))
-
- auth.GET("/Sessions", embySessionsHandler(svc))
- auth.GET("/System/Configuration", embyServerConfigurationHandler(svc))
- auth.GET("/System/WakeOnLanInfo", embyEmptyArrayHandler(svc))
- auth.GET("/ScheduledTasks", embyEmptyArrayHandler(svc))
- auth.GET("/LiveTv/Recordings", embyEmptyItemsHandler(svc))
- auth.GET("/System/ActivityLog/Entries", embyEmptyItemsHandler(svc))
- auth.GET("/Web/ConfigurationPages", embyEmptyArrayHandler(svc))
- auth.POST("/Users/:userId/Configuration", embyNoContentHandler(svc))
-
- registerLowercaseEmbyAuthRoutes(auth, svc)
- }
-}
-
-func registerLowercaseEmbyAuthRoutes(auth *gin.RouterGroup, svc *service.Container) {
- auth.GET("/users/me", embyMeHandler(svc))
- auth.GET("/users", embyListUsersHandler(svc))
- auth.GET("/users/:userId", embyGetUserByIDHandler(svc))
- auth.GET("/users/:userId/views", embyViewsHandler(svc))
- auth.GET("/library/mediafolders", embyViewsHandler(svc))
- auth.GET("/library/virtualfolders", embyVirtualFoldersHandler(svc))
- auth.GET("/library/selectablemediafolders", embyVirtualFoldersHandler(svc))
-
- auth.GET("/items", embyItemsHandler(svc))
- auth.GET("/users/:userId/items", embyItemsHandler(svc))
- auth.GET("/items/counts", embyItemsCountsHandler(svc))
- auth.GET("/users/:userId/items/counts", embyItemsCountsHandler(svc))
- auth.GET("/items/latest", embyLatestItemsHandler(svc))
- auth.GET("/items/resume", embyResumeItemsHandler(svc))
- auth.GET("/items/:id", embyItemByIDHandler(svc))
- auth.GET("/users/:userId/items/:id", embyUserItemByIDHandler(svc))
- auth.GET("/shows/:id/seasons", embyShowSeasonsHandler(svc))
- auth.GET("/shows/:id/episodes", embyShowEpisodesHandler(svc))
- auth.GET("/users/:userId/shows/:id/seasons", embyShowSeasonsHandler(svc))
- auth.GET("/users/:userId/shows/:id/episodes", embyShowEpisodesHandler(svc))
- auth.GET("/shows/nextup", embyEmptyItemsHandler(svc))
- auth.GET("/users/:userId/shows/nextup", embyEmptyItemsHandler(svc))
- auth.GET("/mediasegments/:id", embyEmptyItemsHandler(svc))
- auth.GET("/artists", embyEmptyItemsHandler(svc))
- auth.GET("/persons", embyEmptyItemsHandler(svc))
- auth.GET("/genres", embyEmptyItemsHandler(svc))
- auth.GET("/shows/upcoming", embyEmptyItemsHandler(svc))
- auth.GET("/users/:userId/shows/upcoming", embyEmptyItemsHandler(svc))
- auth.GET("/items/:id/similar", embyEmptyItemsHandler(svc))
- auth.GET("/items/:id/thumbnailset", embyEmptyItemsHandler(svc))
- auth.GET("/items/:id/thememedia", embyThemeMediaHandler(svc))
- auth.GET("/users/:userId/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
- auth.GET("/users/:userId/items/:id/intros", embyEmptyItemsHandler(svc))
- auth.GET("/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
- auth.GET("/items/:id/intros", embyEmptyItemsHandler(svc))
-
- auth.GET("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
- auth.POST("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
- auth.GET("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
- auth.POST("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
-
- auth.GET("/videos/:id/stream", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/videos/:id/stream", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/videos/:id/stream.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/videos/:id/stream.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/videos/:id/original", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/videos/:id/original", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/videos/:id/original.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.HEAD("/videos/:id/original.:container", embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy))
- auth.GET("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.HEAD("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.GET("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.HEAD("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
- auth.GET("/videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
-
- auth.POST("/sessions/playing", embyPlayingProgressHandler(svc))
- auth.POST("/sessions/playing/progress", embyPlayingProgressHandler(svc))
- auth.POST("/sessions/playing/stopped", embyPlayingProgressHandler(svc))
-
- auth.POST("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, true))
- auth.DELETE("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, false))
- auth.POST("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, true))
- auth.DELETE("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, false))
-
- auth.GET("/sessions", embySessionsHandler(svc))
- auth.GET("/system/configuration", embyServerConfigurationHandler(svc))
- auth.GET("/system/wakeonlaninfo", embyEmptyArrayHandler(svc))
- auth.GET("/scheduledtasks", embyEmptyArrayHandler(svc))
- auth.GET("/livetv/recordings", embyEmptyItemsHandler(svc))
- auth.GET("/system/activitylog/entries", embyEmptyItemsHandler(svc))
- auth.GET("/web/configurationpages", embyEmptyArrayHandler(svc))
- auth.POST("/users/:userId/configuration", embyNoContentHandler(svc))
-}
diff --git a/internal/handler/emby_auth.go b/internal/handler/emby_auth.go
new file mode 100644
index 0000000..8866b3e
--- /dev/null
+++ b/internal/handler/emby_auth.go
@@ -0,0 +1,241 @@
+package handler
+
+import (
+ "strings"
+ "sync"
+ "time"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+)
+
+// embyError 返回 Emby 风格的错误(顶层 Code/Message)。
+func embyError(c *gin.Context, status int, msg string) {
+ c.JSON(status, gin.H{"Code": status, "Message": msg})
+}
+
+// embyUserID 从中间件中获取 user id。Emby auth middleware 写入 CtxUserID。
+func embyUserID(c *gin.Context) string {
+ if uid, ok := c.Get(middleware.CtxUserID); ok {
+ if s, ok := uid.(string); ok {
+ return s
+ }
+ }
+ return ""
+}
+
+const embyCompatSessionTTL = 30 * time.Minute
+
+type embyCompatSession struct {
+ token string
+ expiresAt time.Time
+}
+
+var embyCompatSessions = struct {
+ sync.RWMutex
+ items map[string]embyCompatSession
+}{items: map[string]embyCompatSession{}}
+
+func embyAuthRequiredWithSessionFallback(secret string) gin.HandlerFunc {
+ required := middleware.EmbyAuthRequired(secret)
+ return func(c *gin.Context) {
+ if embyRequestToken(c) == "" {
+ if token := embyCompatSessionToken(c); token != "" {
+ c.Request.Header.Set("X-Emby-Token", token)
+ }
+ }
+ required(c)
+ }
+}
+
+func embyRememberCompatSession(c *gin.Context, token string) {
+ token = strings.TrimSpace(token)
+ if token == "" {
+ return
+ }
+ keys := embyCompatSessionKeys(c)
+ if len(keys) == 0 {
+ return
+ }
+ expiresAt := time.Now().Add(embyCompatSessionTTL)
+ embyCompatSessions.Lock()
+ defer embyCompatSessions.Unlock()
+ if len(embyCompatSessions.items) > 1000 {
+ now := time.Now()
+ for key, session := range embyCompatSessions.items {
+ if now.After(session.expiresAt) {
+ delete(embyCompatSessions.items, key)
+ }
+ }
+ if len(embyCompatSessions.items) > 1000 {
+ embyCompatSessions.items = map[string]embyCompatSession{}
+ }
+ }
+ for _, key := range keys {
+ embyCompatSessions.items[key] = embyCompatSession{token: token, expiresAt: expiresAt}
+ }
+}
+
+func embyCompatSessionToken(c *gin.Context) string {
+ keys := embyCompatSessionKeys(c)
+ if len(keys) == 0 {
+ return ""
+ }
+ now := time.Now()
+ embyCompatSessions.RLock()
+ defer embyCompatSessions.RUnlock()
+ for _, key := range keys {
+ session, ok := embyCompatSessions.items[key]
+ if ok && now.Before(session.expiresAt) {
+ return session.token
+ }
+ }
+ return ""
+}
+
+func embyCompatSessionKeys(c *gin.Context) []string {
+ if c == nil {
+ return nil
+ }
+ ip := strings.TrimSpace(c.ClientIP())
+ if ip == "" {
+ return nil
+ }
+ keys := []string{}
+ add := func(kind, value string) {
+ value = strings.TrimSpace(value)
+ if value != "" {
+ keys = append(keys, ip+"\x00"+kind+"\x00"+value)
+ }
+ }
+ add("device", firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"))
+ add("ua", c.GetHeader("User-Agent"))
+ return keys
+}
+
+func firstHeaderValue(c *gin.Context, names ...string) string {
+ for _, name := range names {
+ if value := strings.TrimSpace(c.GetHeader(name)); value != "" {
+ return value
+ }
+ }
+ return ""
+}
+
+type embyClientInfo struct {
+ DeviceID string
+ DeviceName string
+ Client string
+}
+
+func embyClientInfoFromRequest(c *gin.Context) embyClientInfo {
+ auth := parseMediaBrowserAuthorization(firstHeaderValue(c,
+ "X-Emby-Authorization",
+ "X-MediaBrowser-Authorization",
+ "Authorization",
+ ))
+ info := embyClientInfo{
+ DeviceID: firstNonEmptyHeaderString(
+ firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"),
+ auth["DeviceId"],
+ auth["DeviceID"],
+ ),
+ DeviceName: firstNonEmptyHeaderString(
+ firstHeaderValue(c, "X-Emby-Device-Name", "X-Emby-DeviceName", "X-MediaBrowser-Device-Name", "X-MediaBrowser-DeviceName"),
+ auth["Device"],
+ ),
+ Client: firstNonEmptyHeaderString(
+ firstHeaderValue(c, "X-Emby-Client", "X-MediaBrowser-Client"),
+ auth["Client"],
+ ),
+ }
+ ua := strings.TrimSpace(c.GetHeader("User-Agent"))
+ if info.Client == "" {
+ info.Client = embyClientFromUserAgent(ua)
+ }
+ if info.DeviceName == "" {
+ info.DeviceName = embyDeviceFromUserAgent(ua)
+ }
+ return info
+}
+
+func parseMediaBrowserAuthorization(raw string) map[string]string {
+ out := map[string]string{}
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return out
+ }
+ for _, prefix := range []string{"MediaBrowser ", "Emby "} {
+ if strings.HasPrefix(raw, prefix) {
+ raw = strings.TrimSpace(strings.TrimPrefix(raw, prefix))
+ break
+ }
+ }
+ for _, part := range strings.Split(raw, ",") {
+ key, value, ok := strings.Cut(strings.TrimSpace(part), "=")
+ if !ok {
+ continue
+ }
+ key = strings.TrimSpace(key)
+ value = strings.Trim(strings.TrimSpace(value), `"`)
+ if key != "" && value != "" {
+ out[key] = value
+ }
+ }
+ return out
+}
+
+func firstNonEmptyHeaderString(values ...string) string {
+ for _, value := range values {
+ if strings.TrimSpace(value) != "" {
+ return strings.TrimSpace(value)
+ }
+ }
+ return ""
+}
+
+func embyClientFromUserAgent(ua string) string {
+ ua = strings.TrimSpace(ua)
+ lower := strings.ToLower(ua)
+ switch {
+ case strings.Contains(lower, "infuse"):
+ return "Infuse"
+ case strings.Contains(lower, "emby"):
+ return "Emby"
+ case strings.Contains(lower, "jellyfin"):
+ return "Jellyfin"
+ case strings.Contains(lower, "yamby"):
+ return "Yamby"
+ case strings.Contains(lower, "vidhub"):
+ return "VidHub"
+ case strings.Contains(lower, "hills"):
+ return "Hills"
+ default:
+ return ua
+ }
+}
+
+func embyDeviceFromUserAgent(ua string) string {
+ lower := strings.ToLower(strings.TrimSpace(ua))
+ switch {
+ case strings.Contains(lower, "android"):
+ return "Android"
+ case strings.Contains(lower, "iphone"):
+ return "iPhone"
+ case strings.Contains(lower, "ipad"):
+ return "iPad"
+ case strings.Contains(lower, "ios"):
+ return "iOS"
+ case strings.Contains(lower, "windows"):
+ return "Windows PC"
+ case strings.Contains(lower, "macintosh") || strings.Contains(lower, "mac os"):
+ return "Mac"
+ case strings.Contains(lower, "linux"):
+ return "Linux PC"
+ case strings.Contains(lower, "appletv") || strings.Contains(lower, "apple tv"):
+ return "Apple TV"
+ default:
+ return ""
+ }
+}
diff --git a/internal/handler/emby_auth_request.go b/internal/handler/emby_auth_request.go
new file mode 100644
index 0000000..9df6c45
--- /dev/null
+++ b/internal/handler/emby_auth_request.go
@@ -0,0 +1,174 @@
+package handler
+
+import (
+ "bytes"
+ "encoding/json"
+ "errors"
+ "io"
+ "net/url"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+)
+
+type embyAuthByNameReq struct {
+ Username string `json:"Username"`
+ Pw string `json:"Pw"`
+ Password string `json:"Password"`
+ PasswordMd5 string `json:"PasswordMd5"`
+ PasswordSha1 string `json:"PasswordSha1"`
+}
+
+func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
+ req := embyAuthByNameReq{}
+ if strings.Contains(strings.ToLower(c.GetHeader("Content-Type")), "json") {
+ var body map[string]any
+ if err := c.ShouldBindJSON(&body); err != nil && !errors.Is(err, io.EOF) {
+ return req, err
+ }
+ fillEmbyAuthFromMap(&req, body)
+ }
+
+ if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
+ _ = c.Request.ParseForm()
+ if req.Username == "" {
+ req.Username = firstFormValue(c, "Username", "username", "Name", "name")
+ }
+ if req.Pw == "" {
+ req.Pw = firstFormValue(c, "Pw", "pw")
+ }
+ if req.Password == "" {
+ req.Password = firstFormValue(c, "Password", "password")
+ }
+ if req.PasswordMd5 == "" {
+ req.PasswordMd5 = firstFormValue(c, "PasswordMd5", "passwordMd5", "password_md5")
+ }
+ if req.PasswordSha1 == "" {
+ req.PasswordSha1 = firstFormValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
+ }
+ }
+
+ if req.Username == "" {
+ req.Username = firstQueryValue(c, "Username", "username", "Name", "name")
+ }
+ if req.Pw == "" {
+ req.Pw = firstQueryValue(c, "Pw", "pw")
+ }
+ if req.Password == "" {
+ req.Password = firstQueryValue(c, "Password", "password")
+ }
+ if req.PasswordMd5 == "" {
+ req.PasswordMd5 = firstQueryValue(c, "PasswordMd5", "passwordMd5", "password_md5")
+ }
+ if req.PasswordSha1 == "" {
+ req.PasswordSha1 = firstQueryValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
+ }
+ if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
+ fillEmbyAuthFromRawBody(c, &req)
+ }
+ return req, nil
+}
+
+func fillEmbyAuthFromMap(req *embyAuthByNameReq, body map[string]any) {
+ if req.Username == "" {
+ req.Username = firstStringFromMap(body, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
+ }
+ if req.Pw == "" {
+ req.Pw = firstStringFromMap(body, "Pw", "pw", "PW")
+ }
+ if req.Password == "" {
+ req.Password = firstStringFromMap(body, "Password", "password", "Pass", "pass", "Pwd", "pwd")
+ }
+ if req.PasswordMd5 == "" {
+ req.PasswordMd5 = firstStringFromMap(body, "PasswordMd5", "passwordMd5", "password_md5")
+ }
+ if req.PasswordSha1 == "" {
+ req.PasswordSha1 = firstStringFromMap(body, "PasswordSha1", "passwordSha1", "password_sha1")
+ }
+}
+
+func fillEmbyAuthFromRawBody(c *gin.Context, req *embyAuthByNameReq) {
+ if c.Request == nil || c.Request.Body == nil {
+ return
+ }
+ raw, err := io.ReadAll(io.LimitReader(c.Request.Body, 1<<20))
+ if err != nil {
+ return
+ }
+ c.Request.Body = io.NopCloser(bytes.NewReader(raw))
+ raw = bytes.TrimSpace(raw)
+ if len(raw) == 0 {
+ return
+ }
+ if bytes.HasPrefix(raw, []byte("{")) {
+ var body map[string]any
+ if err := json.Unmarshal(raw, &body); err == nil {
+ fillEmbyAuthFromMap(req, body)
+ }
+ return
+ }
+ if values, err := url.ParseQuery(string(raw)); err == nil {
+ fillEmbyAuthFromValues(req, values)
+ }
+}
+
+func fillEmbyAuthFromValues(req *embyAuthByNameReq, values url.Values) {
+ if req.Username == "" {
+ req.Username = firstValue(values, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
+ }
+ if req.Pw == "" {
+ req.Pw = firstValue(values, "Pw", "pw", "PW")
+ }
+ if req.Password == "" {
+ req.Password = firstValue(values, "Password", "password", "Pass", "pass", "Pwd", "pwd")
+ }
+ if req.PasswordMd5 == "" {
+ req.PasswordMd5 = firstValue(values, "PasswordMd5", "passwordMd5", "password_md5")
+ }
+ if req.PasswordSha1 == "" {
+ req.PasswordSha1 = firstValue(values, "PasswordSha1", "passwordSha1", "password_sha1")
+ }
+}
+
+func firstValue(values url.Values, keys ...string) string {
+ for _, key := range keys {
+ if value := strings.TrimSpace(values.Get(key)); value != "" {
+ return value
+ }
+ }
+ return ""
+}
+
+func firstStringFromMap(body map[string]any, keys ...string) string {
+ if len(body) == 0 {
+ return ""
+ }
+ for _, key := range keys {
+ if value, ok := body[key]; ok {
+ if s, ok := value.(string); ok {
+ return strings.TrimSpace(s)
+ }
+ }
+ }
+ return ""
+}
+
+func firstFormValue(c *gin.Context, keys ...string) string {
+ for _, key := range keys {
+ if values, ok := c.Request.PostForm[key]; ok && len(values) > 0 {
+ if value := strings.TrimSpace(values[0]); value != "" {
+ return value
+ }
+ }
+ }
+ return ""
+}
+
+func firstQueryValue(c *gin.Context, keys ...string) string {
+ for _, key := range keys {
+ if value := strings.TrimSpace(c.Query(key)); value != "" {
+ return value
+ }
+ }
+ return ""
+}
diff --git a/internal/handler/emby_auth_test.go b/internal/handler/emby_auth_test.go
new file mode 100644
index 0000000..61d3f72
--- /dev/null
+++ b/internal/handler/emby_auth_test.go
@@ -0,0 +1,168 @@
+package handler
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestParseEmbyAuthByNameReqAcceptsLowercaseJSON(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader(`{"username":"alice","password":"secret"}`))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ req, err := parseEmbyAuthByNameReq(c)
+ if err != nil {
+ t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
+ }
+ if req.Username != "alice" || req.Password != "secret" {
+ t.Fatalf("unexpected request: %#v", req)
+ }
+}
+
+func TestParseEmbyAuthByNameReqAcceptsFormBody(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader("Username=bob&Pw=secret"))
+ c.Request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
+
+ req, err := parseEmbyAuthByNameReq(c)
+ if err != nil {
+ t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
+ }
+ if req.Username != "bob" || req.Pw != "secret" {
+ t.Fatalf("unexpected request: %#v", req)
+ }
+}
+
+func TestParseEmbyAuthByNameReqAcceptsJSONWithoutContentType(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"UserName":"carol","PW":"secret"}`))
+
+ req, err := parseEmbyAuthByNameReq(c)
+ if err != nil {
+ t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
+ }
+ if req.Username != "carol" || req.Pw != "secret" {
+ t.Fatalf("unexpected request: %#v", req)
+ }
+}
+
+func TestEmbyAuthenticateByNameAcceptsCaseVariantUsernameAndPath(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.Setting{}); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.JWTSecret = "test-secret"
+ log := zap.NewNop()
+ permissions := service.NewPermissionService(log, repos)
+ auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
+ if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
+ t.Fatalf("register: %v", err)
+ }
+
+ router := gin.New()
+ registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
+ Repo: repos,
+ Auth: auth,
+ Emby: service.NewEmbyService(cfg, log, repos),
+ Audit: service.NewAuditService(log, repos),
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"Username":"Viewer","Pw":"secret-pass"}`))
+ req.Header.Set("Content-Type", "application/json")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ var payload map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if payload["AccessToken"] == "" {
+ t.Fatalf("missing AccessToken: %#v", payload)
+ }
+}
+
+func TestEmbyAuthenticateRecordsMediaBrowserClientInfo(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if sqlDB, err := db.DB(); err == nil {
+ sqlDB.SetMaxOpenConns(1)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.JWTSecret = "test-secret"
+ log := zap.NewNop()
+ permissions := service.NewPermissionService(log, repos)
+ auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
+ if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
+ t.Fatalf("register: %v", err)
+ }
+
+ router := gin.New()
+ registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
+ Repo: repos,
+ Auth: auth,
+ Emby: service.NewEmbyService(cfg, log, repos),
+ Device: service.NewDeviceService(log, repos),
+ Audit: service.NewAuditService(log, repos),
+ Permissions: permissions,
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/emby/Users/AuthenticateByName", strings.NewReader(`{"Username":"viewer","Pw":"secret-pass"}`))
+ req.Header.Set("Content-Type", "application/json")
+ req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", DeviceId="device-42"`)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ user, err := repos.User.FindByUsername(context.Background(), "viewer")
+ if err != nil {
+ t.Fatalf("find user: %v", err)
+ }
+ devices, err := repos.UserDevice.ListByUser(context.Background(), user.ID)
+ if err != nil {
+ t.Fatalf("list devices: %v", err)
+ }
+ if len(devices) != 1 {
+ t.Fatalf("devices = %#v, want one recorded device", devices)
+ }
+ if devices[0].DeviceID != "device-42" || devices[0].DeviceName != "PC" || devices[0].Client != "Infuse" {
+ t.Fatalf("device info not parsed from MediaBrowser header: %#v", devices[0])
+ }
+}
diff --git a/internal/handler/emby_images.go b/internal/handler/emby_images.go
new file mode 100644
index 0000000..49e1915
--- /dev/null
+++ b/internal/handler/emby_images.go
@@ -0,0 +1,72 @@
+package handler
+
+import (
+ "context"
+ "net/http"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+var embyPlaceholderPNG = []byte{
+ 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
+ 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
+ 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
+ 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
+ 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
+ 0x54, 0x78, 0x9c, 0x63, 0x50, 0xd1, 0x30, 0xf8,
+ 0x0f, 0x00, 0x02, 0x6c, 0x01, 0x7c, 0x30, 0xed,
+ 0x6e, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45,
+ 0x4e, 0x44, 0xae, 0x42, 0x60, 0x82,
+}
+
+// embyItemImageHandler 把 /Items/{id}/Images/Primary 等请求直接输出为图片。
+// Emby 客户端缓存图片 URL 时经常不会继续携带 token;如果重定向到受保护的
+// /api/img 会变成 401,所以这里复用 ImageProxy 但不再走 /api 路由。
+func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ clearEmbyImageNoStoreHeaders(c)
+ ctx, cancel := context.WithTimeout(c.Request.Context(), 8*time.Second)
+ defer cancel()
+ req := c.Request.WithContext(ctx)
+ id := c.Param("id")
+ imgType := strings.ToLower(c.Param("type"))
+ raw, err := svc.Emby.ImageURL(ctx, id, imgType)
+ if err != nil || raw == "" {
+ embyServePlaceholderImage(c)
+ return
+ }
+ if typ, ref, ok := service.ParseCloudArtworkURL(raw); ok {
+ c.Request = req
+ serveCloudResolvedLink(svc, c, typ, ref)
+ return
+ }
+ if svc.ImageProxy == nil {
+ embyServePlaceholderImage(c)
+ return
+ }
+ if err := svc.ImageProxy.Serve(ctx, c.Writer, req, raw); err != nil {
+ embyServePlaceholderImage(c)
+ }
+ }
+}
+
+func clearEmbyImageNoStoreHeaders(c *gin.Context) {
+ c.Writer.Header().Del("Pragma")
+ c.Writer.Header().Del("Expires")
+}
+
+func embyServePlaceholderImage(c *gin.Context) {
+ c.Header("Content-Type", "image/png")
+ c.Header("Cache-Control", "public, max-age=3600")
+ c.Header("Content-Length", strconv.Itoa(len(embyPlaceholderPNG)))
+ if c.Request.Method == http.MethodHead {
+ c.Status(http.StatusOK)
+ return
+ }
+ c.Data(http.StatusOK, "image/png", embyPlaceholderPNG)
+}
diff --git a/internal/handler/emby_items_handlers.go b/internal/handler/emby_items_handlers.go
new file mode 100644
index 0000000..d182880
--- /dev/null
+++ b/internal/handler/emby_items_handlers.go
@@ -0,0 +1,231 @@
+package handler
+
+import (
+ "net/http"
+ "strconv"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func parseEmbyItemsParams(c *gin.Context) service.ItemsParams {
+ limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "50"))
+ offset, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "StartIndex", "startIndex", "startindex"), "0"))
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = firstQueryValue(c, "UserId", "userId", "userid")
+ }
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ splitOpt := func(s string) []string {
+ if s == "" {
+ return nil
+ }
+ parts := strings.Split(s, ",")
+ out := make([]string, 0, len(parts))
+ for _, p := range parts {
+ p = strings.TrimSpace(p)
+ if p != "" {
+ out = append(out, p)
+ }
+ }
+ return out
+ }
+ return service.ItemsParams{
+ UserID: uid,
+ ParentID: firstQueryValue(c, "ParentId", "parentId", "parentid"),
+ IDs: splitOpt(firstQueryValue(c, "Ids", "ids")),
+ SearchTerm: firstQueryValue(c, "SearchTerm", "searchTerm", "searchterm"),
+ IncludeItemTypes: splitOpt(firstQueryValue(c, "IncludeItemTypes", "includeItemTypes", "includeitemtypes")),
+ Filters: splitOpt(firstQueryValue(c, "Filters", "filters")),
+ Recursive: strings.EqualFold(firstQueryValue(c, "Recursive", "recursive"), "true"),
+ SortBy: firstQueryValue(c, "SortBy", "sortBy", "sortby"),
+ SortOrder: firstQueryValue(c, "SortOrder", "sortOrder", "sortorder"),
+ Limit: limit,
+ StartIndex: offset,
+ }
+}
+
+func embyFirstNonEmptyString(values ...string) string {
+ for _, value := range values {
+ if strings.TrimSpace(value) != "" {
+ return strings.TrimSpace(value)
+ }
+ }
+ return ""
+}
+
+func embyItemsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ out, err := svc.Emby.Items(c.Request.Context(), parseEmbyItemsParams(c))
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyItemByIDHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ id := c.Param("id")
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ out, err := svc.Emby.Item(c.Request.Context(), id, uid)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ if out == nil {
+ embyError(c, http.StatusNotFound, "item not found")
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyUserItemByIDHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ switch strings.ToLower(c.Param("id")) {
+ case "latest":
+ embyLatestItemsHandler(svc)(c)
+ case "resume":
+ embyResumeItemsHandler(svc)(c)
+ default:
+ embyItemByIDHandler(svc)(c)
+ }
+ }
+}
+
+func embyLatestItemsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = firstQueryValue(c, "UserId", "userId", "userid")
+ }
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
+ out, err := svc.Emby.LatestItems(c.Request.Context(), uid, firstQueryValue(c, "ParentId", "parentId", "parentid"), limit)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyResumeItemsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = firstQueryValue(c, "UserId", "userId", "userid")
+ }
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
+ out, err := svc.Emby.ResumeItems(c.Request.Context(), uid, limit)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyItemsCountsHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "MovieCount": 0,
+ "SeriesCount": 0,
+ "EpisodeCount": 0,
+ "ItemCount": 0,
+ })
+ }
+}
+
+func embyDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "Id": c.Param("id"),
+ "ViewType": "Poster",
+ "SortBy": "SortName",
+ "SortOrder": "Ascending",
+ "IndexBy": "SortName",
+ "RememberIndexing": false,
+ "PrimaryImageHeight": 250,
+ "PrimaryImageWidth": 250,
+ "ScrollDirection": "Horizontal",
+ "ShowSidebar": true,
+ "CustomPrefs": gin.H{
+ "homeexploresection": "1",
+ "homesection0": "smalllibrarytiles",
+ "homesection1": "resume",
+ "homesection2": "latestmedia",
+ "homesection3": "nextup",
+ "homesection4": "none",
+ "homesection5": "none",
+ "homesection6": "none",
+ "latestItems": "true",
+ "landing-livetv": "false",
+ },
+ })
+ }
+}
+
+func embySaveDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Status(http.StatusNoContent)
+ }
+}
+
+func embyShowSeasonsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ params := service.ItemsParams{
+ UserID: firstQueryValue(c, "UserId", "userId"),
+ ParentID: c.Param("id"),
+ Limit: 500,
+ }
+ out, err := svc.Emby.Items(c.Request.Context(), params)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyShowEpisodesHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ parentID := firstQueryValue(c, "SeasonId", "seasonId")
+ if parentID == "" {
+ parentID = c.Param("id")
+ }
+ params := service.ItemsParams{
+ UserID: firstQueryValue(c, "UserId", "userId"),
+ ParentID: parentID,
+ IncludeItemTypes: []string{"Episode"},
+ Recursive: true,
+ Limit: 500,
+ }
+ out, err := svc.Emby.Items(c.Request.Context(), params)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
diff --git a/internal/handler/emby_items_test.go b/internal/handler/emby_items_test.go
new file mode 100644
index 0000000..fcb4b0b
--- /dev/null
+++ b/internal/handler/emby_items_test.go
@@ -0,0 +1,301 @@
+package handler
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "os"
+ "path/filepath"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(&model.Media{}); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+
+ posterPath := filepath.Join(t.TempDir(), "poster.png")
+ if err := os.WriteFile(posterPath, []byte{
+ 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
+ 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
+ 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
+ 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
+ 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
+ 0x54, 0x78, 0x9c, 0x63, 0x00, 0x01, 0x00, 0x00,
+ 0x05, 0x00, 0x01, 0x0d, 0x0a, 0x2d, 0xb4, 0x00,
+ 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae,
+ 0x42, 0x60, 0x82,
+ }, 0o644); err != nil {
+ t.Fatalf("write poster: %v", err)
+ }
+
+ repos := repository.New(db)
+ cfg := &config.Config{
+ App: config.AppConfig{DataDir: filepath.Dir(posterPath)},
+ Cache: config.CacheConfig{CacheDir: t.TempDir()},
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ Title: "Poster Test",
+ Path: "D:\\media\\poster-test.mp4",
+ PosterURL: posterPath,
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ router := gin.New()
+ registerEmbyRoutes(router, "test-secret", &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
+ ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/Items/media-1/Images/Primary", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if location := w.Header().Get("Location"); location != "" {
+ t.Fatalf("expected direct image response, got redirect to %q", location)
+ }
+ if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
+ t.Fatalf("expected png content type, got %q", contentType)
+ }
+ if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
+ t.Fatalf("image Cache-Control = %q, want long browser cache", got)
+ }
+ if got := w.Header().Get("Pragma"); got != "" {
+ t.Fatalf("image Pragma = %q, want empty", got)
+ }
+ if got := w.Header().Get("Expires"); got != "" {
+ t.Fatalf("image Expires = %q, want empty", got)
+ }
+}
+
+func TestEmbyItemImageServesCachedCloudArtworkWithoutResolve(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("emby-cached-cloud-poster"))
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(&model.Media{}); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+
+ cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
+ imageProxy := service.NewImageProxy(cfg, zap.NewNop())
+ ref := "/Movies/Cloud Movie/poster.jpg"
+ if err := imageProxy.PrefetchCloudResolved(t.Context(), "openlist:"+ref, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
+ t.Fatalf("prefetch cloud poster: %v", err)
+ }
+ repos := repository.New(db)
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-media-1"},
+ Title: "Cloud Poster Test",
+ Path: "cloud://openlist/Movies/Cloud Movie/movie.mkv",
+ PosterURL: service.CloudArtworkURL("openlist", ref),
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ router := gin.New()
+ registerEmbyRoutes(router, "test-secret", &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
+ ImageProxy: imageProxy,
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/Items/cloud-media-1/Images/Primary", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if got := w.Body.String(); got != "emby-cached-cloud-poster" {
+ t.Fatalf("body = %q, want cached cloud poster", got)
+ }
+ if location := w.Header().Get("Location"); location != "" {
+ t.Fatalf("expected direct cached image response, got redirect to %q", location)
+ }
+ if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
+ t.Fatalf("image Cache-Control = %q, want long browser cache", got)
+ }
+}
+
+func TestEmbyMissingItemImageReturnsTransparentPlaceholder(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
+ router := gin.New()
+ registerEmbyRoutes(router, "test-secret", &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
+ ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
+ })
+
+ req := httptest.NewRequest(http.MethodHead, "/Items/missing/Images/Primary", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("expected placeholder status 200, got %d body=%s", w.Code, w.Body.String())
+ }
+ if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
+ t.Fatalf("expected png content type, got %q", contentType)
+ }
+ if length := w.Header().Get("Content-Length"); length == "" || length == "0" {
+ t.Fatalf("expected placeholder content length, got %q", length)
+ }
+ if got := w.Header().Get("Pragma"); got != "" {
+ t.Fatalf("placeholder Pragma = %q, want empty", got)
+ }
+ if got := w.Header().Get("Expires"); got != "" {
+ t.Fatalf("placeholder Expires = %q, want empty", got)
+ }
+}
+
+func TestEmbyUserItemByIDRouteReturnsJSON(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "episode-1"},
+ LibraryID: lib.ID,
+ Title: "Test Show",
+ Path: "D:\\media\\tv\\Test Show\\Season 01\\Test Show - S01E01.mkv",
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/episode-1", nil)
+ req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
+ req.Header.Set("If-None-Match", `"stale-client-cache"`)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "application/json") {
+ t.Fatalf("expected JSON content type, got %q body=%s", contentType, w.Body.String())
+ }
+ var item map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
+ t.Fatalf("decode item: %v", err)
+ }
+ if item["Id"] != "episode-1" || item["Type"] != "Episode" {
+ t.Fatalf("unexpected item payload: %#v", item)
+ }
+}
+
+func TestEmbyUserItemByIDRouteReturnsLibraryView(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Base: model.Base{ID: "lib-tv"}, Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/lib-tv", nil)
+ req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ var item map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
+ t.Fatalf("decode item: %v", err)
+ }
+ if item["Id"] != "lib-tv" || item["Type"] != "CollectionFolder" || item["CollectionType"] != "tvshows" {
+ t.Fatalf("unexpected library payload: %#v", item)
+ }
+}
diff --git a/internal/handler/emby_playback.go b/internal/handler/emby_playback.go
new file mode 100644
index 0000000..89a45a9
--- /dev/null
+++ b/internal/handler/emby_playback.go
@@ -0,0 +1,272 @@
+package handler
+
+import (
+ "errors"
+ "net/http"
+ "net/url"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ out, err := svc.Emby.PlaybackInfo(c.Request.Context(), c.Param("id"), uid)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ if out == nil {
+ embyError(c, http.StatusNotFound, "not found")
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyAttachRequestTokenToMediaSources(c *gin.Context, out any) {
+ token := embyRequestToken(c)
+ if token == "" || out == nil {
+ return
+ }
+ embyAttachTokenToMediaSourcesValue(out, token)
+}
+
+func embyAttachTokenToMediaSourcesValue(value any, token string) {
+ switch typed := value.(type) {
+ case map[string]any:
+ embyAttachTokenToMediaSourcesMap(typed, token)
+ case gin.H:
+ embyAttachTokenToMediaSourcesMap(map[string]any(typed), token)
+ case []map[string]any:
+ for _, item := range typed {
+ embyAttachTokenToMediaSourcesMap(item, token)
+ }
+ case []any:
+ for _, item := range typed {
+ embyAttachTokenToMediaSourcesValue(item, token)
+ }
+ }
+}
+
+func embyAttachTokenToMediaSourcesMap(out map[string]any, token string) {
+ if out == nil {
+ return
+ }
+ if sources, ok := out["MediaSources"].([]map[string]any); ok {
+ embyAttachTokenToMediaSources(sources, token)
+ } else if sources, ok := out["MediaSources"].([]any); ok {
+ for _, source := range sources {
+ if sourceMap, ok := source.(map[string]any); ok {
+ embyAttachTokenToMediaSources([]map[string]any{sourceMap}, token)
+ }
+ }
+ }
+ if items, ok := out["Items"]; ok {
+ embyAttachTokenToMediaSourcesValue(items, token)
+ }
+}
+
+func embyAttachTokenToMediaSources(sources []map[string]any, token string) {
+ for _, source := range sources {
+ for _, key := range []string{"DirectStreamUrl", "TranscodingUrl"} {
+ raw, ok := source[key].(string)
+ if !ok {
+ continue
+ }
+ source[key] = embyAppendAPIKey(raw, token)
+ }
+ }
+}
+
+func embyRequestToken(c *gin.Context) string {
+ if c == nil {
+ return ""
+ }
+ for _, key := range []string{"api_key", "apiKey", "ApiKey", "token", "X-Emby-Token", "X-MediaBrowser-Token"} {
+ if value := strings.TrimSpace(c.Query(key)); value != "" {
+ return value
+ }
+ }
+ for _, header := range []string{"X-Emby-Token", "X-MediaBrowser-Token"} {
+ if value := strings.TrimSpace(c.GetHeader(header)); value != "" {
+ return value
+ }
+ }
+ for _, header := range []string{"Authorization", "X-Emby-Authorization", "X-MediaBrowser-Authorization"} {
+ if token := embyTokenFromAuthHeader(c.GetHeader(header)); token != "" {
+ return token
+ }
+ }
+ return ""
+}
+
+func embyTokenFromAuthHeader(value string) string {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ return ""
+ }
+ for _, prefix := range []string{"Bearer ", "Emby "} {
+ if strings.HasPrefix(value, prefix) {
+ return strings.TrimSpace(strings.TrimPrefix(value, prefix))
+ }
+ }
+ for _, part := range strings.Split(value, ",") {
+ part = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(part), "MediaBrowser "))
+ if !strings.HasPrefix(part, "Token=") {
+ continue
+ }
+ token := strings.TrimSpace(strings.TrimPrefix(part, "Token="))
+ return strings.Trim(token, `"`)
+ }
+ if strings.Contains(value, "Token=") {
+ return ""
+ }
+ return value
+}
+
+func embyAppendAPIKey(raw, token string) string {
+ raw = strings.TrimSpace(raw)
+ token = strings.TrimSpace(token)
+ if raw == "" || token == "" {
+ return raw
+ }
+ if strings.HasPrefix(raw, "//") {
+ return raw
+ }
+ u, err := url.Parse(raw)
+ if err != nil || u.IsAbs() {
+ return raw
+ }
+ q := u.Query()
+ if q.Get("api_key") == "" && q.Get("apiKey") == "" && q.Get("token") == "" {
+ q.Set("api_key", token)
+ u.RawQuery = q.Encode()
+ }
+ return u.String()
+}
+
+// embyVideoStreamHandler 是 GET /Videos/{id}/stream 的入口,
+// 直接代理到我们的 /api/stream/{id}(同一个 ServeFile)。
+func embyVideoStreamHandler(svc *service.Container, cloudMode string) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := embyUserID(c)
+ item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ if item == nil {
+ c.Status(http.StatusNotFound)
+ return
+ }
+ if embyShouldRedirectVideoStreamToSTRM(c, svc, c.Param("id"), cloudMode) {
+ target := "/api/stream/" + url.PathEscape(strings.TrimSpace(c.Param("id")))
+ if token := embyPlaybackRedirectToken(c, svc); token != "" {
+ target = embyAppendAPIKey(target, token)
+ }
+ setRedirectNoStoreHeaders(c)
+ c.Redirect(http.StatusFound, absoluteRequestURL(c, target))
+ return
+ }
+ // 直接调用 Stream service 写入 response。
+ // 此前这里把所有错误一律吞成 404:云盘 Cookie 过期、直链解析失败、
+ // STRM 播放被关闭……在第三方播放器上全部表现为「404 不存在」,
+ // 无法排查。现在区分:行不存在→404;云盘播放不可用/上游故障→502+原因。
+ err = svc.Stream.ServeFileWithCloudMode(c.Writer, c.Request, c.Param("id"), cloudMode)
+ switch {
+ case err == nil:
+ case errors.Is(err, service.ErrMediaNotFound):
+ c.Status(http.StatusNotFound)
+ case errors.Is(err, service.ErrCloudPlaybackDisabled):
+ if !c.Writer.Written() {
+ c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
+ }
+ default:
+ if !c.Writer.Written() {
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ }
+ }
+ }
+}
+
+func embyPlaybackRedirectToken(c *gin.Context, svc *service.Container) string {
+ if token := embyRequestToken(c); token != "" {
+ return token
+ }
+ if c == nil || svc == nil || svc.Auth == nil || svc.Repo == nil || svc.Repo.User == nil {
+ return ""
+ }
+ uid := embyUserID(c)
+ if uid == "" {
+ return ""
+ }
+ u, err := svc.Repo.User.FindByID(c.Request.Context(), uid)
+ if err != nil || u == nil {
+ return ""
+ }
+ token, err := svc.Auth.IssueEmbyToken(u)
+ if err != nil {
+ return ""
+ }
+ return token
+}
+
+func embyShouldRedirectVideoStreamToSTRM(c *gin.Context, svc *service.Container, mediaID, cloudMode string) bool {
+ if c == nil || svc == nil || svc.Repo == nil || svc.Repo.Media == nil || cloudMode != service.CloudPlaybackModeRedirectProxy {
+ return false
+ }
+ settings := service.CloudPlaybackSettings(c.Request.Context(), svc.Repo)
+ if settings.PreferredMode != service.CloudPlaybackModeSTRM || !settings.STRMEnabled {
+ return false
+ }
+ m, err := svc.Repo.Media.FindByID(c.Request.Context(), mediaID)
+ if err != nil || m == nil {
+ return false
+ }
+ return strings.TrimSpace(m.STRMURL) != ""
+}
+
+func embyVideoHLSPlaylistHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := embyUserID(c)
+ item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
+ if err != nil || item == nil || svc.Stream == nil {
+ c.Status(http.StatusNotFound)
+ return
+ }
+ err = svc.Stream.ServeHLSPlaylist(c.Writer, c.Request, c.Param("id"))
+ if errors.Is(err, service.ErrTranscodeDisabled) {
+ c.JSON(http.StatusConflict, gin.H{"error": "transcode disabled"})
+ return
+ }
+ if errors.Is(err, service.ErrTranscodeBusy) {
+ c.JSON(http.StatusTooManyRequests, gin.H{"error": "transcode busy"})
+ return
+ }
+ if err != nil {
+ c.Status(http.StatusNotFound)
+ }
+ }
+}
+
+func embyVideoHLSSegmentHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := embyUserID(c)
+ item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
+ if err != nil || item == nil || svc.Stream == nil {
+ c.Status(http.StatusNotFound)
+ return
+ }
+ if err := svc.Stream.ServeHLSSegment(c.Writer, c.Request, c.Param("id"), c.Param("seg")); err != nil {
+ c.Status(http.StatusNotFound)
+ }
+ }
+}
diff --git a/internal/handler/emby_playback_routes_test.go b/internal/handler/emby_playback_routes_test.go
new file mode 100644
index 0000000..e915a73
--- /dev/null
+++ b/internal/handler/emby_playback_routes_test.go
@@ -0,0 +1,690 @@
+package handler
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "os"
+ "path/filepath"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestEmbyLowercasePlaybackInfoRouteReturnsJSON(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: t.TempDir(), Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ LibraryID: lib.ID,
+ Title: "Lowercase Playback",
+ Path: filepath.Join(lib.Path, "lowercase-playback.mp4"),
+ Container: "mp4",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/media-1/playbackinfo", nil)
+ req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ var body map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
+ t.Fatalf("decode playback info: %v", err)
+ }
+ if _, ok := body["MediaSources"]; !ok {
+ t.Fatalf("missing MediaSources: %#v", body)
+ }
+ sources, ok := body["MediaSources"].([]any)
+ if !ok || len(sources) == 0 {
+ t.Fatalf("unexpected MediaSources: %#v", body["MediaSources"])
+ }
+ source, ok := sources[0].(map[string]any)
+ if !ok {
+ t.Fatalf("unexpected MediaSource: %#v", sources[0])
+ }
+ directURL, _ := source["DirectStreamUrl"].(string)
+ if !strings.Contains(directURL, "api_key=") {
+ t.Fatalf("DirectStreamUrl should carry api_key for clients that do not repeat auth headers: %#v", source)
+ }
+ transcodeURL, _ := source["TranscodingUrl"].(string)
+ if transcodeURL != "" && !strings.Contains(transcodeURL, "api_key=") {
+ t.Fatalf("TranscodingUrl should carry api_key: %#v", source)
+ }
+}
+
+func TestEmbyPlaybackInfoDoesNotExposeTokenInCloudPath(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: "cloud://openlist/Movies/Movie.mkv",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/cloud-1/playbackinfo", nil)
+ req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ var body map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
+ t.Fatalf("decode playback info: %v", err)
+ }
+ source := body["MediaSources"].([]any)[0].(map[string]any)
+ pathURL, _ := source["Path"].(string)
+ if pathURL != "/api/stream/cloud-1" {
+ t.Fatalf("cloud Path should stay as non-tokenized display stream URL, got %#v", source)
+ }
+ if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
+ t.Fatalf("cloud Path must not expose auth key/token: %#v", source)
+ }
+ if strings.Contains(pathURL, "/api/cloud/play/") {
+ t.Fatalf("cloud Path should not expose naked cloud play URL: %#v", source)
+ }
+ directURL, _ := source["DirectStreamUrl"].(string)
+ if !strings.HasPrefix(directURL, "/api/stream/cloud-1") || !strings.Contains(directURL, "api_key=") {
+ t.Fatalf("DirectStreamUrl should stay tokenized: %#v", source)
+ }
+ if source["SupportsDirectPlay"] != true {
+ t.Fatalf("cloud media should advertise DirectPlay when tokenized Path is playable: %#v", source)
+ }
+ if source["SupportsTranscoding"] != false {
+ t.Fatalf("cloud media should not advertise host transcoding: %#v", source)
+ }
+}
+
+func TestEmbyItemsDoNotExposeTokenInEmbeddedCloudPath(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: "cloud://openlist/Movies/Movie.mkv",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ token := signedTestToken(t, secret)
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/emby/Users/user-1/Items?IncludeItemTypes=Movie&Recursive=true&Limit=5&X-Emby-Token="+token, nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ var body map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
+ t.Fatalf("decode items: %v", err)
+ }
+ items := body["Items"].([]any)
+ if len(items) != 1 {
+ t.Fatalf("unexpected items: %#v", body["Items"])
+ }
+ source := items[0].(map[string]any)["MediaSources"].([]any)[0].(map[string]any)
+ pathURL, _ := source["Path"].(string)
+ if pathURL != "/api/stream/cloud-1" {
+ t.Fatalf("embedded cloud Path should stay as non-tokenized display stream URL, got %#v", source)
+ }
+ if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
+ t.Fatalf("embedded cloud Path must not expose auth key/token: %#v", source)
+ }
+}
+
+func TestEmbyVideoStreamUsesSTRMWhenRedirectProxyDisabled(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
+ t.Fatalf("enable strm playback: %v", err)
+ }
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
+ t.Fatalf("disable redirect playback: %v", err)
+ }
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: "cloud://openlist/Movies/Movie.mkv",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
+ Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
+ })
+
+ token := signedTestToken(t, secret)
+ req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream?api_key="+token, nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusFound {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ loc := w.Header().Get("Location")
+ if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
+ t.Fatalf("STRM mode should redirect /Videos fallback to tokenized /api/stream, got %q", loc)
+ }
+ if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
+ t.Fatalf("STRM fallback redirect Cache-Control = %q, want no-store", got)
+ }
+ if strings.Contains(loc, "/api/cloud/play/") {
+ t.Fatalf("STRM mode should not expose cloud play directly from /Videos fallback: %q", loc)
+ }
+}
+
+func TestEmbyVideoStreamIssuesTokenForSessionFallbackSTRMRedirect(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
+ t.Fatalf("enable strm playback: %v", err)
+ }
+ if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
+ t.Fatalf("disable redirect playback: %v", err)
+ }
+ user := model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }
+ if err := repos.User.Create(t.Context(), &user); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: "cloud://openlist/Movies/Movie.mkv",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
+ svc := &service.Container{
+ Repo: repos,
+ Auth: service.NewAuthService(cfg, zap.NewNop(), repos, nil, nil),
+ Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
+ Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
+ }
+ router := gin.New()
+ router.GET("/videos/:id/stream", func(c *gin.Context) {
+ c.Set(middleware.CtxUserID, user.ID)
+ c.Set(middleware.CtxUserRole, user.Role)
+ embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)(c)
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusFound {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ loc := w.Header().Get("Location")
+ if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
+ t.Fatalf("session fallback redirect should include api_key for /api/stream, got %q", loc)
+ }
+}
+
+func TestEmbyLowercaseVideoStreamRouteServesMedia(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ dir := t.TempDir()
+ mediaPath := filepath.Join(dir, "sample.mp4")
+ if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
+ t.Fatalf("write media: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ LibraryID: lib.ID,
+ Title: "Lowercase Stream",
+ Path: mediaPath,
+ Container: "mp4",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/videos/media-1/stream?api_key="+signedTestToken(t, secret), nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if got := w.Body.String(); got != "fake-video-bytes" {
+ t.Fatalf("unexpected stream body: %q", got)
+ }
+}
+
+func TestEmbyPrefixedAPIStreamRouteServesMedia(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ dir := t.TempDir()
+ mediaPath := filepath.Join(dir, "sample.mp4")
+ if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
+ t.Fatalf("write media: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ LibraryID: lib.ID,
+ Title: "Prefixed API Stream",
+ Path: mediaPath,
+ Container: "mp4",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/emby/api/stream/media-1?api_key="+signedTestToken(t, secret), nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if got := w.Body.String(); got != "fake-video-bytes" {
+ t.Fatalf("unexpected stream body: %q", got)
+ }
+}
+
+func TestEmbyVideoStreamRedirectKeepsMediaBrowserAuthorizationToken(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: "cloud://openlist/Movies/Movie.mkv",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
+ Container: "mkv",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
+ })
+
+ token := signedTestToken(t, secret)
+ req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
+ req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", Token="`+token+`"`)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusFound {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ loc := w.Header().Get("Location")
+ if !strings.Contains(loc, "/api/cloud/play/openlist?") || !strings.Contains(loc, "token=") {
+ t.Fatalf("redirect Location should target tokenized cloud play endpoint, got %q", loc)
+ }
+}
+
+func TestEmbyLowercaseOriginalHeadRouteServesHeaders(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ dir := t.TempDir()
+ mediaPath := filepath.Join(dir, "sample.mp4")
+ if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
+ t.Fatalf("write media: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ LibraryID: lib.ID,
+ Title: "Lowercase Original",
+ Path: mediaPath,
+ Container: "mp4",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
+ })
+
+ req := httptest.NewRequest(http.MethodHead, "/videos/media-1/original.mp4?api_key="+signedTestToken(t, secret), nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
+ }
+ if w.Body.Len() != 0 {
+ t.Fatalf("HEAD response should not include body, got %q", w.Body.String())
+ }
+}
+
+func TestEmbyLowercaseVideoHLSRouteDoesNot404WhenDirectOnly(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("open db: %v", err)
+ }
+ if err := db.AutoMigrate(model.AllModels()...); err != nil {
+ t.Fatalf("migrate: %v", err)
+ }
+ repos := repository.New(db)
+ if err := repos.User.Create(t.Context(), &model.User{
+ Base: model.Base{ID: "user-1"},
+ Username: "tester",
+ PasswordHash: "x",
+ Role: "admin",
+ Tier: "plus",
+ IsActive: true,
+ }); err != nil {
+ t.Fatalf("create user: %v", err)
+ }
+ dir := t.TempDir()
+ mediaPath := filepath.Join(dir, "sample.mp4")
+ if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
+ t.Fatalf("write media: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := db.Create(&model.Media{
+ Base: model.Base{ID: "media-1"},
+ LibraryID: lib.ID,
+ Title: "Lowercase HLS",
+ Path: mediaPath,
+ Container: "mp4",
+ }).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+ if err := repos.Setting.Set(t.Context(), service.PlaybackDirectOnlySettingKey, "true"); err != nil {
+ t.Fatalf("set direct-only: %v", err)
+ }
+
+ const secret = "test-secret"
+ router := gin.New()
+ registerEmbyRoutes(router, secret, &service.Container{
+ Repo: repos,
+ Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
+ Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/videos/media-1/master.m3u8?api_key="+signedTestToken(t, secret), nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code == http.StatusNotFound {
+ t.Fatalf("lowercase HLS route should be registered, got 404")
+ }
+ if w.Code != http.StatusConflict {
+ t.Fatalf("direct-only HLS should return 409, got %d body=%s", w.Code, w.Body.String())
+ }
+}
diff --git a/internal/handler/emby_playstate_handlers.go b/internal/handler/emby_playstate_handlers.go
new file mode 100644
index 0000000..96023f7
--- /dev/null
+++ b/internal/handler/emby_playstate_handlers.go
@@ -0,0 +1,119 @@
+package handler
+
+import (
+ "net/http"
+ "strconv"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+type embyPlayingReq struct {
+ ItemId string `json:"ItemId"`
+ PositionTicks int64 `json:"PositionTicks"`
+ RunTimeTicks int64 `json:"RunTimeTicks"`
+}
+
+func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := embyUserID(c)
+ if uid == "" {
+ c.Status(http.StatusUnauthorized)
+ return
+ }
+ var req embyPlayingReq
+ _ = c.ShouldBindJSON(&req)
+ if req.ItemId == "" {
+ req.ItemId = c.Query("ItemId")
+ }
+ if req.PositionTicks == 0 {
+ req.PositionTicks, _ = strconv.ParseInt(c.Query("PositionTicks"), 10, 64)
+ }
+ if req.RunTimeTicks == 0 {
+ req.RunTimeTicks, _ = strconv.ParseInt(c.Query("RunTimeTicks"), 10, 64)
+ }
+ if req.ItemId == "" {
+ c.Status(http.StatusOK)
+ return
+ }
+ clientInfo := embyClientInfoFromRequest(c)
+ if svc.Device != nil && svc.Device.IsDeviceKicked(c.Request.Context(), uid, clientInfo.DeviceID) {
+ c.Status(http.StatusUnauthorized)
+ return
+ }
+ _ = svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks)
+ stopped := strings.Contains(strings.ToLower(c.FullPath()+" "+c.Request.URL.Path), "stopped")
+ if svc.Sessions != nil {
+ svc.Sessions.RecordPlayback(c.Request.Context(), uid, "",
+ clientInfo.DeviceID,
+ clientInfo.DeviceName,
+ clientInfo.Client,
+ c.ClientIP(),
+ req.ItemId,
+ req.PositionTicks,
+ req.RunTimeTicks,
+ stopped)
+ }
+ if svc.Device != nil && !stopped {
+ svc.Device.RecordPlayback(c.Request.Context(), uid,
+ clientInfo.DeviceID,
+ clientInfo.DeviceName,
+ clientInfo.Client)
+ }
+ c.Status(http.StatusNoContent)
+ }
+}
+
+func embyFavoriteHandler(svc *service.Container, fav bool) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ mid := c.Param("itemId")
+ if uid == "" || mid == "" {
+ c.Status(http.StatusBadRequest)
+ return
+ }
+ if err := svc.Emby.SetFavorite(c.Request.Context(), uid, mid, fav); err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
+ if out != nil {
+ c.JSON(http.StatusOK, out["UserData"])
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"IsFavorite": fav})
+ }
+}
+
+func embyMarkPlayedHandler(svc *service.Container, played bool) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ mid := c.Param("itemId")
+ if uid == "" || mid == "" {
+ c.Status(http.StatusBadRequest)
+ return
+ }
+ if err := svc.Emby.MarkPlayed(c.Request.Context(), uid, mid, played); err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ if played && svc.Device != nil {
+ clientInfo := embyClientInfoFromRequest(c)
+ svc.Device.RecordPlayback(c.Request.Context(), uid, clientInfo.DeviceID, clientInfo.DeviceName, clientInfo.Client)
+ }
+ out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
+ if out != nil {
+ c.JSON(http.StatusOK, out["UserData"])
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"Played": played})
+ }
+}
diff --git a/internal/handler/emby_routes.go b/internal/handler/emby_routes.go
new file mode 100644
index 0000000..967ee12
--- /dev/null
+++ b/internal/handler/emby_routes.go
@@ -0,0 +1,235 @@
+package handler
+
+import (
+ "time"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+// registerEmbyRoutes 在 r 上挂双前缀("" + "/emby")的 Emby 兼容路由。
+func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) {
+ for _, prefix := range []string{"/emby", ""} {
+ grp := r.Group(prefix)
+ grp.Use(embyNoStoreHeaders())
+
+ registerEmbyRootRoutes(grp, prefix, svc)
+ registerEmbyPublicRoutes(grp, svc)
+ registerEmbyPublicImageRoutes(grp, svc)
+
+ // 鉴权后端点
+ auth := grp.Group("", embyAuthRequiredWithSessionFallback(jwtSecret), activeEmbyUserRequired(svc))
+ registerEmbyAuthenticatedRoutes(auth, prefix, svc)
+ }
+}
+
+type embyRouteHandlerFactory func(*service.Container) gin.HandlerFunc
+
+func embyNoStoreHeaders() gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Header("Cache-Control", "no-store")
+ c.Header("Pragma", "no-cache")
+ c.Header("Expires", "0")
+ c.Next()
+ }
+}
+
+func registerEmbyRootRoutes(grp *gin.RouterGroup, prefix string, svc *service.Container) {
+ if prefix != "/emby" {
+ return
+ }
+ grp.GET("", embyRootHandler(svc))
+ grp.HEAD("", embyRootHandler(svc))
+ grp.GET("/", embyRootHandler(svc))
+ grp.HEAD("/", embyRootHandler(svc))
+}
+
+func registerEmbyPublicRoutes(grp *gin.RouterGroup, svc *service.Container) {
+ registerEmbyPublicSystemRoutes(grp, svc)
+ registerEmbyPublicSessionRoutes(grp, svc)
+ registerEmbyPublicClientRoutes(grp, svc)
+}
+
+func registerEmbyPublicSystemRoutes(grp *gin.RouterGroup, svc *service.Container) {
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Info/Public", "/system/info/public"}, embySystemInfoPublicHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Info", "/system/info"}, embySystemInfoHandler)
+ registerEmbyGetRoutes(grp, svc, []string{"/System/Endpoint", "/system/endpoint"}, embySystemEndpointHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Ext/ServerDomains", "/system/ext/serverdomains"}, embyServerDomainsHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Configuration/Public", "/system/configuration/public"}, embyPublicServerConfigurationHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/Startup/Configuration", "/startup/configuration"}, embyStartupConfigurationHandler)
+ registerEmbyPostRoutes(grp, svc, []string{"/Startup/Complete", "/startup/complete"}, embyNoContentHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/QuickConnect/Enabled", "/quickconnect/enabled"}, embyQuickConnectEnabledHandler)
+ for _, path := range []string{"/System/Ping", "/system/ping"} {
+ grp.GET(path, embyPingHandler(svc))
+ grp.HEAD(path, embyPingHandler(svc))
+ grp.POST(path, embyPingHandler(svc))
+ }
+}
+
+func registerEmbyPublicSessionRoutes(grp *gin.RouterGroup, svc *service.Container) {
+ registerEmbyPostRoutes(grp, svc, []string{
+ "/Sessions/Capabilities", "/Sessions/Capabilities/Full",
+ "/sessions/capabilities", "/sessions/capabilities/full",
+ }, embyNoContentHandler)
+
+ // 30/min per IP: many Emby clients sit behind a single NAT/reverse-proxy
+ // IP, so a low limit would throttle legitimate logins into 429s.
+ embyLoginLimiter := middleware.NewRateLimiter(30, 1*time.Minute)
+ for _, path := range []string{"/Users/AuthenticateByName", "/Users/authenticatebyname", "/users/AuthenticateByName", "/users/authenticatebyname"} {
+ grp.POST(path, middleware.RateLimit(embyLoginLimiter), embyAuthByNameHandler(svc))
+ }
+
+ registerEmbyGetRoutes(grp, svc, []string{"/Users/Public", "/users/public"}, embyPublicUsersHandler)
+}
+
+func registerEmbyPublicClientRoutes(grp *gin.RouterGroup, svc *service.Container) {
+ registerEmbyGetRoutes(grp, svc, []string{"/Branding/Configuration", "/branding/configuration"}, embyBrandingConfigHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/Branding/Css", "/branding/css"}, embyBrandingCSSHandler)
+ registerEmbyGetRoutes(grp, svc, []string{"/Localization/Options", "/localization/options"}, embyLocalizationOptionsHandler)
+ registerEmbyGetRoutes(grp, svc, []string{"/Localization/Cultures", "/Localization/cultures", "/localization/cultures"}, embyLocalizationCulturesHandler)
+ registerEmbyGetHeadRoutes(grp, svc, []string{"/CustomCssJS/Scripts", "/customcssjs/scripts"}, embyCustomCSSJSScriptsHandler)
+ for _, path := range []string{"/embywebsocket", "/EmbyWebSocket"} {
+ grp.GET(path, embyWebSocketHandler(svc))
+ grp.HEAD(path, embyNoContentHandler(svc))
+ }
+ registerEmbyPostRoutes(grp, svc, []string{"/Sessions/Logout", "/sessions/logout"}, embySessionLogoutHandler)
+ grp.GET("/DisplayPreferences/:id", embyDisplayPreferencesHandler(svc))
+ grp.POST("/DisplayPreferences/:id", embySaveDisplayPreferencesHandler(svc))
+ grp.GET("/displaypreferences/:id", embyDisplayPreferencesHandler(svc))
+ grp.POST("/displaypreferences/:id", embySaveDisplayPreferencesHandler(svc))
+}
+
+func registerEmbyPublicImageRoutes(grp *gin.RouterGroup, svc *service.Container) {
+ // 图片公开(Infuse 缓存 URL 时会丢 token)
+ grp.GET("/Items/:id/Images/:type", embyItemImageHandler(svc))
+ grp.GET("/Items/:id/Images/:type/:index", embyItemImageHandler(svc))
+ grp.HEAD("/Items/:id/Images/:type", embyItemImageHandler(svc))
+ grp.GET("/items/:id/images/:type", embyItemImageHandler(svc))
+ grp.GET("/items/:id/images/:type/:index", embyItemImageHandler(svc))
+ grp.HEAD("/items/:id/images/:type", embyItemImageHandler(svc))
+}
+
+func registerEmbyGetRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
+ for _, path := range paths {
+ grp.GET(path, factory(svc))
+ }
+}
+
+func registerEmbyGetHeadRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
+ for _, path := range paths {
+ grp.GET(path, factory(svc))
+ grp.HEAD(path, factory(svc))
+ }
+}
+
+func registerEmbyPostRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
+ for _, path := range paths {
+ grp.POST(path, factory(svc))
+ }
+}
+
+func registerEmbyAuthenticatedRoutes(auth *gin.RouterGroup, prefix string, svc *service.Container) {
+ registerEmbyAuthenticatedUserRoutes(auth, svc)
+ registerEmbyAuthenticatedItemRoutes(auth, svc)
+ registerEmbyAuthenticatedPlaybackRoutes(auth, prefix, svc)
+ registerEmbyAuthenticatedProgressRoutes(auth, svc)
+ registerEmbyAuthenticatedUserDataRoutes(auth, svc)
+ registerEmbyAuthenticatedSystemRoutes(auth, svc)
+ registerLowercaseEmbyAuthRoutes(auth, svc)
+}
+
+func registerEmbyAuthenticatedUserRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/Users/Me", embyMeHandler(svc))
+ auth.GET("/Users", embyListUsersHandler(svc))
+ auth.GET("/Users/:userId", embyGetUserByIDHandler(svc))
+ auth.GET("/Users/:userId/Views", embyViewsHandler(svc))
+ auth.GET("/Library/MediaFolders", embyViewsHandler(svc))
+ auth.GET("/Library/VirtualFolders", embyVirtualFoldersHandler(svc))
+ auth.GET("/Library/SelectableMediaFolders", embyVirtualFoldersHandler(svc))
+}
+
+func registerEmbyAuthenticatedItemRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/Items", embyItemsHandler(svc))
+ auth.GET("/Users/:userId/Items", embyItemsHandler(svc))
+ auth.GET("/Items/Counts", embyItemsCountsHandler(svc))
+ auth.GET("/Users/:userId/Items/Counts", embyItemsCountsHandler(svc))
+ auth.GET("/Items/Latest", embyLatestItemsHandler(svc))
+ auth.GET("/Items/Resume", embyResumeItemsHandler(svc))
+ auth.GET("/Items/:id", embyItemByIDHandler(svc))
+ auth.GET("/Users/:userId/Items/:id", embyUserItemByIDHandler(svc))
+ auth.GET("/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
+ auth.GET("/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
+ auth.GET("/Users/:userId/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
+ auth.GET("/Users/:userId/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
+ auth.GET("/Shows/NextUp", embyEmptyItemsHandler(svc))
+ auth.GET("/Users/:userId/Shows/NextUp", embyEmptyItemsHandler(svc))
+ auth.GET("/MediaSegments/:id", embyEmptyItemsHandler(svc))
+ auth.GET("/Artists", embyEmptyItemsHandler(svc))
+ auth.GET("/Persons", embyEmptyItemsHandler(svc))
+ auth.GET("/Genres", embyEmptyItemsHandler(svc))
+ auth.GET("/Shows/Upcoming", embyEmptyItemsHandler(svc))
+ auth.GET("/Users/:userId/Shows/Upcoming", embyEmptyItemsHandler(svc))
+ auth.GET("/Items/:id/Similar", embyEmptyItemsHandler(svc))
+ auth.GET("/Items/:id/ThumbnailSet", embyEmptyItemsHandler(svc))
+ auth.GET("/Items/:id/ThemeMedia", embyThemeMediaHandler(svc))
+ auth.GET("/Users/:userId/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
+ auth.GET("/Users/:userId/Items/:id/Intros", embyEmptyItemsHandler(svc))
+ auth.GET("/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
+ auth.GET("/Items/:id/Intros", embyEmptyItemsHandler(svc))
+ auth.GET("/api/danmu/:id/raw", embyDanmuRawHandler(svc))
+}
+
+func registerEmbyAuthenticatedPlaybackRoutes(auth *gin.RouterGroup, prefix string, svc *service.Container) {
+ auth.GET("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
+ auth.POST("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
+ auth.GET("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
+ auth.POST("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
+
+ registerEmbyVideoStreamRoutes(auth, svc, "/Videos")
+ if prefix == "/emby" {
+ auth.GET("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
+ auth.HEAD("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
+ }
+ auth.GET("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.HEAD("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.GET("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.HEAD("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.GET("/Videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
+}
+
+func registerEmbyVideoStreamRoutes(auth *gin.RouterGroup, svc *service.Container, basePath string) {
+ streamHandler := func() gin.HandlerFunc {
+ return embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)
+ }
+ for _, path := range []string{"/:id/stream", "/:id/stream.:container", "/:id/original", "/:id/original.:container"} {
+ fullPath := basePath + path
+ auth.GET(fullPath, streamHandler())
+ auth.HEAD(fullPath, streamHandler())
+ }
+}
+
+func registerEmbyAuthenticatedProgressRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.POST("/Sessions/Playing", embyPlayingProgressHandler(svc))
+ auth.POST("/Sessions/Playing/Progress", embyPlayingProgressHandler(svc))
+ auth.POST("/Sessions/Playing/Stopped", embyPlayingProgressHandler(svc))
+}
+
+func registerEmbyAuthenticatedUserDataRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.POST("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, true))
+ auth.DELETE("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, false))
+ auth.POST("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, true))
+ auth.DELETE("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, false))
+}
+
+func registerEmbyAuthenticatedSystemRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/Sessions", embySessionsHandler(svc))
+ auth.GET("/System/Configuration", embyServerConfigurationHandler(svc))
+ auth.GET("/System/WakeOnLanInfo", embyEmptyArrayHandler(svc))
+ auth.GET("/ScheduledTasks", embyEmptyArrayHandler(svc))
+ auth.GET("/LiveTv/Recordings", embyEmptyItemsHandler(svc))
+ auth.GET("/System/ActivityLog/Entries", embyEmptyItemsHandler(svc))
+ auth.GET("/Web/ConfigurationPages", embyEmptyArrayHandler(svc))
+ auth.POST("/Users/:userId/Configuration", embyNoContentHandler(svc))
+}
diff --git a/internal/handler/emby_routes_lowercase.go b/internal/handler/emby_routes_lowercase.go
new file mode 100644
index 0000000..baabe4a
--- /dev/null
+++ b/internal/handler/emby_routes_lowercase.go
@@ -0,0 +1,94 @@
+package handler
+
+import (
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func registerLowercaseEmbyAuthRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ registerLowercaseEmbyUserRoutes(auth, svc)
+ registerLowercaseEmbyItemRoutes(auth, svc)
+ registerLowercaseEmbyPlaybackRoutes(auth, svc)
+ registerLowercaseEmbyProgressRoutes(auth, svc)
+ registerLowercaseEmbyUserDataRoutes(auth, svc)
+ registerLowercaseEmbySystemRoutes(auth, svc)
+}
+
+func registerLowercaseEmbyUserRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/users/me", embyMeHandler(svc))
+ auth.GET("/users", embyListUsersHandler(svc))
+ auth.GET("/users/:userId", embyGetUserByIDHandler(svc))
+ auth.GET("/users/:userId/views", embyViewsHandler(svc))
+ auth.GET("/library/mediafolders", embyViewsHandler(svc))
+ auth.GET("/library/virtualfolders", embyVirtualFoldersHandler(svc))
+ auth.GET("/library/selectablemediafolders", embyVirtualFoldersHandler(svc))
+}
+
+func registerLowercaseEmbyItemRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/items", embyItemsHandler(svc))
+ auth.GET("/users/:userId/items", embyItemsHandler(svc))
+ auth.GET("/items/counts", embyItemsCountsHandler(svc))
+ auth.GET("/users/:userId/items/counts", embyItemsCountsHandler(svc))
+ auth.GET("/items/latest", embyLatestItemsHandler(svc))
+ auth.GET("/items/resume", embyResumeItemsHandler(svc))
+ auth.GET("/items/:id", embyItemByIDHandler(svc))
+ auth.GET("/users/:userId/items/:id", embyUserItemByIDHandler(svc))
+ auth.GET("/shows/:id/seasons", embyShowSeasonsHandler(svc))
+ auth.GET("/shows/:id/episodes", embyShowEpisodesHandler(svc))
+ auth.GET("/users/:userId/shows/:id/seasons", embyShowSeasonsHandler(svc))
+ auth.GET("/users/:userId/shows/:id/episodes", embyShowEpisodesHandler(svc))
+ auth.GET("/shows/nextup", embyEmptyItemsHandler(svc))
+ auth.GET("/users/:userId/shows/nextup", embyEmptyItemsHandler(svc))
+ auth.GET("/mediasegments/:id", embyEmptyItemsHandler(svc))
+ auth.GET("/artists", embyEmptyItemsHandler(svc))
+ auth.GET("/persons", embyEmptyItemsHandler(svc))
+ auth.GET("/genres", embyEmptyItemsHandler(svc))
+ auth.GET("/shows/upcoming", embyEmptyItemsHandler(svc))
+ auth.GET("/users/:userId/shows/upcoming", embyEmptyItemsHandler(svc))
+ auth.GET("/items/:id/similar", embyEmptyItemsHandler(svc))
+ auth.GET("/items/:id/thumbnailset", embyEmptyItemsHandler(svc))
+ auth.GET("/items/:id/thememedia", embyThemeMediaHandler(svc))
+ auth.GET("/users/:userId/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
+ auth.GET("/users/:userId/items/:id/intros", embyEmptyItemsHandler(svc))
+ auth.GET("/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
+ auth.GET("/items/:id/intros", embyEmptyItemsHandler(svc))
+}
+
+func registerLowercaseEmbyPlaybackRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
+ auth.POST("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
+ auth.GET("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
+ auth.POST("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
+
+ registerEmbyVideoStreamRoutes(auth, svc, "/videos")
+ auth.GET("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.HEAD("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.GET("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.HEAD("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
+ auth.GET("/videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
+}
+
+func registerLowercaseEmbyProgressRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.POST("/sessions/playing", embyPlayingProgressHandler(svc))
+ auth.POST("/sessions/playing/progress", embyPlayingProgressHandler(svc))
+ auth.POST("/sessions/playing/stopped", embyPlayingProgressHandler(svc))
+}
+
+func registerLowercaseEmbyUserDataRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.POST("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, true))
+ auth.DELETE("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, false))
+ auth.POST("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, true))
+ auth.DELETE("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, false))
+}
+
+func registerLowercaseEmbySystemRoutes(auth *gin.RouterGroup, svc *service.Container) {
+ auth.GET("/sessions", embySessionsHandler(svc))
+ auth.GET("/system/configuration", embyServerConfigurationHandler(svc))
+ auth.GET("/system/wakeonlaninfo", embyEmptyArrayHandler(svc))
+ auth.GET("/scheduledtasks", embyEmptyArrayHandler(svc))
+ auth.GET("/livetv/recordings", embyEmptyItemsHandler(svc))
+ auth.GET("/system/activitylog/entries", embyEmptyItemsHandler(svc))
+ auth.GET("/web/configurationpages", embyEmptyArrayHandler(svc))
+ auth.POST("/users/:userId/configuration", embyNoContentHandler(svc))
+}
diff --git a/internal/handler/emby_sessions.go b/internal/handler/emby_sessions.go
new file mode 100644
index 0000000..a52540b
--- /dev/null
+++ b/internal/handler/emby_sessions.go
@@ -0,0 +1,59 @@
+package handler
+
+import (
+ "net/http"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func embySessionsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ if svc.Sessions == nil {
+ c.JSON(http.StatusOK, []any{})
+ return
+ }
+ out := make([]gin.H, 0)
+ for _, sess := range svc.Sessions.List(c.Request.Context()) {
+ last := sess.LastActivityAt
+ itemID := sess.ItemID
+ playState := gin.H{
+ "PositionTicks": sess.PositionTicks,
+ "IsPaused": sess.IsPaused,
+ "PlayMethod": "DirectStream",
+ "CanSeek": true,
+ }
+ row := gin.H{
+ "Id": sess.ID,
+ "ServerId": "mediastation-go-001",
+ "Client": sess.Client,
+ "DeviceId": sess.DeviceID,
+ "DeviceName": sess.DeviceName,
+ "UserId": sess.UserID,
+ "UserName": sess.UserName,
+ "LastActivityDate": last,
+ "RemoteEndPoint": sess.RemoteEndPoint,
+ "PlayState": playState,
+ "SupportsRemoteControl": true,
+ }
+ if itemID != "" && sess.IsPlaying {
+ row["NowPlayingItem"] = gin.H{"Id": itemID}
+ }
+ out = append(out, row)
+ }
+ c.Header("Cache-Control", "no-store")
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embySessionLogoutHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ if svc.Sessions != nil {
+ uid := embyUserID(c)
+ clientInfo := embyClientInfoFromRequest(c)
+ svc.Sessions.Logout(c.Request.Context(), uid, clientInfo.DeviceID, c.ClientIP())
+ }
+ c.Status(http.StatusNoContent)
+ }
+}
diff --git a/internal/handler/emby_sessions_test.go b/internal/handler/emby_sessions_test.go
new file mode 100644
index 0000000..bd4d485
--- /dev/null
+++ b/internal/handler/emby_sessions_test.go
@@ -0,0 +1,47 @@
+package handler
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/gin-gonic/gin"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestEmbySessionsReturnsRealtimeSession(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ tracker := service.NewSessionTrackerService(zap.NewNop())
+ tracker.RecordPlayback(t.Context(), "user-1", "viewer", "dev-1", "Apple TV", "Yamby", "10.0.0.8", "media-1", 1000, 2000, false)
+ svc := &service.Container{Sessions: tracker}
+ router := gin.New()
+ router.GET("/Sessions", embySessionsHandler(svc))
+
+ req := httptest.NewRequest(http.MethodGet, "/Sessions", nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ if got := w.Header().Get("Cache-Control"); got != "no-store" {
+ t.Fatalf("cache-control = %q, want no-store", got)
+ }
+ var rows []map[string]any
+ if err := json.Unmarshal(w.Body.Bytes(), &rows); err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 1 {
+ t.Fatalf("sessions = %d, want 1: %s", len(rows), w.Body.String())
+ }
+ if rows[0]["UserId"] != "user-1" || rows[0]["DeviceId"] != "dev-1" || rows[0]["Client"] != "Yamby" {
+ t.Fatalf("session payload = %#v", rows[0])
+ }
+ if _, err := time.Parse(time.RFC3339Nano, rows[0]["LastActivityDate"].(string)); err != nil {
+ t.Fatalf("LastActivityDate should be RFC3339 time, got %#v", rows[0]["LastActivityDate"])
+ }
+}
diff --git a/internal/handler/emby_static.go b/internal/handler/emby_static.go
new file mode 100644
index 0000000..20740d8
--- /dev/null
+++ b/internal/handler/emby_static.go
@@ -0,0 +1,189 @@
+package handler
+
+import (
+ "net/http"
+ "time"
+
+ "github.com/gin-gonic/gin"
+ "github.com/gorilla/websocket"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func embyNoContentHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Status(http.StatusNoContent)
+ }
+}
+
+func embyWebSocketHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ if !websocket.IsWebSocketUpgrade(c.Request) {
+ c.Status(http.StatusNoContent)
+ return
+ }
+ conn, err := wsUpgrader.Upgrade(c.Writer, c.Request, nil)
+ if err != nil {
+ return
+ }
+ defer conn.Close()
+
+ done := make(chan struct{})
+ go func() {
+ defer close(done)
+ for {
+ if _, _, err := conn.NextReader(); err != nil {
+ return
+ }
+ }
+ }()
+
+ ticker := time.NewTicker(30 * time.Second)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-done:
+ return
+ case <-ticker.C:
+ _ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
+ if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
+ return
+ }
+ }
+ }
+ }
+}
+
+func embyServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "EnableFolderView": true,
+ "EnableGroupingIntoCollections": true,
+ "EnableExternalContentInSuggestions": false,
+ "ImageSavingConvention": "Compatible",
+ })
+ }
+}
+
+func embyPublicServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "IsStartupWizardCompleted": true,
+ "EnableRemoteAccess": true,
+ "EnableUPnP": false,
+ "EnableHttps": false,
+ "RequireHttps": false,
+ "LocalNetworkSubnets": []string{},
+ "LocalNetworkAddresses": []string{},
+ "RemoteClientBitrateLimit": 0,
+ })
+ }
+}
+
+func embyStartupConfigurationHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "IsStartupWizardCompleted": true,
+ "StartupWizardCompleted": true,
+ "EnableRemoteAccess": true,
+ "UICulture": "zh-CN",
+ "MetadataCountryCode": "CN",
+ "PreferredMetadataLanguage": "zh-CN",
+ })
+ }
+}
+
+func embyQuickConnectEnabledHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, false)
+ }
+}
+
+func embyEmptyItemsHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"Items": []any{}, "TotalRecordCount": 0})
+ }
+}
+
+func embyEmptyArrayHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, []any{})
+ }
+}
+
+func embyCustomCSSJSScriptsHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Data(http.StatusOK, "application/javascript; charset=utf-8", nil)
+ }
+}
+
+func embyLocalizationCulturesHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, []gin.H{
+ {
+ "DisplayName": "简体中文",
+ "Name": "zh-CN",
+ "ThreeLetterISOLanguageName": "zho",
+ "TwoLetterISOLanguageName": "zh",
+ "ThreeLetterISOLanguageNames": []string{"zho", "chi"},
+ "IsRightToLeft": false,
+ },
+ {
+ "DisplayName": "English",
+ "Name": "en-US",
+ "ThreeLetterISOLanguageName": "eng",
+ "TwoLetterISOLanguageName": "en",
+ "ThreeLetterISOLanguageNames": []string{"eng"},
+ "IsRightToLeft": false,
+ },
+ })
+ }
+}
+
+func embyThemeMediaHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ empty := gin.H{"Items": []any{}, "TotalRecordCount": 0}
+ c.JSON(http.StatusOK, gin.H{
+ "ThemeVideosResult": empty,
+ "ThemeSongsResult": empty,
+ "SoundtrackSongsResult": empty,
+ })
+ }
+}
+
+func embyServerDomainsHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, []any{})
+ }
+}
+
+func embyDanmuRawHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Data(http.StatusOK, "text/plain; charset=utf-8", nil)
+ }
+}
+
+func embyBrandingConfigHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "LoginDisclaimer": "",
+ "CustomCss": "",
+ "SplashscreenEnabled": false,
+ })
+ }
+}
+
+func embyBrandingCSSHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Data(http.StatusOK, "text/css; charset=utf-8", []byte(""))
+ }
+}
+
+func embyLocalizationOptionsHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, []map[string]any{
+ {"Name": "简体中文", "Value": "zh-CN"},
+ {"Name": "English", "Value": "en-US"},
+ })
+ }
+}
diff --git a/internal/handler/emby_system.go b/internal/handler/emby_system.go
new file mode 100644
index 0000000..c70c024
--- /dev/null
+++ b/internal/handler/emby_system.go
@@ -0,0 +1,98 @@
+package handler
+
+import (
+ "net/http"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func embySystemInfoHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfo()))
+ }
+}
+
+func embySystemInfoPublicHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfoPublic()))
+ }
+}
+
+func embyRequestBaseURL(c *gin.Context) string {
+ proto := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto"))
+ if proto == "" {
+ if c.Request != nil && c.Request.TLS != nil {
+ proto = "https"
+ } else {
+ proto = "http"
+ }
+ }
+ if comma := strings.Index(proto, ","); comma >= 0 {
+ proto = strings.TrimSpace(proto[:comma])
+ }
+
+ host := strings.TrimSpace(c.GetHeader("X-Forwarded-Host"))
+ if host == "" && c.Request != nil {
+ host = strings.TrimSpace(c.Request.Host)
+ }
+ if host == "" {
+ return ""
+ }
+ return strings.TrimRight(proto+"://"+host, "/")
+}
+
+func embyWithRequestAddress(c *gin.Context, payload map[string]any) map[string]any {
+ out := make(map[string]any, len(payload)+2)
+ for key, value := range payload {
+ out[key] = value
+ }
+ if address := embyRequestBaseURL(c); address != "" {
+ out["LocalAddress"] = address
+ out["WanAddress"] = address
+ out["PublishedServerUrl"] = address
+ }
+ return out
+}
+
+func embySystemEndpointHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "IsLocal": true,
+ "IsInNetwork": true,
+ })
+ }
+}
+
+func embyPingHandler(_ *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ // Emby/Jellyfin 期望 plain text "Emby Server"
+ c.String(http.StatusOK, "Emby Server")
+ }
+}
+
+func embyRootHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.JSON(http.StatusOK, embyPublicSystemInfoPayload(c, svc))
+ }
+}
+
+func embyPublicSystemInfoPayload(c *gin.Context, svc *service.Container) map[string]any {
+ if svc != nil && svc.Emby != nil {
+ return embyWithRequestAddress(c, svc.Emby.SystemInfoPublic())
+ }
+ return embyWithRequestAddress(c, map[string]any{
+ "Id": "mediastation-go-001",
+ "ServerId": "mediastation-go-001",
+ "ServerName": "MediaStationGo",
+ "Version": "4.8.10.0",
+ "ServerVersion": "4.8.10.0",
+ "ProductName": "Emby Server",
+ "OperatingSystem": "Windows",
+ "SupportsHttps": false,
+ "SupportsAutoDiscovery": true,
+ "StartupWizardCompleted": true,
+ })
+}
diff --git a/internal/handler/emby_test.go b/internal/handler/emby_test.go
index f78d07a..5207b3c 100644
--- a/internal/handler/emby_test.go
+++ b/internal/handler/emby_test.go
@@ -5,8 +5,6 @@ import (
"encoding/json"
"net/http"
"net/http/httptest"
- "os"
- "path/filepath"
"strings"
"testing"
"time"
@@ -25,154 +23,6 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/service"
)
-func TestParseEmbyAuthByNameReqAcceptsLowercaseJSON(t *testing.T) {
- gin.SetMode(gin.TestMode)
- w := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(w)
- c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader(`{"username":"alice","password":"secret"}`))
- c.Request.Header.Set("Content-Type", "application/json")
-
- req, err := parseEmbyAuthByNameReq(c)
- if err != nil {
- t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
- }
- if req.Username != "alice" || req.Password != "secret" {
- t.Fatalf("unexpected request: %#v", req)
- }
-}
-
-func TestParseEmbyAuthByNameReqAcceptsFormBody(t *testing.T) {
- gin.SetMode(gin.TestMode)
- w := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(w)
- c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader("Username=bob&Pw=secret"))
- c.Request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
-
- req, err := parseEmbyAuthByNameReq(c)
- if err != nil {
- t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
- }
- if req.Username != "bob" || req.Pw != "secret" {
- t.Fatalf("unexpected request: %#v", req)
- }
-}
-
-func TestParseEmbyAuthByNameReqAcceptsJSONWithoutContentType(t *testing.T) {
- gin.SetMode(gin.TestMode)
- w := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(w)
- c.Request = httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"UserName":"carol","PW":"secret"}`))
-
- req, err := parseEmbyAuthByNameReq(c)
- if err != nil {
- t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
- }
- if req.Username != "carol" || req.Pw != "secret" {
- t.Fatalf("unexpected request: %#v", req)
- }
-}
-
-func TestEmbyAuthenticateByNameAcceptsCaseVariantUsernameAndPath(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.Setting{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- cfg := &config.Config{}
- cfg.Secrets.JWTSecret = "test-secret"
- log := zap.NewNop()
- permissions := service.NewPermissionService(log, repos)
- auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
- if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
- t.Fatalf("register: %v", err)
- }
-
- router := gin.New()
- registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
- Repo: repos,
- Auth: auth,
- Emby: service.NewEmbyService(cfg, log, repos),
- Audit: service.NewAuditService(log, repos),
- })
-
- req := httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"Username":"Viewer","Pw":"secret-pass"}`))
- req.Header.Set("Content-Type", "application/json")
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- var payload map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
- t.Fatalf("decode response: %v", err)
- }
- if payload["AccessToken"] == "" {
- t.Fatalf("missing AccessToken: %#v", payload)
- }
-}
-
-func TestEmbyAuthenticateRecordsMediaBrowserClientInfo(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if sqlDB, err := db.DB(); err == nil {
- sqlDB.SetMaxOpenConns(1)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- cfg := &config.Config{}
- cfg.Secrets.JWTSecret = "test-secret"
- log := zap.NewNop()
- permissions := service.NewPermissionService(log, repos)
- auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
- if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
- t.Fatalf("register: %v", err)
- }
-
- router := gin.New()
- registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
- Repo: repos,
- Auth: auth,
- Emby: service.NewEmbyService(cfg, log, repos),
- Device: service.NewDeviceService(log, repos),
- Audit: service.NewAuditService(log, repos),
- Permissions: permissions,
- })
-
- req := httptest.NewRequest(http.MethodPost, "/emby/Users/AuthenticateByName", strings.NewReader(`{"Username":"viewer","Pw":"secret-pass"}`))
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", DeviceId="device-42"`)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- user, err := repos.User.FindByUsername(context.Background(), "viewer")
- if err != nil {
- t.Fatalf("find user: %v", err)
- }
- devices, err := repos.UserDevice.ListByUser(context.Background(), user.ID)
- if err != nil {
- t.Fatalf("list devices: %v", err)
- }
- if len(devices) != 1 {
- t.Fatalf("devices = %#v, want one recorded device", devices)
- }
- if devices[0].DeviceID != "device-42" || devices[0].DeviceName != "PC" || devices[0].Client != "Infuse" {
- t.Fatalf("device info not parsed from MediaBrowser header: %#v", devices[0])
- }
-}
-
func TestEmbyMarkPlayedRefreshesPlaybackDevice(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -762,878 +612,6 @@ func TestEmbyWebSocketRouteUpgradesForOfficialClients(t *testing.T) {
}
}
-func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
-
- posterPath := filepath.Join(t.TempDir(), "poster.png")
- if err := os.WriteFile(posterPath, []byte{
- 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
- 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
- 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
- 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
- 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
- 0x54, 0x78, 0x9c, 0x63, 0x00, 0x01, 0x00, 0x00,
- 0x05, 0x00, 0x01, 0x0d, 0x0a, 0x2d, 0xb4, 0x00,
- 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae,
- 0x42, 0x60, 0x82,
- }, 0o644); err != nil {
- t.Fatalf("write poster: %v", err)
- }
-
- repos := repository.New(db)
- cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- Title: "Poster Test",
- Path: "D:\\media\\poster-test.mp4",
- PosterURL: posterPath,
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- router := gin.New()
- registerEmbyRoutes(router, "test-secret", &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
- ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/Items/media-1/Images/Primary", nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- if location := w.Header().Get("Location"); location != "" {
- t.Fatalf("expected direct image response, got redirect to %q", location)
- }
- if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
- t.Fatalf("expected png content type, got %q", contentType)
- }
-}
-
-func TestEmbyMissingItemImageReturnsTransparentPlaceholder(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
- router := gin.New()
- registerEmbyRoutes(router, "test-secret", &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
- ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
- })
-
- req := httptest.NewRequest(http.MethodHead, "/Items/missing/Images/Primary", nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("expected placeholder status 200, got %d body=%s", w.Code, w.Body.String())
- }
- if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
- t.Fatalf("expected png content type, got %q", contentType)
- }
- if length := w.Header().Get("Content-Length"); length == "" || length == "0" {
- t.Fatalf("expected placeholder content length, got %q", length)
- }
-}
-
-func TestEmbyUserItemByIDRouteReturnsJSON(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "episode-1"},
- LibraryID: lib.ID,
- Title: "Test Show",
- Path: "D:\\media\\tv\\Test Show\\Season 01\\Test Show - S01E01.mkv",
- SeasonNum: 1,
- EpisodeNum: 1,
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/episode-1", nil)
- req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
- req.Header.Set("If-None-Match", `"stale-client-cache"`)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "application/json") {
- t.Fatalf("expected JSON content type, got %q body=%s", contentType, w.Body.String())
- }
- var item map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
- t.Fatalf("decode item: %v", err)
- }
- if item["Id"] != "episode-1" || item["Type"] != "Episode" {
- t.Fatalf("unexpected item payload: %#v", item)
- }
-}
-
-func TestEmbyUserItemByIDRouteReturnsLibraryView(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Base: model.Base{ID: "lib-tv"}, Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/lib-tv", nil)
- req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- var item map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
- t.Fatalf("decode item: %v", err)
- }
- if item["Id"] != "lib-tv" || item["Type"] != "CollectionFolder" || item["CollectionType"] != "tvshows" {
- t.Fatalf("unexpected library payload: %#v", item)
- }
-}
-
-func TestEmbyLowercasePlaybackInfoRouteReturnsJSON(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "电影", Path: t.TempDir(), Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- LibraryID: lib.ID,
- Title: "Lowercase Playback",
- Path: filepath.Join(lib.Path, "lowercase-playback.mp4"),
- Container: "mp4",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/media-1/playbackinfo", nil)
- req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- var body map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
- t.Fatalf("decode playback info: %v", err)
- }
- if _, ok := body["MediaSources"]; !ok {
- t.Fatalf("missing MediaSources: %#v", body)
- }
- sources, ok := body["MediaSources"].([]any)
- if !ok || len(sources) == 0 {
- t.Fatalf("unexpected MediaSources: %#v", body["MediaSources"])
- }
- source, ok := sources[0].(map[string]any)
- if !ok {
- t.Fatalf("unexpected MediaSource: %#v", sources[0])
- }
- directURL, _ := source["DirectStreamUrl"].(string)
- if !strings.Contains(directURL, "api_key=") {
- t.Fatalf("DirectStreamUrl should carry api_key for clients that do not repeat auth headers: %#v", source)
- }
- transcodeURL, _ := source["TranscodingUrl"].(string)
- if transcodeURL != "" && !strings.Contains(transcodeURL, "api_key=") {
- t.Fatalf("TranscodingUrl should carry api_key: %#v", source)
- }
-}
-
-func TestEmbyPlaybackInfoDoesNotExposeTokenInCloudPath(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: "cloud://openlist/Movies/Movie.mkv",
- STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/cloud-1/playbackinfo", nil)
- req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- var body map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
- t.Fatalf("decode playback info: %v", err)
- }
- source := body["MediaSources"].([]any)[0].(map[string]any)
- pathURL, _ := source["Path"].(string)
- if pathURL != "/api/stream/cloud-1" {
- t.Fatalf("cloud Path should stay as non-tokenized display stream URL, got %#v", source)
- }
- if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
- t.Fatalf("cloud Path must not expose auth key/token: %#v", source)
- }
- if strings.Contains(pathURL, "/api/cloud/play/") {
- t.Fatalf("cloud Path should not expose naked cloud play URL: %#v", source)
- }
- directURL, _ := source["DirectStreamUrl"].(string)
- if !strings.HasPrefix(directURL, "/api/stream/cloud-1") || !strings.Contains(directURL, "api_key=") {
- t.Fatalf("DirectStreamUrl should stay tokenized: %#v", source)
- }
- if source["SupportsDirectPlay"] != true {
- t.Fatalf("cloud media should advertise DirectPlay when tokenized Path is playable: %#v", source)
- }
- if source["SupportsTranscoding"] != false {
- t.Fatalf("cloud media should not advertise host transcoding: %#v", source)
- }
-}
-
-func TestEmbyItemsDoNotExposeTokenInEmbeddedCloudPath(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: "cloud://openlist/Movies/Movie.mkv",
- STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- token := signedTestToken(t, secret)
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/emby/Users/user-1/Items?IncludeItemTypes=Movie&Recursive=true&Limit=5&X-Emby-Token="+token, nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- var body map[string]any
- if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
- t.Fatalf("decode items: %v", err)
- }
- items := body["Items"].([]any)
- if len(items) != 1 {
- t.Fatalf("unexpected items: %#v", body["Items"])
- }
- source := items[0].(map[string]any)["MediaSources"].([]any)[0].(map[string]any)
- pathURL, _ := source["Path"].(string)
- if pathURL != "/api/stream/cloud-1" {
- t.Fatalf("embedded cloud Path should stay as non-tokenized display stream URL, got %#v", source)
- }
- if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
- t.Fatalf("embedded cloud Path must not expose auth key/token: %#v", source)
- }
-}
-
-func TestEmbyVideoStreamUsesSTRMWhenRedirectProxyDisabled(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
- t.Fatalf("enable strm playback: %v", err)
- }
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
- t.Fatalf("disable redirect playback: %v", err)
- }
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: "cloud://openlist/Movies/Movie.mkv",
- STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
- Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
- })
-
- token := signedTestToken(t, secret)
- req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream?api_key="+token, nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusFound {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- loc := w.Header().Get("Location")
- if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
- t.Fatalf("STRM mode should redirect /Videos fallback to tokenized /api/stream, got %q", loc)
- }
- if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
- t.Fatalf("STRM fallback redirect Cache-Control = %q, want no-store", got)
- }
- if strings.Contains(loc, "/api/cloud/play/") {
- t.Fatalf("STRM mode should not expose cloud play directly from /Videos fallback: %q", loc)
- }
-}
-
-func TestEmbyVideoStreamIssuesTokenForSessionFallbackSTRMRedirect(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
- t.Fatalf("enable strm playback: %v", err)
- }
- if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
- t.Fatalf("disable redirect playback: %v", err)
- }
- user := model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }
- if err := repos.User.Create(t.Context(), &user); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: "cloud://openlist/Movies/Movie.mkv",
- STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
- svc := &service.Container{
- Repo: repos,
- Auth: service.NewAuthService(cfg, zap.NewNop(), repos, nil, nil),
- Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
- Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
- }
- router := gin.New()
- router.GET("/videos/:id/stream", func(c *gin.Context) {
- c.Set(middleware.CtxUserID, user.ID)
- c.Set(middleware.CtxUserRole, user.Role)
- embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)(c)
- })
-
- req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusFound {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- loc := w.Header().Get("Location")
- if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
- t.Fatalf("session fallback redirect should include api_key for /api/stream, got %q", loc)
- }
-}
-
-func TestEmbyLowercaseVideoStreamRouteServesMedia(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- dir := t.TempDir()
- mediaPath := filepath.Join(dir, "sample.mp4")
- if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
- t.Fatalf("write media: %v", err)
- }
- lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- LibraryID: lib.ID,
- Title: "Lowercase Stream",
- Path: mediaPath,
- Container: "mp4",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/videos/media-1/stream?api_key="+signedTestToken(t, secret), nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- if got := w.Body.String(); got != "fake-video-bytes" {
- t.Fatalf("unexpected stream body: %q", got)
- }
-}
-
-func TestEmbyPrefixedAPIStreamRouteServesMedia(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- dir := t.TempDir()
- mediaPath := filepath.Join(dir, "sample.mp4")
- if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
- t.Fatalf("write media: %v", err)
- }
- lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- LibraryID: lib.ID,
- Title: "Prefixed API Stream",
- Path: mediaPath,
- Container: "mp4",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/emby/api/stream/media-1?api_key="+signedTestToken(t, secret), nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- if got := w.Body.String(); got != "fake-video-bytes" {
- t.Fatalf("unexpected stream body: %q", got)
- }
-}
-
-func TestEmbyVideoStreamRedirectKeepsMediaBrowserAuthorizationToken(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: "cloud://openlist/Movies/Movie.mkv",
- STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
- Container: "mkv",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
- })
-
- token := signedTestToken(t, secret)
- req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
- req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", Token="`+token+`"`)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusFound {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- loc := w.Header().Get("Location")
- if !strings.Contains(loc, "/api/cloud/play/openlist?") || !strings.Contains(loc, "token=") {
- t.Fatalf("redirect Location should target tokenized cloud play endpoint, got %q", loc)
- }
-}
-
-func TestEmbyLowercaseOriginalHeadRouteServesHeaders(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- dir := t.TempDir()
- mediaPath := filepath.Join(dir, "sample.mp4")
- if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
- t.Fatalf("write media: %v", err)
- }
- lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- LibraryID: lib.ID,
- Title: "Lowercase Original",
- Path: mediaPath,
- Container: "mp4",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
- })
-
- req := httptest.NewRequest(http.MethodHead, "/videos/media-1/original.mp4?api_key="+signedTestToken(t, secret), nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code != http.StatusOK {
- t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
- }
- if w.Body.Len() != 0 {
- t.Fatalf("HEAD response should not include body, got %q", w.Body.String())
- }
-}
-
-func TestEmbyLowercaseVideoHLSRouteDoesNot404WhenDirectOnly(t *testing.T) {
- gin.SetMode(gin.TestMode)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatalf("migrate: %v", err)
- }
- repos := repository.New(db)
- if err := repos.User.Create(t.Context(), &model.User{
- Base: model.Base{ID: "user-1"},
- Username: "tester",
- PasswordHash: "x",
- Role: "admin",
- Tier: "plus",
- IsActive: true,
- }); err != nil {
- t.Fatalf("create user: %v", err)
- }
- dir := t.TempDir()
- mediaPath := filepath.Join(dir, "sample.mp4")
- if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
- t.Fatalf("write media: %v", err)
- }
- lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := db.Create(&model.Media{
- Base: model.Base{ID: "media-1"},
- LibraryID: lib.ID,
- Title: "Lowercase HLS",
- Path: mediaPath,
- Container: "mp4",
- }).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
- if err := repos.Setting.Set(t.Context(), service.PlaybackDirectOnlySettingKey, "true"); err != nil {
- t.Fatalf("set direct-only: %v", err)
- }
-
- const secret = "test-secret"
- router := gin.New()
- registerEmbyRoutes(router, secret, &service.Container{
- Repo: repos,
- Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
- Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
- })
-
- req := httptest.NewRequest(http.MethodGet, "/videos/media-1/master.m3u8?api_key="+signedTestToken(t, secret), nil)
- w := httptest.NewRecorder()
- router.ServeHTTP(w, req)
-
- if w.Code == http.StatusNotFound {
- t.Fatalf("lowercase HLS route should be registered, got 404")
- }
- if w.Code != http.StatusConflict {
- t.Fatalf("direct-only HLS should return 409, got %d body=%s", w.Code, w.Body.String())
- }
-}
-
func signedTestToken(t *testing.T, secret string) string {
t.Helper()
claims := middleware.Claims{
diff --git a/internal/handler/emby_users.go b/internal/handler/emby_users.go
new file mode 100644
index 0000000..d99ee04
--- /dev/null
+++ b/internal/handler/emby_users.go
@@ -0,0 +1,174 @@
+package handler
+
+import (
+ "net/http"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+// ─── Users / Auth ────────────────────────────────────────────────────────────
+
+// embyAuthByNameHandler 处理 POST /Users/AuthenticateByName。
+//
+// 这是 Emby 客户端登录的唯一入口(Infuse / Yamby / Hills 等都走这里)。
+// 用户名+密码 → 调用我们已有的 AuthService.Login → 返回 AccessToken + User。
+func embyAuthByNameHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ req, err := parseEmbyAuthByNameReq(c)
+ if err != nil {
+ embyError(c, http.StatusBadRequest, "invalid body")
+ return
+ }
+ password := req.Pw
+ if password == "" {
+ password = req.Password
+ }
+ if strings.TrimSpace(req.Username) == "" || password == "" {
+ if req.PasswordMd5 != "" || req.PasswordSha1 != "" {
+ embyError(c, http.StatusBadRequest, "plain password required")
+ return
+ }
+ embyError(c, http.StatusBadRequest, "missing username or password")
+ return
+ }
+ resp, err := svc.Auth.Login(c.Request.Context(), req.Username, password)
+ if err != nil {
+ embyError(c, http.StatusUnauthorized, err.Error())
+ return
+ }
+ // 记录登录设备会话并执行防共享检测(登录客户端数 / 设备指纹)。
+ clientInfo := embyClientInfoFromRequest(c)
+ if svc.Sessions != nil {
+ svc.Sessions.RecordLogin(c.Request.Context(), resp.User.ID, resp.User.Username,
+ clientInfo.DeviceID,
+ clientInfo.DeviceName,
+ clientInfo.Client,
+ c.ClientIP())
+ }
+ if svc.Device != nil {
+ svc.Device.RecordLogin(c.Request.Context(), resp.User.ID,
+ clientInfo.DeviceID,
+ clientInfo.DeviceName,
+ clientInfo.Client,
+ c.ClientIP())
+ }
+ userPayload, _ := svc.Emby.FindUser(c.Request.Context(), resp.User.ID)
+ // Emby/Jellyfin 客户端没有 refresh token 机制:它们把这里返回的
+ // AccessToken 长期保存并反复使用。若返回 60 分钟的普通 access
+ // token,客户端每小时就会掉登录、无法播放、媒体库无法刷新。因此
+ // 签发长期令牌(IssueEmbyToken)匹配 Emby 持久化令牌语义。
+ accessToken := resp.Tokens.AccessToken
+ if longLived, err := svc.Auth.IssueEmbyToken(resp.User); err == nil && longLived != "" {
+ accessToken = longLived
+ }
+ embyRememberCompatSession(c, accessToken)
+ c.JSON(http.StatusOK, gin.H{
+ "AccessToken": accessToken,
+ "ServerId": "mediastation-go-001",
+ "User": userPayload,
+ "SessionInfo": gin.H{
+ "Id": resp.User.ID,
+ "UserId": resp.User.ID,
+ "UserName": resp.User.Username,
+ "Client": clientInfo.Client,
+ "DeviceId": clientInfo.DeviceID,
+ "DeviceName": clientInfo.DeviceName,
+ },
+ })
+ }
+}
+
+func embyPublicUsersHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ // 公开用户列表(Emby Web 客户端登录页拉这个,列出可见用户)。
+ users, err := svc.Emby.ListUsers(c.Request.Context())
+ if err != nil {
+ c.JSON(http.StatusOK, []any{})
+ return
+ }
+ // 公开版本只暴露 Id + Name,不包含 Policy。
+ out := make([]map[string]any, 0, len(users))
+ for _, u := range users {
+ out = append(out, map[string]any{
+ "Id": u["Id"],
+ "Name": u["Name"],
+ "ServerId": u["ServerId"],
+ "HasPassword": true,
+ })
+ }
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyListUsersHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ users, err := svc.Emby.ListUsers(c.Request.Context())
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, users)
+ }
+}
+
+func embyMeHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := embyUserID(c)
+ if uid == "" {
+ embyError(c, http.StatusUnauthorized, "not authenticated")
+ return
+ }
+ u, err := svc.Emby.FindUser(c.Request.Context(), uid)
+ if err != nil || u == nil {
+ embyError(c, http.StatusNotFound, "user not found")
+ return
+ }
+ c.JSON(http.StatusOK, u)
+ }
+}
+
+func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ u, err := svc.Emby.FindUser(c.Request.Context(), c.Param("userId"))
+ if err == nil && u != nil {
+ c.JSON(http.StatusOK, u)
+ return
+ }
+ if authUID := embyUserID(c); authUID != "" && authUID != c.Param("userId") {
+ u, err = svc.Emby.FindUser(c.Request.Context(), authUID)
+ if err == nil && u != nil {
+ c.JSON(http.StatusOK, u)
+ return
+ }
+ }
+ c.JSON(http.StatusOK, embyFallbackUser(c.Param("userId")))
+ }
+}
+
+func embyFallbackUser(id string) gin.H {
+ if strings.TrimSpace(id) == "" {
+ id = "mediastation-user"
+ }
+ return gin.H{
+ "Id": id,
+ "Name": "MediaStationGo",
+ "ServerId": "mediastation-go-001",
+ "HasPassword": true,
+ "HasConfiguredPassword": true,
+ "HasConfiguredEasyPassword": false,
+ "EnableAutoLogin": false,
+ "Policy": gin.H{
+ "IsAdministrator": true,
+ "EnableContentDeletion": true,
+ "EnableRemoteControlOfOtherUsers": true,
+ "EnableSharedDeviceControl": true,
+ "EnableRemoteAccess": true,
+ "EnableAllDevices": true,
+ "EnableAllChannels": true,
+ "EnableAllFolders": true,
+ },
+ }
+}
diff --git a/internal/handler/emby_views_handlers.go b/internal/handler/emby_views_handlers.go
new file mode 100644
index 0000000..8553dfa
--- /dev/null
+++ b/internal/handler/emby_views_handlers.go
@@ -0,0 +1,63 @@
+package handler
+
+import (
+ "net/http"
+
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func embyViewsHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ uid := c.Param("userId")
+ if uid == "" {
+ uid = embyUserID(c)
+ }
+ out, err := svc.Emby.Views(c.Request.Context(), uid)
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ embyAttachRequestTokenToMediaSources(c, out)
+ c.JSON(http.StatusOK, out)
+ }
+}
+
+func embyVirtualFoldersHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ c.Header("Cache-Control", "no-store")
+ libs, err := svc.Repo.Library.List(c.Request.Context())
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
+ uid := embyUserID(c)
+ visibility := service.UserDefaultMediaVisibility(c.Request.Context(), svc.Repo, uid)
+ out := make([]gin.H, 0, len(libs))
+ for _, lib := range libs {
+ if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
+ continue
+ }
+ collectionType := "movies"
+ switch lib.Type {
+ case "tv", "anime", "variety":
+ collectionType = "tvshows"
+ case "music":
+ collectionType = "music"
+ }
+ out = append(out, gin.H{
+ "Name": lib.Name,
+ "Locations": []string{lib.Path},
+ "CollectionType": collectionType,
+ "ItemId": lib.ID,
+ "Id": lib.ID,
+ "PrimaryImageItemId": lib.ID,
+ "RefreshStatus": "Idle",
+ "LibraryOptions": gin.H{},
+ })
+ }
+ c.JSON(http.StatusOK, out)
+ }
+}
diff --git a/internal/handler/manual_scrape.go b/internal/handler/manual_scrape.go
index 316b984..3b0f19b 100644
--- a/internal/handler/manual_scrape.go
+++ b/internal/handler/manual_scrape.go
@@ -12,8 +12,20 @@ import (
)
type manualScrapeApplyReq struct {
- MediaIDs []string `json:"media_ids"`
- Match service.ManualScrapeRequest `json:"match"`
+ MediaIDs []string `json:"media_ids"`
+ Match service.ManualScrapeRequest `json:"match"`
+ EpisodeArtwork *bool `json:"episode_artwork"`
+ EpisodeImages *bool `json:"episode_images"`
+}
+
+func (r manualScrapeApplyReq) episodeArtworkOption() *bool {
+ if r.EpisodeImages != nil {
+ return r.EpisodeImages
+ }
+ if r.EpisodeArtwork != nil {
+ return r.EpisodeArtwork
+ }
+ return r.Match.EpisodeArtworkOption()
}
const manualScrapeApplyTimeout = 5 * time.Minute
@@ -72,10 +84,11 @@ func manualScrapeApplyBatchHandler(svc *service.Container) gin.HandlerFunc {
}
applyCtx, cancel := manualScrapeApplyContext(c)
defer cancel()
+ options := service.ScrapeOptions{EpisodeArtwork: req.episodeArtworkOption()}
applied := 0
errorsOut := make([]string, 0)
for _, id := range ids {
- if _, err := svc.Scraper.ApplyManualMatch(applyCtx, id, req.Match); err != nil {
+ if _, err := svc.Scraper.ApplyManualMatchWithOptions(applyCtx, id, req.Match, options); err != nil {
errorsOut = append(errorsOut, id+": "+err.Error())
continue
}
diff --git a/internal/handler/media.go b/internal/handler/media.go
index 5b27eb9..c22ccc8 100644
--- a/internal/handler/media.go
+++ b/internal/handler/media.go
@@ -27,6 +27,7 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
+ libs = service.FilterDeprecatedNativeCloudLibraries(libs)
role, _ := c.Get(middleware.CtxUserRole)
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
if !includeHidden {
@@ -109,7 +110,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
})
}
- _, _, _ = svc.Scan.StartCloudLibraryScan(id, false)
+ _, _, _ = svc.Scan.StartCloudLibraryScan(id, true)
finishHTTPTask(task, nil, "queued", "云盘扫描已加入后台队列", map[string]int64{"queued": 1}, nil)
c.JSON(http.StatusAccepted, gin.H{
"library_id": id,
@@ -119,7 +120,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
"probed": 0,
"queued": true,
"cloud": true,
- "message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库",
+ "message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库;若已开启自动刮削,会在扫描后补齐元数据",
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
})
return
diff --git a/internal/handler/media_favorite.go b/internal/handler/media_favorite.go
index 3ef9782..9dbeddc 100644
--- a/internal/handler/media_favorite.go
+++ b/internal/handler/media_favorite.go
@@ -80,16 +80,22 @@ func listFavoritesAliasHandler(svc *service.Container) gin.HandlerFunc {
// path; the AI hint comes from svc.AI when configured.
func aiScrapeMediaHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ options, err := scrapeOptionsFromRequest(c, false)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
+ return
+ }
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
return
}
- if err := svc.Scraper.EnrichOne(c.Request.Context(), m); err != nil {
+ if err := svc.Scraper.EnrichOneWithOptions(c.Request.Context(), m, options); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
- c.JSON(http.StatusOK, m)
+ refreshed, _ := svc.Repo.Media.FindByID(c.Request.Context(), m.ID)
+ c.JSON(http.StatusOK, refreshed)
}
}
diff --git a/internal/handler/media_test.go b/internal/handler/media_test.go
index 6603e77..4cb0f09 100644
--- a/internal/handler/media_test.go
+++ b/internal/handler/media_test.go
@@ -1,9 +1,13 @@
package handler
import (
+ "bytes"
"encoding/json"
+ "fmt"
"net/http"
"net/http/httptest"
+ "net/url"
+ "strings"
"testing"
"time"
@@ -154,6 +158,82 @@ func TestListMediaGroupsMultipleVersionsByDefault(t *testing.T) {
}
}
+func TestListLibrarySeriesDoesNotTruncateLargeEpisodeLibraries(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ lib := model.Library{Name: "国漫", Path: "cloud://openlist/国漫", Type: "anime", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ rows := make([]model.Media, 0, 2001)
+ for i := 1; i <= 2001; i++ {
+ rows = append(rows, model.Media{
+ Base: model.Base{ID: fmt.Sprintf("ep-%04d", i), CreatedAt: time.Now().Add(time.Duration(i) * time.Second)},
+ LibraryID: lib.ID,
+ Title: "大剧",
+ Path: fmt.Sprintf("cloud://openlist/国漫/大剧 (2026) {tmdb-123}/Season 1/大剧.S01E%04d.mkv", i),
+ SeasonNum: 1,
+ EpisodeNum: i,
+ })
+ }
+ if err := repos.DB.CreateInBatches(rows, 500).Error; err != nil {
+ t.Fatal(err)
+ }
+ svc := &service.Container{
+ Repo: repos,
+ Media: service.NewMediaService(&config.Config{}, zap.NewNop(), repos),
+ }
+
+ series := requestLibrarySeries(t, svc, "/api/libraries/"+lib.ID+"/series", lib.ID)
+ if series.Total != 1 || len(series.Items) != 1 {
+ t.Fatalf("series response total=%d len=%d body=%#v", series.Total, len(series.Items), series)
+ }
+ if series.Items[0].Count != 2001 {
+ t.Fatalf("series count = %d, want 2001", series.Items[0].Count)
+ }
+ if !strings.HasPrefix(series.Items[0].Key, "series:") ||
+ strings.Contains(series.Items[0].Key, "lib:") ||
+ strings.Contains(series.Items[0].Key, "show:") {
+ t.Fatalf("series key = %q, want compact non-raw key", series.Items[0].Key)
+ }
+ episodes := requestLibrarySeriesEpisodes(t, svc, "/api/libraries/"+lib.ID+"/series/episodes?key="+url.QueryEscape(series.Items[0].Key), lib.ID)
+ if episodes.Total != 2001 || len(episodes.Items) != 2001 {
+ t.Fatalf("episodes total=%d len=%d, want 2001", episodes.Total, len(episodes.Items))
+ }
+ if episodes.Items[0].EpisodeNum != 1 || episodes.Items[len(episodes.Items)-1].EpisodeNum != 2001 {
+ t.Fatalf("episode order first=%d last=%d", episodes.Items[0].EpisodeNum, episodes.Items[len(episodes.Items)-1].EpisodeNum)
+ }
+}
+
+func TestScrapeOptionsFromRequestPreservesEpisodeImagesFalse(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Request = httptest.NewRequest(http.MethodPost, "/api/media/ep-1/scrape", bytes.NewBufferString(`{"episode_images":false,"refresh_matched":true}`))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ options, err := scrapeOptionsFromRequest(c, false)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if options.EpisodeArtwork == nil {
+ t.Fatal("EpisodeArtwork is nil, want explicit false")
+ }
+ if *options.EpisodeArtwork {
+ t.Fatal("EpisodeArtwork = true, want false")
+ }
+ if !options.IncludeMatched {
+ t.Fatal("IncludeMatched = false, want true from refresh_matched")
+ }
+}
+
func requestLibraries(t *testing.T, svc *service.Container, userID, role, path string) []model.Library {
t.Helper()
w := httptest.NewRecorder()
@@ -177,6 +257,16 @@ type mediaListResponse struct {
Total int64 `json:"total"`
}
+type seriesListResponse struct {
+ Items []service.SeriesCard `json:"items"`
+ Total int64 `json:"total"`
+}
+
+type seriesEpisodesResponse struct {
+ Items []model.Media `json:"items"`
+ Total int64 `json:"total"`
+}
+
func requestMediaList(t *testing.T, svc *service.Container, path, libraryID string) mediaListResponse {
t.Helper()
w := httptest.NewRecorder()
@@ -195,3 +285,41 @@ func requestMediaList(t *testing.T, svc *service.Container, path, libraryID stri
}
return payload
}
+
+func requestLibrarySeries(t *testing.T, svc *service.Container, path, libraryID string) seriesListResponse {
+ t.Helper()
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Set(middleware.CtxUserID, "user-1")
+ c.Set(middleware.CtxUserRole, "user")
+ c.Params = gin.Params{{Key: "id", Value: libraryID}}
+ c.Request = httptest.NewRequest(http.MethodGet, path, nil)
+ listLibrarySeriesHandler(svc)(c)
+ if w.Code != http.StatusOK {
+ t.Fatalf("GET %s status = %d body=%s", path, w.Code, w.Body.String())
+ }
+ var payload seriesListResponse
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode series list: %v", err)
+ }
+ return payload
+}
+
+func requestLibrarySeriesEpisodes(t *testing.T, svc *service.Container, path, libraryID string) seriesEpisodesResponse {
+ t.Helper()
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+ c.Set(middleware.CtxUserID, "user-1")
+ c.Set(middleware.CtxUserRole, "user")
+ c.Params = gin.Params{{Key: "id", Value: libraryID}}
+ c.Request = httptest.NewRequest(http.MethodGet, path, nil)
+ listLibrarySeriesEpisodesHandler(svc)(c)
+ if w.Code != http.StatusOK {
+ t.Fatalf("GET %s status = %d body=%s", path, w.Code, w.Body.String())
+ }
+ var payload seriesEpisodesResponse
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode series episodes: %v", err)
+ }
+ return payload
+}
diff --git a/internal/handler/playback_extra.go b/internal/handler/playback_extra.go
index 3ef7429..e335098 100644
--- a/internal/handler/playback_extra.go
+++ b/internal/handler/playback_extra.go
@@ -74,7 +74,7 @@ func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
return
}
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
- streamURL := absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c))
+ streamURL := externalPlaybackURL(c, svc, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c))
escapedStream := url.QueryEscape(streamURL)
c.JSON(http.StatusOK, gin.H{
"url": streamURL,
@@ -100,7 +100,7 @@ func externalURLHandler(svc *service.Container) gin.HandlerFunc {
}
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
c.JSON(http.StatusOK, gin.H{
- "url": absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c)),
+ "url": externalPlaybackURL(c, svc, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c)),
})
}
}
@@ -137,6 +137,71 @@ func externalPlaybackToken(c *gin.Context, svc *service.Container, mediaID strin
return token
}
+func externalPlaybackURL(c *gin.Context, svc *service.Container, path string) string {
+ if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
+ return path
+ }
+ headerOrigin := sanitizedPublicOrigin(c.GetHeader("X-MediaStation-Public-Origin"))
+ if headerOrigin != "" && !isLocalPublicOrigin(headerOrigin) {
+ return joinOriginPath(headerOrigin, path)
+ }
+ if svc != nil {
+ if origin := sanitizedPublicOrigin(service.PublicServerURL(c.Request.Context(), svc.Repo, svc.Cfg)); origin != "" {
+ return joinOriginPath(origin, path)
+ }
+ }
+ if headerOrigin != "" {
+ return joinOriginPath(headerOrigin, path)
+ }
+ return absoluteRequestURL(c, path)
+}
+
+func isLocalPublicOrigin(origin string) bool {
+ u, err := url.Parse(origin)
+ if err != nil || u == nil {
+ return false
+ }
+ host := strings.ToLower(strings.Trim(u.Hostname(), "[]"))
+ switch host {
+ case "localhost", "127.0.0.1", "::1":
+ return true
+ default:
+ return strings.HasPrefix(host, "127.")
+ }
+}
+
+func sanitizedPublicOrigin(raw string) string {
+ raw = strings.TrimSpace(strings.Split(raw, ",")[0])
+ if raw == "" {
+ return ""
+ }
+ u, err := url.Parse(raw)
+ if err != nil || u == nil {
+ return ""
+ }
+ scheme := strings.ToLower(strings.TrimSpace(u.Scheme))
+ if scheme != "http" && scheme != "https" {
+ return ""
+ }
+ if strings.TrimSpace(u.Host) == "" {
+ return ""
+ }
+ u.Scheme = scheme
+ u.User = nil
+ u.Path = ""
+ u.RawPath = ""
+ u.RawQuery = ""
+ u.Fragment = ""
+ return strings.TrimRight(u.String(), "/")
+}
+
+func joinOriginPath(origin, path string) string {
+ if !strings.HasPrefix(path, "/") {
+ path = "/" + path
+ }
+ return strings.TrimRight(origin, "/") + path
+}
+
func absoluteRequestURL(c *gin.Context, path string) string {
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
return path
diff --git a/internal/handler/playback_extra_test.go b/internal/handler/playback_extra_test.go
index db07390..d3a6127 100644
--- a/internal/handler/playback_extra_test.go
+++ b/internal/handler/playback_extra_test.go
@@ -84,6 +84,137 @@ func TestExternalURLUsesMediaScopedPlaybackToken(t *testing.T) {
}
}
+func TestExternalURLPrefersBrowserPublicOriginOverForwardedSource(t *testing.T) {
+ router, _, secret := newPlaybackScopeTestRouter(t)
+ loginToken := signedTestToken(t, secret)
+
+ req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-url", nil)
+ req.Header.Set("Authorization", "Bearer "+loginToken)
+ req.Header.Set("X-Forwarded-Proto", "https")
+ req.Header.Set("X-Forwarded-Host", "media.v6.agonyz.dpdns.org")
+ req.Header.Set("X-MediaStation-Public-Origin", "https://media.agonyz.dpdns.org")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ var payload struct {
+ URL string `json:"url"`
+ }
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode: %v", err)
+ }
+ streamURL, err := url.Parse(payload.URL)
+ if err != nil {
+ t.Fatalf("parse stream url: %v", err)
+ }
+ if got, want := streamURL.Scheme+"://"+streamURL.Host, "https://media.agonyz.dpdns.org"; got != want {
+ t.Fatalf("external url origin = %q, want %q; full url=%s", got, want, payload.URL)
+ }
+ if strings.Contains(payload.URL, "media.v6.agonyz.dpdns.org") {
+ t.Fatalf("external url should not use forwarded source host: %s", payload.URL)
+ }
+}
+
+func TestExternalPlayersSanitizeBrowserPublicOrigin(t *testing.T) {
+ router, _, secret := newPlaybackScopeTestRouter(t)
+ loginToken := signedTestToken(t, secret)
+
+ req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-players", nil)
+ req.Header.Set("Authorization", "Bearer "+loginToken)
+ req.Header.Set("X-MediaStation-Public-Origin", "https://user:pass@media.agonyz.dpdns.org/sneaky/path?x=1#frag")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ var payload struct {
+ URL string `json:"url"`
+ Players []struct {
+ Name string `json:"name"`
+ URL string `json:"url"`
+ } `json:"players"`
+ }
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode: %v", err)
+ }
+ if !strings.HasPrefix(payload.URL, "https://media.agonyz.dpdns.org/api/stream/media-1?") {
+ t.Fatalf("sanitized stream url = %q", payload.URL)
+ }
+ if strings.Contains(payload.URL, "user:pass") || strings.Contains(payload.URL, "sneaky") || strings.Contains(payload.URL, "x=1") || strings.Contains(payload.URL, "#frag") {
+ t.Fatalf("stream url contains unsafe origin components: %s", payload.URL)
+ }
+ for _, player := range payload.Players {
+ if !strings.Contains(player.URL, "media.agonyz.dpdns.org") {
+ t.Fatalf("%s player url does not include sanitized public host: %s", player.Name, player.URL)
+ }
+ if strings.Contains(player.URL, "user:pass") || strings.Contains(player.URL, "sneaky") {
+ t.Fatalf("%s player url contains unsafe origin components: %s", player.Name, player.URL)
+ }
+ }
+}
+
+func TestExternalURLFallsBackToConfiguredPublicServerURL(t *testing.T) {
+ router, svc, secret := newPlaybackScopeTestRouter(t)
+ loginToken := signedTestToken(t, secret)
+ if err := svc.Repo.Setting.Set(t.Context(), "app.server_url", "https://public.example.test"); err != nil {
+ t.Fatalf("set public url: %v", err)
+ }
+
+ req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-url", nil)
+ req.Header.Set("Authorization", "Bearer "+loginToken)
+ req.Header.Set("X-Forwarded-Proto", "https")
+ req.Header.Set("X-Forwarded-Host", "source.example.test")
+ req.Header.Set("X-MediaStation-Public-Origin", "javascript:alert(1)")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ var payload struct {
+ URL string `json:"url"`
+ }
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode: %v", err)
+ }
+ if !strings.HasPrefix(payload.URL, "https://public.example.test/api/stream/media-1?") {
+ t.Fatalf("external url = %q, want configured public origin", payload.URL)
+ }
+}
+
+func TestExternalURLPrefersConfiguredPublicServerURLOverLocalBrowserOrigin(t *testing.T) {
+ router, svc, secret := newPlaybackScopeTestRouter(t)
+ loginToken := signedTestToken(t, secret)
+ if err := svc.Repo.Setting.Set(t.Context(), "app.server_url", "https://media.example.test"); err != nil {
+ t.Fatalf("set public url: %v", err)
+ }
+
+ req := httptest.NewRequest(http.MethodGet, "http://127.0.0.1:8080/api/playback/media-1/external-url", nil)
+ req.Header.Set("Authorization", "Bearer "+loginToken)
+ req.Header.Set("X-MediaStation-Public-Origin", "http://127.0.0.1:8080")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ var payload struct {
+ URL string `json:"url"`
+ }
+ if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
+ t.Fatalf("decode: %v", err)
+ }
+ if !strings.HasPrefix(payload.URL, "https://media.example.test/api/stream/media-1?") {
+ t.Fatalf("external url = %q, want configured public origin instead of localhost", payload.URL)
+ }
+ if strings.Contains(payload.URL, "127.0.0.1:8080") {
+ t.Fatalf("external url should not keep local browser origin when public url is configured: %s", payload.URL)
+ }
+}
+
func TestScopedPlaybackTokenCannotStreamAnotherMedia(t *testing.T) {
router, svc, _ := newPlaybackScopeTestRouter(t)
user, err := svc.Repo.User.FindByID(t.Context(), "user-1")
@@ -248,6 +379,7 @@ func newPlaybackScopeTestRouter(t *testing.T) (*gin.Engine, *service.Container,
api := router.Group("/api")
api.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
api.GET("/playback/:id/external-url", externalURLHandler(svc))
+ api.GET("/playback/:id/external-players", externalPlayersHandler(svc))
api.GET("/stream/:id", streamHandler(svc))
api.GET("/cloud/play/:type", cloudPlayHandler(svc))
return router, svc, cfg.Secrets.JWTSecret
diff --git a/internal/handler/refresh_handler.go b/internal/handler/refresh_handler.go
index b2cd360..f5c6653 100644
--- a/internal/handler/refresh_handler.go
+++ b/internal/handler/refresh_handler.go
@@ -57,6 +57,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
return
}
+ setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "ok",
@@ -72,6 +73,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
// Logout 登出当前用户。
// POST /api/auth/logout
func (h *RefreshHandler) Logout(c *gin.Context) {
+ clearAccessTokenCookie(c)
userID := c.GetString("ctx_user_id")
if userID == "" {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
diff --git a/internal/handler/repair_rescrape.go b/internal/handler/repair_rescrape.go
index 92fb1c5..017d876 100644
--- a/internal/handler/repair_rescrape.go
+++ b/internal/handler/repair_rescrape.go
@@ -17,14 +17,22 @@ import (
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ options, err := scrapeOptionsFromRequest(c, true)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
+ return
+ }
task := startScrapeHTTPTask(svc, "全库修复并重刮", "", "")
- go func() {
- result, err := svc.RepairAndRescrapeAllLibraries(context.Background())
+ go func(options service.ScrapeOptions) {
+ result, err := svc.RepairAndRescrapeAllLibraries(context.Background(), options)
metrics := map[string]int64{
- "repaired": int64(result.Repaired),
- "libraries": int64(result.Libraries),
- "matched": int64(result.Matched),
- "reset": int64(result.Reset),
+ "repaired": int64(result.Repaired),
+ "reclassified": int64(result.Reclassified),
+ "libraries": int64(result.Libraries),
+ "matched": int64(result.Matched),
+ "processed": int64(result.Processed),
+ "errors": int64(result.Errors),
+ "reset": int64(result.Reset),
}
stage := "completed"
message := "全库修复并重刮完成"
@@ -33,7 +41,7 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
message = "全库修复并重刮失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
- }()
+ }(options)
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
}
}
@@ -46,14 +54,22 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libraryID := c.Param("id")
+ options, err := scrapeOptionsFromRequest(c, true)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
+ return
+ }
task := startScrapeHTTPTask(svc, "媒体库修复并重刮", "", "")
- go func() {
- result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID)
+ go func(options service.ScrapeOptions) {
+ result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID, options)
metrics := map[string]int64{
- "repaired": int64(result.Repaired),
- "libraries": int64(result.Libraries),
- "matched": int64(result.Matched),
- "reset": int64(result.Reset),
+ "repaired": int64(result.Repaired),
+ "reclassified": int64(result.Reclassified),
+ "libraries": int64(result.Libraries),
+ "matched": int64(result.Matched),
+ "processed": int64(result.Processed),
+ "errors": int64(result.Errors),
+ "reset": int64(result.Reset),
}
stage := "completed"
message := "媒体库修复并重刮完成"
@@ -62,7 +78,7 @@ func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
message = "媒体库修复并重刮失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
- }()
+ }(options)
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
}
}
diff --git a/internal/handler/routes_admin.go b/internal/handler/routes_admin.go
index e054e85..27beb11 100644
--- a/internal/handler/routes_admin.go
+++ b/internal/handler/routes_admin.go
@@ -10,98 +10,120 @@ import (
)
func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
- // Admin-only endpoints.
admin := api.Group("/admin")
admin.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
- {
- admin.GET("/users", listUsersHandler(svc))
- admin.POST("/users", createUserHandler(svc))
- admin.PATCH("/users/:id", updateUserHandler(svc))
- admin.PATCH("/users/:id/password", resetUserPasswordHandler(svc))
- admin.PATCH("/users/:id/status", updateUserStatusHandler(svc))
- admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
- admin.DELETE("/users/:id", deleteUserHandler(svc))
- admin.GET("/settings", listSettingsHandler(svc))
- admin.PUT("/settings", updateSettingHandler(svc))
- admin.GET("/logs", recentLogsHandler(svc))
-
- // Permissions admin.
- admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
- admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
- admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
-
- // Storage configs (Alist / S3 / WebDAV / 网盘).
- admin.GET("/storage/status", listStorageConfigsHandler(svc))
- admin.GET("/storage/:type", getStorageConfigHandler(svc))
- admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
- admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
- admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
- admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
-
- // Cloud disk (115 / 夸克) browsing, QR login and 302 import.
- admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
- admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
- admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
- admin.GET("/cloud/:type/list", cloudListHandler(svc))
- admin.POST("/cloud/:type/import", cloudImportHandler(svc))
- admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
- admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
- admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
-
- // Download client CRUD.
- admin.GET("/download/clients", listDownloadClientsHandler(svc))
- admin.POST("/download/clients", createDownloadClientHandler(svc))
- admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
- admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
- admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
- admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
-
- // System scheduler trigger alias.
- admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
-
- // Database backup.
- admin.GET("/backups", listBackupsHandler(svc))
- admin.POST("/backups", createBackupHandler(svc))
- admin.DELETE("/backups", deleteBackupHandler(svc))
- admin.POST("/backups/restore", restoreBackupHandler(svc))
-
- // Notifications (test endpoint).
- admin.POST("/notify/test", notifyTestHandler(svc))
-
- // Notify channels CRUD + per-channel test.
- admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
- admin.POST("/notify/channels", createNotifyChannelHandler(svc))
- admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
- admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
- admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
-
- // Telegram Bot webhook management.
- admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
- admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
- admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
- admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
-
- // File organizer.
- admin.POST("/media/:id/organize", organizeMediaHandler(svc))
- admin.POST("/libraries/:id/organize", organizeLibraryHandler(svc))
- admin.GET("/organize/sources", organizeSourcesHandler(svc))
- admin.POST("/organize/source", organizeDirectoryHandler(svc))
-
- // 全库修复+重刮:从路径占位符回填缺失外部 ID,然后批量重刮整库。
- admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
- // 单库修复+重刮:只对指定媒体库回填占位符外部 ID 并重刮。
- admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
-
- // API key management (encrypted at rest).
- admin.GET("/api-configs", listAPIConfigsHandler(svc))
- admin.GET("/api-configs/:provider", getAPIConfigHandler(svc))
- admin.PUT("/api-configs/:provider", updateAPIConfigHandler(svc))
- admin.DELETE("/api-configs/:provider", deleteAPIConfigHandler(svc))
-
- // Scheduled jobs.
- admin.GET("/scheduler", schedulerStatusHandler(svc))
- admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
-
- }
-
+ registerAdminUserRoutes(admin, svc)
+ registerAdminPermissionRoutes(admin, svc)
+ registerAdminStorageRoutes(admin, svc)
+ registerAdminCloudRoutes(admin, svc)
+ registerAdminDownloadClientRoutes(admin, svc)
+ registerAdminSystemRoutes(admin, svc)
+ registerAdminBackupRoutes(admin, svc)
+ registerAdminNotificationRoutes(admin, svc)
+ registerAdminTelegramRoutes(admin, svc)
+ registerAdminOrganizerRoutes(admin, svc)
+ registerAdminRepairRoutes(admin, svc)
+ registerAdminAPIConfigRoutes(admin, svc)
+ registerAdminSchedulerRoutes(admin, svc)
+}
+
+func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/users", listUsersHandler(svc))
+ admin.POST("/users", createUserHandler(svc))
+ admin.PATCH("/users/:id", updateUserHandler(svc))
+ admin.PATCH("/users/:id/password", resetUserPasswordHandler(svc))
+ admin.PATCH("/users/:id/status", updateUserStatusHandler(svc))
+ admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
+ admin.DELETE("/users/:id", deleteUserHandler(svc))
+ admin.GET("/settings", listSettingsHandler(svc))
+ admin.PUT("/settings", updateSettingHandler(svc))
+ admin.GET("/logs", recentLogsHandler(svc))
+}
+
+func registerAdminPermissionRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
+ admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
+ admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
+}
+
+func registerAdminStorageRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/storage/status", listStorageConfigsHandler(svc))
+ admin.GET("/storage/:type", getStorageConfigHandler(svc))
+ admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
+ admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
+ admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
+ admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
+}
+
+func registerAdminCloudRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
+ admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
+ admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
+ admin.GET("/cloud/:type/list", cloudListHandler(svc))
+ admin.POST("/cloud/:type/mkdir", cloudMkdirHandler(svc))
+ admin.PUT("/cloud/:type/rename", cloudRenameHandler(svc))
+ admin.POST("/cloud/:type/import", cloudImportHandler(svc))
+ admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
+ admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
+ admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
+}
+
+func registerAdminDownloadClientRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/download/clients", listDownloadClientsHandler(svc))
+ admin.POST("/download/clients", createDownloadClientHandler(svc))
+ admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
+ admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
+ admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
+ admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
+}
+
+func registerAdminSystemRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
+}
+
+func registerAdminBackupRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/backups", listBackupsHandler(svc))
+ admin.POST("/backups", createBackupHandler(svc))
+ admin.DELETE("/backups", deleteBackupHandler(svc))
+ admin.POST("/backups/restore", restoreBackupHandler(svc))
+}
+
+func registerAdminNotificationRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.POST("/notify/test", notifyTestHandler(svc))
+ admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
+ admin.POST("/notify/channels", createNotifyChannelHandler(svc))
+ admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
+ admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
+ admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
+}
+
+func registerAdminTelegramRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
+ admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
+ admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
+ admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
+}
+
+func registerAdminOrganizerRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.POST("/media/:id/organize", organizeMediaHandler(svc))
+ admin.POST("/libraries/:id/organize", organizeLibraryHandler(svc))
+ admin.GET("/organize/sources", organizeSourcesHandler(svc))
+ admin.POST("/organize/source", organizeDirectoryHandler(svc))
+}
+
+func registerAdminRepairRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
+ admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
+}
+
+func registerAdminAPIConfigRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/api-configs", listAPIConfigsHandler(svc))
+ admin.GET("/api-configs/:provider", getAPIConfigHandler(svc))
+ admin.PUT("/api-configs/:provider", updateAPIConfigHandler(svc))
+ admin.DELETE("/api-configs/:provider", deleteAPIConfigHandler(svc))
+}
+
+func registerAdminSchedulerRoutes(admin *gin.RouterGroup, svc *service.Container) {
+ admin.GET("/scheduler", schedulerStatusHandler(svc))
+ admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
}
diff --git a/internal/handler/routes_admin_test.go b/internal/handler/routes_admin_test.go
new file mode 100644
index 0000000..5f54e73
--- /dev/null
+++ b/internal/handler/routes_admin_test.go
@@ -0,0 +1,45 @@
+package handler
+
+import (
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestAdminRouteSurfacesAreRegistered(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ router := gin.New()
+ Register(router, &config.Config{
+ Secrets: config.SecretsConfig{JWTSecret: "test-secret"},
+ }, zap.NewNop(), &service.Container{Log: zap.NewNop()})
+
+ routes := map[string]bool{}
+ for _, route := range router.Routes() {
+ routes[route.Method+" "+route.Path] = true
+ }
+
+ for _, want := range []string{
+ "GET /api/admin/users",
+ "GET /api/admin/users/:id/permissions",
+ "GET /api/admin/storage/status",
+ "GET /api/admin/cloud/:type/list",
+ "GET /api/admin/download/clients",
+ "POST /api/admin/system/scheduler/:name/trigger",
+ "POST /api/admin/backups",
+ "GET /api/admin/notify/channels",
+ "GET /api/admin/telegram/webhook",
+ "GET /api/admin/organize/sources",
+ "POST /api/admin/media/repair-rescrape",
+ "GET /api/admin/api-configs",
+ "POST /api/admin/scheduler/:name/run",
+ } {
+ if !routes[want] {
+ t.Fatalf("%s route is not registered", want)
+ }
+ }
+}
diff --git a/internal/handler/routes_authenticated.go b/internal/handler/routes_authenticated.go
index 67f4832..e810cea 100644
--- a/internal/handler/routes_authenticated.go
+++ b/internal/handler/routes_authenticated.go
@@ -10,261 +10,35 @@ import (
)
func registerAuthenticatedRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
- // Authenticated endpoints.
authed := api.Group("/")
authed.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
authed.Use(activeUserRequired(svc))
- {
- authed.GET("/me", meHandler(svc))
- authed.PATCH("/me", updateProfileHandler(svc))
- authed.POST("/me/password", changePasswordHandler(svc))
- authed.POST("/me/logout", logoutHandler(svc))
-
- // Permissions.
- authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
-
- // License activation bridge (admin only; talks to the configured license server).
- authed.GET("/license/status", middleware.AdminRequired(), licenseStatusHandler(svc))
- authed.POST("/license/activate", middleware.AdminRequired(), licenseActivateHandler(svc))
- authed.POST("/license/heartbeat", middleware.AdminRequired(), licenseHeartbeatHandler(svc))
-
- // Libraries.
- authed.GET("/libraries", listLibrariesHandler(svc))
- authed.POST("/libraries", middleware.AdminRequired(), createLibraryHandler(svc))
- authed.DELETE("/libraries/:id", middleware.AdminRequired(), deleteLibraryHandler(svc))
- authed.POST("/libraries/:id/scan", middleware.AdminRequired(), scanLibraryHandler(svc))
- authed.POST("/libraries/:id/scrape", middleware.AdminRequired(), scrapeLibraryHandler(svc))
-
- authed.GET("/libraries/:id/media", listMediaHandler(svc))
- authed.GET("/libraries/:id/seasons", listSeasonsHandler(svc))
-
- // Media.
- authed.GET("/media/:id", getMediaHandler(svc))
- authed.GET("/media", searchMediaHandler(svc))
- authed.PATCH("/media/:id/metadata", middleware.AdminRequired(), updateMediaMetadataHandler(svc))
- authed.POST("/media/:id/scrape", middleware.AdminRequired(), scrapeOneHandler(svc))
- authed.GET("/media/:id/scrape/search", middleware.AdminRequired(), manualScrapeSearchHandler(svc))
- authed.POST("/media/:id/scrape/apply", middleware.AdminRequired(), manualScrapeApplyOneHandler(svc))
- authed.POST("/media/scrape/apply", middleware.AdminRequired(), manualScrapeApplyBatchHandler(svc))
- authed.POST("/media/:id/probe", middleware.AdminRequired(), reprobeHandler(svc))
- authed.DELETE("/media/:id", middleware.AdminRequired(), deleteMediaHandler(svc))
- authed.POST("/media/:id/restore", middleware.AdminRequired(), restoreMediaHandler(svc))
- authed.DELETE("/media/:id/purge", middleware.AdminRequired(), purgeMediaHandler(svc))
- authed.GET("/media/:id/subtitles", listSubtitlesHandler(svc))
- authed.GET("/subtitles/:id", serveSubtitleHandler(svc))
- authed.POST("/media/:id/nfo", middleware.AdminRequired(), exportNFOHandler(svc))
- authed.POST("/libraries/:id/nfo", middleware.AdminRequired(), exportLibraryNFOHandler(svc))
-
- // Streaming.
- authed.GET("/stream/:id", streamHandler(svc))
- authed.HEAD("/stream/:id", streamHandler(svc))
- authed.GET("/hls/:id/index.m3u8", hlsPlaylistHandler(svc))
- authed.GET("/hls/:id/:seg", hlsSegmentHandler(svc))
- authed.DELETE("/hls/:id", stopTranscodeHandler(svc))
-
- // Cloud-disk 302 playback redirect (resolves a fresh direct link).
- authed.GET("/cloud/play/:type", cloudPlayHandler(svc))
- authed.HEAD("/cloud/play/:type", cloudPlayHandler(svc))
-
- // Image proxy (URL passed as ?url=...).
- authed.GET("/img", imageProxyHandler(svc))
-
- // History / favourites / playlists.
- authed.GET("/history", recentHistoryHandler(svc))
- authed.POST("/history", recordProgressHandler(svc))
-
- authed.GET("/favourites", listFavouritesHandler(svc))
- authed.POST("/favourites/:id", toggleFavouriteHandler(svc))
-
- // Storage breakdown.
- authed.GET("/storage", storageBreakdownHandler(svc))
-
- authed.GET("/playlists", listPlaylistsHandler(svc))
- authed.POST("/playlists", createPlaylistHandler(svc))
- authed.GET("/playlists/:id", getPlaylistHandler(svc))
- authed.POST("/playlists/:id/items", addPlaylistItemHandler(svc))
- authed.DELETE("/playlists/:id/items/:media_id", removePlaylistItemHandler(svc))
- authed.DELETE("/playlists/:id", deletePlaylistHandler(svc))
-
- // Downloads.
- authed.GET("/downloads", requirePermission(svc, "can_manage_downloads"), listDownloadsHandler(svc))
- authed.POST("/downloads", requirePermission(svc, "can_manage_downloads"), addDownloadHandler(svc))
- authed.DELETE("/downloads/:hash", requirePermission(svc, "can_manage_downloads"), deleteDownloadHandler(svc))
- authed.POST("/downloads/relocate", requirePermission(svc, "can_manage_downloads"), relocateDownloadHandler(svc))
- authed.POST("/downloads/reload", requirePermission(svc, "can_manage_downloads"), reloadDownloadConfigHandler(svc))
-
- // Subscriptions.
- authed.GET("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionsHandler(svc))
- authed.GET("/subscriptions/history", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionHistoryHandler(svc))
- authed.POST("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), createSubscriptionHandler(svc))
- authed.DELETE("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), deleteSubscriptionHandler(svc))
- authed.POST("/subscriptions/:id/restore", requirePermission(svc, "can_manage_subscriptions"), restoreSubscriptionHandler(svc))
- authed.POST("/subscriptions/:id/run", requirePermission(svc, "can_manage_subscriptions"), runSubscriptionHandler(svc))
-
- // Stats / dashboard.
- authed.GET("/stats", statsHandler(svc))
- authed.GET("/tasks", middleware.AdminRequired(), tasksHandler(svc))
-
- // Discover (TMDb trending / popular).
- authed.GET("/discover/trending", requirePermission(svc, "can_view_discover"), trendingHandler(svc))
- authed.GET("/discover/popular", requirePermission(svc, "can_view_discover"), popularHandler(svc))
-
- // AI.
- authed.GET("/ai/status", requirePermission(svc, "can_use_ai"), aiStatusHandler(svc))
- authed.POST("/ai/search", requirePermission(svc, "can_use_ai"), smartSearchHandler(svc))
- authed.GET("/ai/recommend", requirePermission(svc, "can_use_ai"), aiRecommendHandler(svc))
-
- // File browser (used by the library-path picker).
- authed.GET("/files", middleware.AdminRequired(), browseFilesHandler(svc))
- authed.POST("/files/folders", middleware.AdminRequired(), createFolderHandler(svc))
- authed.PUT("/files/rename", middleware.AdminRequired(), renameFileHandler(svc))
- authed.DELETE("/files", middleware.AdminRequired(), deleteFileHandler(svc))
- authed.POST("/files/transfer", middleware.AdminRequired(), transferFileHandler(svc))
-
- // DLNA discovery + cast.
- authed.GET("/dlna/devices", dlnaListHandler(svc))
- authed.POST("/dlna/cast", dlnaCastHandler(svc))
-
- // STRM (URL-as-file).
- authed.PUT("/media/:id/strm", middleware.AdminRequired(), setSTRMHandler(svc))
- authed.DELETE("/media/:id/strm", middleware.AdminRequired(), clearSTRMHandler(svc))
- authed.POST("/strm/import", middleware.AdminRequired(), importSTRMHandler(svc))
- authed.POST("/strm/generate", middleware.AdminRequired(), generateSTRMHandler(svc))
-
- // Duplicate finder.
- authed.GET("/duplicates", middleware.AdminRequired(), listDuplicatesHandler(svc))
- authed.POST("/duplicates/scan", middleware.AdminRequired(), detectDuplicatesHandler(svc))
- authed.POST("/duplicates/unmark", middleware.AdminRequired(), unmarkDuplicatesHandler(svc))
-
- // Site management + cross-site torrent search (via SiteHandler).
- siteHandler := NewSiteHandler(svc)
- authed.GET("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.ListSites)
- authed.GET("/sites/types", requirePermission(svc, "can_manage_sites"), siteHandler.GetSiteTypes)
- authed.GET("/sites/auth-types", requirePermission(svc, "can_manage_sites"), siteHandler.GetAuthTypes)
- authed.POST("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.CreateSite)
- authed.GET("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.GetSite)
- authed.PUT("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.UpdateSite)
- authed.DELETE("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.DeleteSite)
- authed.POST("/sites/:id/test", requirePermission(svc, "can_manage_sites"), siteHandler.TestSite)
- authed.GET("/sites/search", requirePermission(svc, "can_manage_sites"), siteSearchHandler(svc))
-
- // Recycle bin.
- authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
- authed.POST("/recycle/restore", middleware.AdminRequired(), restoreMediaBatchHandler(svc))
- authed.POST("/recycle/purge", middleware.AdminRequired(), purgeMediaBatchHandler(svc))
-
- authed.GET("/ws", wsHandler(svc))
-
- // SSE event stream.
- authed.GET("/events", sseHandler(svc))
-
- // Scheduler.
- authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
- authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
- authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
-
- // ── Auxiliary endpoints used by the React UI rails ──
- authed.GET("/media/recent", recentMediaHandler(svc))
- authed.GET("/media/stats", mediaStatsHandler(svc))
-
- // Watch history (extra surface beyond /history).
- authed.GET("/watch-history", historyListHandler(svc))
- authed.GET("/watch-history/stats", historyStatsHandler(svc))
- authed.GET("/watch-history/continue", historyContinueHandler(svc))
- authed.DELETE("/watch-history", historyDeleteHandler(svc))
- authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
-
- // Multi-section TMDb feed used by DiscoverPage.
- authed.GET("/discover/sections", requirePermission(svc, "can_view_discover"), discoverSectionsHandler(svc))
- authed.GET("/discover/feed", requirePermission(svc, "can_view_discover"), discoverFeedHandler(svc))
-
- // System metadata + read-only scheduler view.
- authed.GET("/system/info", systemInfoHandler(svc))
- authed.GET("/system/status", systemStatusHandler(svc))
- authed.GET("/system/scheduler", systemSchedulerHandler(svc))
-
- // Richer dashboard rails.
- authed.GET("/stats/overview", statsOverviewHandler(svc))
- authed.GET("/stats/trend", statsTrendHandler(svc))
- authed.GET("/stats/top-content", statsTopContentHandler(svc))
- authed.GET("/stats/libraries", statsLibrariesHandler(svc))
- authed.GET("/stats/monitor", statsMonitorHandler(svc))
-
- // Multi-persona play profiles (caller-scoped).
- authed.GET("/play-profiles", listPlayProfilesHandler(svc))
- authed.POST("/play-profiles", createPlayProfileHandler(svc))
- authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
- authed.POST("/play-profiles/:id/verify-pin", verifyPlayProfilePINHandler(svc))
- authed.DELETE("/play-profiles/:id", deletePlayProfileHandler(svc))
-
- // ── Search aliases ──
- authed.GET("/search", searchUnifiedHandler(svc))
- authed.GET("/search/advanced", searchAdvancedHandler(svc))
- authed.GET("/search/tmdb", searchTMDbHandler(svc))
- authed.GET("/search/sites", searchSitesHandler(svc))
-
- // ── System extras ──
- authed.GET("/system/config", listSystemConfigHandler(svc))
- authed.GET("/settings/schema", schemaHandler(svc))
- authed.GET("/system/events/ticket", systemEventsTicketHandler(svc))
-
- // ── Per-user stats ──
- authed.GET("/stats/user/:id", statsUserHandler(svc))
- authed.GET("/stats/top-users", statsTopUsersHandler(svc))
- authed.POST("/stats/play", statsPlayHandler(svc))
-
- // ── Sites extras ──
- authed.GET("/sites/:id/resource", requirePermission(svc, "can_manage_sites"), siteResourceHandler(svc))
- authed.GET("/sites/:id/userdata", requirePermission(svc, "can_manage_sites"), siteUserdataHandler(svc))
-
- // ── Subscription extras ──
- authed.PUT("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), updateSubscriptionHandler(svc))
- authed.POST("/subscriptions/:id/search", requirePermission(svc, "can_manage_subscriptions"), searchSubscriptionHandler(svc))
-
- // ── Playlist extras ──
- authed.POST("/playlists/:id/reorder", reorderPlaylistHandler(svc))
- authed.DELETE("/playlists/:id/items/by-id/:item_id", deletePlaylistItemByIDHandler(svc))
-
- // ── DLNA per-renderer control ──
- authed.POST("/dlna/:uuid/play", dlnaPlayHandler(svc))
- authed.POST("/dlna/:uuid/pause", dlnaPauseHandler(svc))
- authed.POST("/dlna/:uuid/stop", dlnaStopHandler(svc))
- authed.GET("/dlna/:uuid/status", dlnaStatusHandler(svc))
-
- // ── Media favourite alias surface ──
- authed.GET("/favorites", listFavoritesAliasHandler(svc))
- authed.POST("/media/:id/favorite", addMediaFavoriteHandler(svc))
- authed.DELETE("/media/:id/favorite", removeMediaFavoriteHandler(svc))
- authed.GET("/media/:id/favorite/status", getMediaFavoriteStatusHandler(svc))
- authed.POST("/media/:id/ai-scrape", requirePermission(svc, "can_rescrape"), aiScrapeMediaHandler(svc))
- authed.POST("/media/scrape/test", requirePermission(svc, "can_rescrape"), scrapeTestHandler(svc))
- authed.POST("/media/organize", requirePermission(svc, "can_manage_files"), organizeBulkHandler(svc))
-
- // ── Playback metadata + external player handoff ──
- authed.GET("/playback/:id/info", playbackInfoHandler(svc))
- authed.POST("/playback/:id/progress", playbackProgressHandler(svc))
- authed.GET("/playback/:id/external-players", externalPlayersHandler(svc))
- authed.GET("/playback/:id/external-url", externalURLHandler(svc))
- authed.GET("/playback/transcode/:job_id/status", transcodeStatusHandler(svc))
-
- // ── Download task ops + sync triggers ──
- authed.POST("/download/:id/pause", requirePermission(svc, "can_manage_downloads"), downloadPauseHandler(svc))
- authed.POST("/download/:id/resume", requirePermission(svc, "can_manage_downloads"), downloadResumeHandler(svc))
- authed.POST("/download/:id/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeOneHandler(svc))
- authed.POST("/download/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeAllHandler(svc))
- authed.POST("/download/sync", requirePermission(svc, "can_manage_downloads"), downloadSyncHandler(svc))
- authed.POST("/download/start-auto-sync", requirePermission(svc, "can_manage_downloads"), downloadAutoSyncHandler(svc))
- authed.GET("/download/tasks", requirePermission(svc, "can_manage_downloads"), downloadTasksAliasHandler(svc))
-
- // ── Assistant (multi-turn AI chat) ──
- authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
- authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
- authed.GET("/admin/assistant/session/:id", getAssistantSessionHandler(svc))
- authed.DELETE("/admin/assistant/session/:id", deleteAssistantSessionHandler(svc))
- authed.POST("/admin/assistant/chat", assistantChatHandler(svc))
- authed.POST("/admin/assistant/execute", assistantExecuteHandler(svc))
- authed.POST("/admin/assistant/undo/:op_id", assistantUndoHandler(svc))
- authed.GET("/admin/assistant/history", assistantHistoryHandler(svc))
- }
+ registerAuthedUserAndLicenseRoutes(authed, svc)
+ registerAuthedLibraryRoutes(authed, svc)
+ registerAuthedMediaRoutes(authed, svc)
+ registerAuthedPlaybackAndProxyRoutes(authed, svc)
+ registerAuthedCollectionRoutes(authed, svc)
+ registerAuthedDownloadRoutes(authed, svc)
+ registerAuthedSubscriptionRoutes(authed, svc)
+ registerAuthedStatsDiscoveryAndAIRoutes(authed, svc)
+ registerAuthedFileRoutes(authed, svc)
+ registerAuthedDLNARoutes(authed, svc)
+ registerAuthedSTRMRoutes(authed, svc)
+ registerAuthedDuplicateRoutes(authed, svc)
+ registerAuthedSiteRoutes(authed, svc)
+ registerAuthedRecycleAndRealtimeRoutes(authed, svc)
+ registerAuthedSchedulerRoutes(authed, svc)
+ registerAuthedUISurfaceRoutes(authed, svc)
+ registerAuthedSearchRoutes(authed, svc)
+ registerAuthedSystemExtraRoutes(authed, svc)
+ registerAuthedStatsExtraRoutes(authed, svc)
+ registerAuthedSitesExtraRoutes(authed, svc)
+ registerAuthedSubscriptionExtraRoutes(authed, svc)
+ registerAuthedPlaylistExtraRoutes(authed, svc)
+ registerAuthedDLNAControlRoutes(authed, svc)
+ registerAuthedFavoriteAndMediaActionRoutes(authed, svc)
+ registerAuthedPlaybackExtraRoutes(authed, svc)
+ registerAuthedDownloadOpsRoutes(authed, svc)
+ registerAuthedAssistantRoutes(authed, svc)
}
diff --git a/internal/handler/routes_authenticated_core.go b/internal/handler/routes_authenticated_core.go
new file mode 100644
index 0000000..ecc99b8
--- /dev/null
+++ b/internal/handler/routes_authenticated_core.go
@@ -0,0 +1,84 @@
+package handler
+
+import (
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func registerAuthedUserAndLicenseRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/me", meHandler(svc))
+ authed.PATCH("/me", updateProfileHandler(svc))
+ authed.POST("/me/password", changePasswordHandler(svc))
+ authed.POST("/me/logout", logoutHandler(svc))
+
+ authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
+
+ authed.GET("/license/status", middleware.AdminRequired(), licenseStatusHandler(svc))
+ authed.POST("/license/activate", middleware.AdminRequired(), licenseActivateHandler(svc))
+ authed.POST("/license/heartbeat", middleware.AdminRequired(), licenseHeartbeatHandler(svc))
+}
+
+func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/libraries", listLibrariesHandler(svc))
+ authed.POST("/libraries", middleware.AdminRequired(), createLibraryHandler(svc))
+ authed.DELETE("/libraries/:id", middleware.AdminRequired(), deleteLibraryHandler(svc))
+ authed.POST("/libraries/:id/scan", middleware.AdminRequired(), scanLibraryHandler(svc))
+ authed.POST("/libraries/:id/scrape", middleware.AdminRequired(), scrapeLibraryHandler(svc))
+
+ authed.GET("/libraries/:id/media", listMediaHandler(svc))
+ authed.GET("/libraries/:id/series", listLibrarySeriesHandler(svc))
+ authed.GET("/libraries/:id/series/episodes", listLibrarySeriesEpisodesHandler(svc))
+ authed.GET("/libraries/:id/seasons", listSeasonsHandler(svc))
+}
+
+func registerAuthedMediaRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/media/:id", getMediaHandler(svc))
+ authed.GET("/media", searchMediaHandler(svc))
+ authed.PATCH("/media/:id/metadata", middleware.AdminRequired(), updateMediaMetadataHandler(svc))
+ authed.POST("/media/:id/scrape", middleware.AdminRequired(), scrapeOneHandler(svc))
+ authed.GET("/media/:id/scrape/search", middleware.AdminRequired(), manualScrapeSearchHandler(svc))
+ authed.POST("/media/:id/scrape/apply", middleware.AdminRequired(), manualScrapeApplyOneHandler(svc))
+ authed.POST("/media/scrape/apply", middleware.AdminRequired(), manualScrapeApplyBatchHandler(svc))
+ authed.POST("/media/:id/probe", middleware.AdminRequired(), reprobeHandler(svc))
+ authed.DELETE("/media/:id", middleware.AdminRequired(), deleteMediaHandler(svc))
+ authed.POST("/media/:id/restore", middleware.AdminRequired(), restoreMediaHandler(svc))
+ authed.DELETE("/media/:id/purge", middleware.AdminRequired(), purgeMediaHandler(svc))
+ authed.GET("/media/:id/subtitles", listSubtitlesHandler(svc))
+ authed.GET("/subtitles/:id", serveSubtitleHandler(svc))
+ authed.POST("/media/:id/nfo", middleware.AdminRequired(), exportNFOHandler(svc))
+ authed.POST("/libraries/:id/nfo", middleware.AdminRequired(), exportLibraryNFOHandler(svc))
+}
+
+func registerAuthedPlaybackAndProxyRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/stream/:id", streamHandler(svc))
+ authed.HEAD("/stream/:id", streamHandler(svc))
+ authed.GET("/hls/:id/index.m3u8", hlsPlaylistHandler(svc))
+ authed.GET("/hls/:id/:seg", hlsSegmentHandler(svc))
+ authed.DELETE("/hls/:id", stopTranscodeHandler(svc))
+
+ authed.GET("/cloud/play/:type", cloudPlayHandler(svc))
+ authed.HEAD("/cloud/play/:type", cloudPlayHandler(svc))
+
+ authed.GET("/img/cloud/:type", cloudArtworkProxyHandler(svc))
+ authed.HEAD("/img/cloud/:type", cloudArtworkProxyHandler(svc))
+ authed.GET("/img", imageProxyHandler(svc))
+}
+
+func registerAuthedCollectionRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/history", recentHistoryHandler(svc))
+ authed.POST("/history", recordProgressHandler(svc))
+
+ authed.GET("/favourites", listFavouritesHandler(svc))
+ authed.POST("/favourites/:id", toggleFavouriteHandler(svc))
+
+ authed.GET("/storage", storageBreakdownHandler(svc))
+
+ authed.GET("/playlists", listPlaylistsHandler(svc))
+ authed.POST("/playlists", createPlaylistHandler(svc))
+ authed.GET("/playlists/:id", getPlaylistHandler(svc))
+ authed.POST("/playlists/:id/items", addPlaylistItemHandler(svc))
+ authed.DELETE("/playlists/:id/items/:media_id", removePlaylistItemHandler(svc))
+ authed.DELETE("/playlists/:id", deletePlaylistHandler(svc))
+}
diff --git a/internal/handler/routes_authenticated_extras.go b/internal/handler/routes_authenticated_extras.go
new file mode 100644
index 0000000..7e4b256
--- /dev/null
+++ b/internal/handler/routes_authenticated_extras.go
@@ -0,0 +1,117 @@
+package handler
+
+import (
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func registerAuthedUISurfaceRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/media/recent", recentMediaHandler(svc))
+ authed.GET("/media/stats", mediaStatsHandler(svc))
+
+ authed.GET("/watch-history", historyListHandler(svc))
+ authed.GET("/watch-history/stats", historyStatsHandler(svc))
+ authed.GET("/watch-history/continue", historyContinueHandler(svc))
+ authed.DELETE("/watch-history", historyDeleteHandler(svc))
+ authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
+
+ authed.GET("/discover/sections", requirePermission(svc, "can_view_discover"), discoverSectionsHandler(svc))
+ authed.GET("/discover/feed", requirePermission(svc, "can_view_discover"), discoverFeedHandler(svc))
+
+ authed.GET("/system/info", systemInfoHandler(svc))
+ authed.GET("/system/status", systemStatusHandler(svc))
+ authed.GET("/system/scheduler", systemSchedulerHandler(svc))
+
+ authed.GET("/stats/overview", statsOverviewHandler(svc))
+ authed.GET("/stats/trend", statsTrendHandler(svc))
+ authed.GET("/stats/top-content", statsTopContentHandler(svc))
+ authed.GET("/stats/libraries", statsLibrariesHandler(svc))
+ authed.GET("/stats/monitor", statsMonitorHandler(svc))
+
+ authed.GET("/play-profiles", listPlayProfilesHandler(svc))
+ authed.POST("/play-profiles", createPlayProfileHandler(svc))
+ authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
+ authed.POST("/play-profiles/:id/verify-pin", verifyPlayProfilePINHandler(svc))
+ authed.DELETE("/play-profiles/:id", deletePlayProfileHandler(svc))
+}
+
+func registerAuthedSearchRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/search", searchUnifiedHandler(svc))
+ authed.GET("/search/advanced", searchAdvancedHandler(svc))
+ authed.GET("/search/tmdb", searchTMDbHandler(svc))
+ authed.GET("/search/sites", searchSitesHandler(svc))
+}
+
+func registerAuthedSystemExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/system/config", listSystemConfigHandler(svc))
+ authed.GET("/settings/schema", schemaHandler(svc))
+ authed.GET("/system/events/ticket", systemEventsTicketHandler(svc))
+}
+
+func registerAuthedStatsExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/stats/user/:id", statsUserHandler(svc))
+ authed.GET("/stats/top-users", statsTopUsersHandler(svc))
+ authed.POST("/stats/play", statsPlayHandler(svc))
+}
+
+func registerAuthedSitesExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/sites/:id/resource", requirePermission(svc, "can_manage_sites"), siteResourceHandler(svc))
+ authed.GET("/sites/:id/userdata", requirePermission(svc, "can_manage_sites"), siteUserdataHandler(svc))
+}
+
+func registerAuthedSubscriptionExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.PUT("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), updateSubscriptionHandler(svc))
+ authed.POST("/subscriptions/:id/search", requirePermission(svc, "can_manage_subscriptions"), searchSubscriptionHandler(svc))
+}
+
+func registerAuthedPlaylistExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.POST("/playlists/:id/reorder", reorderPlaylistHandler(svc))
+ authed.DELETE("/playlists/:id/items/by-id/:item_id", deletePlaylistItemByIDHandler(svc))
+}
+
+func registerAuthedDLNAControlRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.POST("/dlna/:uuid/play", dlnaPlayHandler(svc))
+ authed.POST("/dlna/:uuid/pause", dlnaPauseHandler(svc))
+ authed.POST("/dlna/:uuid/stop", dlnaStopHandler(svc))
+ authed.GET("/dlna/:uuid/status", dlnaStatusHandler(svc))
+}
+
+func registerAuthedFavoriteAndMediaActionRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/favorites", listFavoritesAliasHandler(svc))
+ authed.POST("/media/:id/favorite", addMediaFavoriteHandler(svc))
+ authed.DELETE("/media/:id/favorite", removeMediaFavoriteHandler(svc))
+ authed.GET("/media/:id/favorite/status", getMediaFavoriteStatusHandler(svc))
+ authed.POST("/media/:id/ai-scrape", requirePermission(svc, "can_rescrape"), aiScrapeMediaHandler(svc))
+ authed.POST("/media/scrape/test", requirePermission(svc, "can_rescrape"), scrapeTestHandler(svc))
+ authed.POST("/media/organize", requirePermission(svc, "can_manage_files"), organizeBulkHandler(svc))
+}
+
+func registerAuthedPlaybackExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/playback/:id/info", playbackInfoHandler(svc))
+ authed.POST("/playback/:id/progress", playbackProgressHandler(svc))
+ authed.GET("/playback/:id/external-players", externalPlayersHandler(svc))
+ authed.GET("/playback/:id/external-url", externalURLHandler(svc))
+ authed.GET("/playback/transcode/:job_id/status", transcodeStatusHandler(svc))
+}
+
+func registerAuthedDownloadOpsRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.POST("/download/:id/pause", requirePermission(svc, "can_manage_downloads"), downloadPauseHandler(svc))
+ authed.POST("/download/:id/resume", requirePermission(svc, "can_manage_downloads"), downloadResumeHandler(svc))
+ authed.POST("/download/:id/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeOneHandler(svc))
+ authed.POST("/download/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeAllHandler(svc))
+ authed.POST("/download/sync", requirePermission(svc, "can_manage_downloads"), downloadSyncHandler(svc))
+ authed.POST("/download/start-auto-sync", requirePermission(svc, "can_manage_downloads"), downloadAutoSyncHandler(svc))
+ authed.GET("/download/tasks", requirePermission(svc, "can_manage_downloads"), downloadTasksAliasHandler(svc))
+}
+
+func registerAuthedAssistantRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
+ authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
+ authed.GET("/admin/assistant/session/:id", getAssistantSessionHandler(svc))
+ authed.DELETE("/admin/assistant/session/:id", deleteAssistantSessionHandler(svc))
+ authed.POST("/admin/assistant/chat", assistantChatHandler(svc))
+ authed.POST("/admin/assistant/execute", assistantExecuteHandler(svc))
+ authed.POST("/admin/assistant/undo/:op_id", assistantUndoHandler(svc))
+ authed.GET("/admin/assistant/history", assistantHistoryHandler(svc))
+}
diff --git a/internal/handler/routes_authenticated_features.go b/internal/handler/routes_authenticated_features.go
new file mode 100644
index 0000000..72fe5c3
--- /dev/null
+++ b/internal/handler/routes_authenticated_features.go
@@ -0,0 +1,91 @@
+package handler
+
+import (
+ "github.com/gin-gonic/gin"
+
+ "github.com/ShukeBta/MediaStationGo/internal/middleware"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func registerAuthedDownloadRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/downloads", requirePermission(svc, "can_manage_downloads"), listDownloadsHandler(svc))
+ authed.POST("/downloads", requirePermission(svc, "can_manage_downloads"), addDownloadHandler(svc))
+ authed.DELETE("/downloads/:hash", requirePermission(svc, "can_manage_downloads"), deleteDownloadHandler(svc))
+ authed.POST("/downloads/relocate", requirePermission(svc, "can_manage_downloads"), relocateDownloadHandler(svc))
+ authed.POST("/downloads/reload", requirePermission(svc, "can_manage_downloads"), reloadDownloadConfigHandler(svc))
+}
+
+func registerAuthedSubscriptionRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionsHandler(svc))
+ authed.GET("/subscriptions/history", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionHistoryHandler(svc))
+ authed.POST("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), createSubscriptionHandler(svc))
+ authed.DELETE("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), deleteSubscriptionHandler(svc))
+ authed.POST("/subscriptions/:id/restore", requirePermission(svc, "can_manage_subscriptions"), restoreSubscriptionHandler(svc))
+ authed.POST("/subscriptions/:id/run", requirePermission(svc, "can_manage_subscriptions"), runSubscriptionHandler(svc))
+}
+
+func registerAuthedStatsDiscoveryAndAIRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/stats", statsHandler(svc))
+ authed.GET("/tasks", middleware.AdminRequired(), tasksHandler(svc))
+
+ authed.GET("/discover/trending", requirePermission(svc, "can_view_discover"), trendingHandler(svc))
+ authed.GET("/discover/popular", requirePermission(svc, "can_view_discover"), popularHandler(svc))
+
+ authed.GET("/ai/status", requirePermission(svc, "can_use_ai"), aiStatusHandler(svc))
+ authed.POST("/ai/search", requirePermission(svc, "can_use_ai"), smartSearchHandler(svc))
+ authed.GET("/ai/recommend", requirePermission(svc, "can_use_ai"), aiRecommendHandler(svc))
+}
+
+func registerAuthedFileRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/files", middleware.AdminRequired(), browseFilesHandler(svc))
+ authed.POST("/files/folders", middleware.AdminRequired(), createFolderHandler(svc))
+ authed.PUT("/files/rename", middleware.AdminRequired(), renameFileHandler(svc))
+ authed.DELETE("/files", middleware.AdminRequired(), deleteFileHandler(svc))
+ authed.POST("/files/transfer", middleware.AdminRequired(), transferFileHandler(svc))
+}
+
+func registerAuthedDLNARoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/dlna/devices", dlnaListHandler(svc))
+ authed.POST("/dlna/cast", dlnaCastHandler(svc))
+}
+
+func registerAuthedSTRMRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.PUT("/media/:id/strm", middleware.AdminRequired(), setSTRMHandler(svc))
+ authed.DELETE("/media/:id/strm", middleware.AdminRequired(), clearSTRMHandler(svc))
+ authed.POST("/strm/import", middleware.AdminRequired(), importSTRMHandler(svc))
+ authed.POST("/strm/generate", middleware.AdminRequired(), generateSTRMHandler(svc))
+}
+
+func registerAuthedDuplicateRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/duplicates", middleware.AdminRequired(), listDuplicatesHandler(svc))
+ authed.POST("/duplicates/scan", middleware.AdminRequired(), detectDuplicatesHandler(svc))
+ authed.POST("/duplicates/unmark", middleware.AdminRequired(), unmarkDuplicatesHandler(svc))
+}
+
+func registerAuthedSiteRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ siteHandler := NewSiteHandler(svc)
+ authed.GET("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.ListSites)
+ authed.GET("/sites/types", requirePermission(svc, "can_manage_sites"), siteHandler.GetSiteTypes)
+ authed.GET("/sites/auth-types", requirePermission(svc, "can_manage_sites"), siteHandler.GetAuthTypes)
+ authed.POST("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.CreateSite)
+ authed.GET("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.GetSite)
+ authed.PUT("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.UpdateSite)
+ authed.DELETE("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.DeleteSite)
+ authed.POST("/sites/:id/test", requirePermission(svc, "can_manage_sites"), siteHandler.TestSite)
+ authed.GET("/sites/search", requirePermission(svc, "can_manage_sites"), siteSearchHandler(svc))
+}
+
+func registerAuthedRecycleAndRealtimeRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
+ authed.POST("/recycle/restore", middleware.AdminRequired(), restoreMediaBatchHandler(svc))
+ authed.POST("/recycle/purge", middleware.AdminRequired(), purgeMediaBatchHandler(svc))
+
+ authed.GET("/ws", wsHandler(svc))
+ authed.GET("/events", sseHandler(svc))
+}
+
+func registerAuthedSchedulerRoutes(authed *gin.RouterGroup, svc *service.Container) {
+ authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
+ authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
+ authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
+}
diff --git a/internal/handler/routes_authenticated_test.go b/internal/handler/routes_authenticated_test.go
new file mode 100644
index 0000000..0b459e1
--- /dev/null
+++ b/internal/handler/routes_authenticated_test.go
@@ -0,0 +1,46 @@
+package handler
+
+import (
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/service"
+)
+
+func TestAuthenticatedRouteSurfacesAreRegistered(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ router := gin.New()
+ Register(router, &config.Config{
+ Secrets: config.SecretsConfig{JWTSecret: "test-secret"},
+ }, zap.NewNop(), &service.Container{Log: zap.NewNop()})
+
+ routes := map[string]bool{}
+ for _, route := range router.Routes() {
+ routes[route.Method+" "+route.Path] = true
+ }
+
+ for _, want := range []string{
+ "GET /api/me",
+ "GET /api/auth/permissions",
+ "GET /api/libraries",
+ "GET /api/media",
+ "GET /api/stream/:id",
+ "GET /api/storage",
+ "GET /api/downloads",
+ "GET /api/subscriptions",
+ "GET /api/sites/search",
+ "GET /api/watch-history",
+ "GET /api/discover/feed",
+ "GET /api/playback/:id/info",
+ "GET /api/download/tasks",
+ "GET /api/admin/assistant/history",
+ } {
+ if !routes[want] {
+ t.Fatalf("%s route is not registered", want)
+ }
+ }
+}
diff --git a/internal/handler/series.go b/internal/handler/series.go
index 92cdf35..5ee1280 100644
--- a/internal/handler/series.go
+++ b/internal/handler/series.go
@@ -8,6 +8,7 @@ package handler
import (
"net/http"
"sort"
+ "strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
@@ -31,15 +32,20 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
return
}
}
- var rows []model.Media
- err := svc.Repo.DB.Where(&model.Media{LibraryID: libID}).
- Order("season_num asc, episode_num asc").
- Find(&rows).Error
- if err != nil && err != gorm.ErrRecordNotFound {
- c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
- return
- }
visibility := mediaVisibilityForRequest(c, svc)
+ var rows []model.Media
+ const pageSize = 2000
+ for page := 1; ; page++ {
+ pageRows, total, err := svc.Media.ListMediaVisible(c.Request.Context(), libID, page, pageSize, visibility)
+ if err != nil && err != gorm.ErrRecordNotFound {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ rows = append(rows, pageRows...)
+ if int64(len(rows)) >= total || len(pageRows) < pageSize {
+ break
+ }
+ }
buckets := make(map[int][]model.Media)
for _, r := range rows {
if !visibility.Allows(&r) {
@@ -55,3 +61,65 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusOK, gin.H{"seasons": out})
}
}
+
+func listLibrarySeriesHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ libID := c.Param("id")
+ if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
+ if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
+ c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
+ return
+ }
+ }
+ items, total, err := svc.Media.ListLibrarySeriesCards(c.Request.Context(), libID, mediaVisibilityForRequest(c, svc))
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
+ size, _ := strconv.Atoi(c.DefaultQuery("page_size", "500"))
+ if page < 1 {
+ page = 1
+ }
+ if size <= 0 || size > 1000 {
+ size = 500
+ }
+ start := (page - 1) * size
+ if start > len(items) {
+ start = len(items)
+ }
+ end := start + size
+ if end > len(items) {
+ end = len(items)
+ }
+ c.JSON(http.StatusOK, gin.H{
+ "items": items[start:end],
+ "total": total,
+ "page": page,
+ "page_size": size,
+ })
+ }
+}
+
+func listLibrarySeriesEpisodesHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ libID := c.Param("id")
+ key := c.Query("key")
+ if key == "" {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "key is required"})
+ return
+ }
+ if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
+ if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
+ c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
+ return
+ }
+ }
+ items, err := svc.Media.ListLibrarySeriesEpisodes(c.Request.Context(), libID, key, mediaVisibilityForRequest(c, svc))
+ if err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
+ }
+}
diff --git a/internal/handler/storage_config.go b/internal/handler/storage_config.go
index 08231dc..481eae5 100644
--- a/internal/handler/storage_config.go
+++ b/internal/handler/storage_config.go
@@ -1,4 +1,4 @@
-// Package handler — Alist / S3 / WebDAV storage config endpoints.
+// Package handler — external storage config endpoints.
package handler
import (
@@ -25,6 +25,10 @@ func listStorageConfigsHandler(svc *service.Container) gin.HandlerFunc {
// getStorageConfigHandler returns one config (with the decrypted body).
func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ if !service.IsAdminStorageConfigurable(c.Param("type")) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
+ return
+ }
row, err := svc.StorageCfg.Get(c.Request.Context(), c.Param("type"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -42,6 +46,10 @@ func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
// the type via URL and the body as a JSON object.
func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ if !service.IsAdminStorageConfigurable(c.Param("type")) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
+ return
+ }
var in service.StorageInput
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -63,6 +71,10 @@ func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
// testStorageConfigHandler probes an unsaved config.
func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ if !service.IsAdminStorageConfigurable(c.Param("type")) {
+ c.JSON(http.StatusBadRequest, gin.H{"ok": false, "error": "unsupported storage type"})
+ return
+ }
var in service.StorageInput
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -80,6 +92,10 @@ func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
+ if !service.IsAdminStorageConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
+ return
+ }
row, err := svc.StorageCfg.Logout(c.Request.Context(), typ)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -94,6 +110,10 @@ func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
func storageUploadLocalHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ if !service.IsAdminStorageConfigurable(c.Param("type")) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
+ return
+ }
var req service.CloudUploadInput
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
diff --git a/internal/handler/storage_config_test.go b/internal/handler/storage_config_test.go
new file mode 100644
index 0000000..c4966cc
--- /dev/null
+++ b/internal/handler/storage_config_test.go
@@ -0,0 +1,28 @@
+package handler
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+)
+
+func TestStorageConfigHandlersRejectQuark(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ router.PUT("/admin/storage/:type", saveStorageConfigHandler(nil))
+
+ req := httptest.NewRequest(http.MethodPut, "/admin/storage/quark", strings.NewReader(`{"type":"quark","config":{"cookie":"x"}}`))
+ req.Header.Set("Content-Type", "application/json")
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
+ }
+ if !strings.Contains(w.Body.String(), "unsupported storage type") {
+ t.Fatalf("body = %s, want unsupported storage type", w.Body.String())
+ }
+}
diff --git a/internal/handler/streaming.go b/internal/handler/streaming.go
index c9541de..7073824 100644
--- a/internal/handler/streaming.go
+++ b/internal/handler/streaming.go
@@ -4,6 +4,7 @@ package handler
import (
"context"
"errors"
+ "io"
"net/http"
"github.com/gin-gonic/gin"
@@ -79,16 +80,99 @@ func imageProxyHandler(svc *service.Container) gin.HandlerFunc {
}
}
+func cloudArtworkProxyHandler(svc *service.Container) gin.HandlerFunc {
+ return func(c *gin.Context) {
+ typ := c.Param("type")
+ ref := c.Query("ref")
+ if !service.IsAdminCloudConfigurable(typ) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
+ return
+ }
+ if ref == "" || !isCloudImageRef(ref) {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "image ref required"})
+ return
+ }
+ if svc == nil || svc.ImageProxy == nil {
+ c.JSON(http.StatusServiceUnavailable, gin.H{"error": "image proxy unavailable"})
+ return
+ }
+ stableKey := typ + ":" + ref
+ if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, stableKey) {
+ return
+ }
+ if svc.StorageCfg == nil {
+ c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
+ return
+ }
+ link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
+ if err != nil {
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ return
+ }
+ if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, stableKey, link); err != nil {
+ c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
+ return
+ }
+ }
+}
+
+type scrapeRequest struct {
+ EpisodeArtwork *bool `json:"episode_artwork"`
+ EpisodeImages *bool `json:"episode_images"`
+ RefreshMatched *bool `json:"refresh_matched"`
+ IncludeMatched *bool `json:"include_matched"`
+}
+
+func (r scrapeRequest) episodeArtworkOption() *bool {
+ if r.EpisodeImages != nil {
+ return r.EpisodeImages
+ }
+ return r.EpisodeArtwork
+}
+
+func (r scrapeRequest) includeMatchedOption() bool {
+ if r.IncludeMatched != nil {
+ return *r.IncludeMatched
+ }
+ if r.RefreshMatched != nil {
+ return *r.RefreshMatched
+ }
+ return false
+}
+
+func scrapeOptionsFromRequest(c *gin.Context, retryNoMatch bool) (service.ScrapeOptions, error) {
+ options := service.ScrapeOptions{RetryNoMatch: retryNoMatch}
+ if c.Request.Body == nil || c.Request.ContentLength == 0 {
+ return options, nil
+ }
+ var req scrapeRequest
+ if err := c.ShouldBindJSON(&req); err != nil {
+ if errors.Is(err, io.EOF) {
+ return options, nil
+ }
+ return options, err
+ }
+ options.EpisodeArtwork = req.episodeArtworkOption()
+ options.IncludeMatched = req.includeMatchedOption()
+ return options, nil
+}
+
// scrapeOneHandler enriches a single media via the configured scraper chain.
func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
+ options, err := scrapeOptionsFromRequest(c, true)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
+ return
+ }
+ options.IncludeMatched = true
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
task := startScrapeHTTPTask(svc, "手动刮削媒体", m.Title, m.Path)
- if err := svc.Scraper.EnrichOne(c.Request.Context(), m); err != nil {
+ if err := svc.Scraper.EnrichOneWithOptions(c.Request.Context(), m, options); err != nil {
finishHTTPTask(task, err, "scrape", "手动刮削媒体失败", nil, nil)
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -103,10 +187,16 @@ func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
}
}
-// scrapeLibraryHandler retries every pending/no_match media in a library.
+// scrapeLibraryHandler manually refreshes every scrapeable row in a library.
func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libID := c.Param("id")
+ options, err := scrapeOptionsFromRequest(c, true)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
+ return
+ }
+ options.IncludeMatched = true
var task *service.TaskHandle
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
task = startScrapeHTTPTask(svc, "手动刮削媒体库", lib.Name, lib.Path)
@@ -115,9 +205,16 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
}
// Run in the background so HTTP returns instantly; the WS hub
// pushes per-item progress on the "scrape" topic.
- go func(libID string, task *service.TaskHandle) {
- matched, err := svc.Scraper.EnrichLibrary(context.Background(), libID, true)
- metrics := map[string]int64{"matched": int64(matched)}
+ go func(libID string, task *service.TaskHandle, options service.ScrapeOptions) {
+ result, err := svc.Scraper.EnrichLibraryDetailedWithOptions(context.Background(), libID, options)
+ metrics := map[string]int64{
+ "matched": int64(result.Matched),
+ "processed": int64(result.Processed),
+ "candidates": int64(result.Candidates),
+ }
+ if result.Failed > 0 {
+ metrics["errors"] = int64(result.Failed)
+ }
stage := "completed"
message := "手动刮削媒体库结束"
if err != nil {
@@ -125,7 +222,7 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
message = "手动刮削媒体库失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
- }(libID, task)
+ }(libID, task, options)
c.JSON(http.StatusAccepted, gin.H{"status": "scraping"})
}
}
diff --git a/internal/handler/strm.go b/internal/handler/strm.go
index d4c2624..52dc22a 100644
--- a/internal/handler/strm.go
+++ b/internal/handler/strm.go
@@ -99,7 +99,7 @@ func importSTRMHandler(svc *service.Container) gin.HandlerFunc {
}
type generateSTRMReq struct {
- LibraryID string `json:"library_id" binding:"required"`
+ LibraryID string `json:"library_id"`
OutputDir string `json:"output_dir"`
BaseURL string `json:"base_url"`
Enabled bool `json:"enabled"`
@@ -122,7 +122,7 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
if baseURL == "" {
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
}
- res, err := strmSvc.GenerateForLibrary(c.Request.Context(), service.GenerateSTRMOptions{
+ options := service.GenerateSTRMOptions{
LibraryID: req.LibraryID,
OutputDir: req.OutputDir,
BaseURL: baseURL,
@@ -130,7 +130,14 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
Overwrite: req.Overwrite,
IncludeLocal: true,
PlaybackToken: strmPlaybackTokenForRequest(c, svc),
- })
+ }
+ var res *service.GenerateSTRMResult
+ var err error
+ if strings.TrimSpace(req.LibraryID) == "*" {
+ res, err = strmSvc.GenerateForAllLibraries(c.Request.Context(), options)
+ } else {
+ res, err = strmSvc.GenerateForLibrary(c.Request.Context(), options)
+ }
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
diff --git a/internal/handler/subscriptions.go b/internal/handler/subscriptions.go
index 53175f5..092e876 100644
--- a/internal/handler/subscriptions.go
+++ b/internal/handler/subscriptions.go
@@ -24,6 +24,8 @@ type subscriptionReq struct {
PosterURL string `json:"poster_url"`
BackdropURL string `json:"backdrop_url"`
Overview string `json:"overview"`
+ OriginalName string `json:"original_name"`
+ Year int `json:"year"`
Resolution string `json:"resolution"`
Quality string `json:"quality"`
Effects string `json:"effects"`
@@ -62,6 +64,8 @@ func createSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
PosterURL: req.PosterURL,
BackdropURL: req.BackdropURL,
Overview: req.Overview,
+ OriginalName: req.OriginalName,
+ Year: req.Year,
Resolution: req.Resolution,
Quality: req.Quality,
Effects: req.Effects,
diff --git a/internal/handler/system_extra.go b/internal/handler/system_extra.go
index 40a91cd..80d97ed 100644
--- a/internal/handler/system_extra.go
+++ b/internal/handler/system_extra.go
@@ -85,12 +85,11 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
{"key": "cloud.boot_scan_enabled", "type": "toggle", "label": "启动后立即扫描网盘"},
{"key": "cloud.upload_auto_enabled", "type": "toggle", "label": "启用自动转存"},
{"key": "cloud.upload_provider", "type": "select", "label": "转存目标", "options": []gin.H{
- {"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里/夸克)"},
- {"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里/夸克)"},
+ {"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里等)"},
+ {"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里等)"},
{"value": "alist", "label": "Alist(可桥接多网盘)"},
{"value": "webdav", "label": "WebDAV"},
{"value": "cloud115", "label": "115 原生(待接分片上传)"},
- {"value": "quark", "label": "夸克原生(待接分片上传)"},
}},
{"key": "cloud.upload_source_dir", "type": "text", "label": "本地源目录"},
{"key": "cloud.upload_dest_path", "type": "text", "label": "网盘目标目录"},
diff --git a/internal/handler/tasks.go b/internal/handler/tasks.go
index e685db4..3d4468d 100644
--- a/internal/handler/tasks.go
+++ b/internal/handler/tasks.go
@@ -8,16 +8,25 @@ package handler
import (
"net/http"
+ "time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
+const tasksLiveTorrentSnapshotMaxAge = 30 * time.Second
+
func tasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
- transcodes := svc.Transcoder.Active()
- _, torrents, _ := svc.Downloads.List(c.Request.Context())
+ var transcodes []service.ActiveJob
+ if svc.Transcoder != nil {
+ transcodes = svc.Transcoder.Active()
+ }
+ var torrents []service.QBitTorrent
+ if svc.Downloads != nil {
+ torrents = svc.Downloads.LiveTorrentSnapshot(tasksLiveTorrentSnapshotMaxAge)
+ }
background := service.TaskSnapshot{}
if svc.Tasks != nil {
background = svc.Tasks.Snapshot()
diff --git a/internal/middleware/middleware.go b/internal/middleware/middleware.go
index 2766216..7ef22b8 100644
--- a/internal/middleware/middleware.go
+++ b/internal/middleware/middleware.go
@@ -21,6 +21,11 @@ const (
CtxUserTier = "ctx_user_tier"
CtxTokenPurpose = "ctx_token_purpose"
CtxTokenMediaID = "ctx_token_media_id"
+
+ // AccessTokenCookieName carries the web access token for browser-managed
+ // resource requests such as
, which cannot attach Authorization.
+ AccessTokenCookieName = "msgo_access_token"
+ AccessTokenCookiePath = "/api"
)
// RequestLogger logs one structured line per request.
@@ -199,6 +204,7 @@ func AuthRequired(secret string) gin.HandlerFunc {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"code": 40304, "message": "token scope denied"})
return
}
+ syncAccessTokenCookie(c, raw, claims)
c.Set(CtxUserID, claims.UserID)
c.Set(CtxUserRole, claims.Role)
c.Set(CtxUserTier, claims.Tier)
@@ -208,6 +214,48 @@ func AuthRequired(secret string) gin.HandlerFunc {
}
}
+func syncAccessTokenCookie(c *gin.Context, raw string, claims *Claims) {
+ if c == nil || claims == nil || strings.TrimSpace(raw) == "" || strings.TrimSpace(claims.Purpose) != "" {
+ return
+ }
+ if existing, err := c.Cookie(AccessTokenCookieName); err == nil && existing == raw {
+ return
+ }
+ maxAge := int(time.Hour.Seconds())
+ expires := time.Now().Add(time.Hour)
+ if claims.ExpiresAt != nil {
+ expires = claims.ExpiresAt.Time
+ ttl := time.Until(expires)
+ if ttl <= 0 {
+ return
+ }
+ maxAge = int(ttl.Seconds())
+ if maxAge < 1 {
+ maxAge = 1
+ }
+ }
+ http.SetCookie(c.Writer, &http.Cookie{
+ Name: AccessTokenCookieName,
+ Value: raw,
+ Path: AccessTokenCookiePath,
+ MaxAge: maxAge,
+ Expires: expires,
+ HttpOnly: true,
+ SameSite: http.SameSiteLaxMode,
+ Secure: requestIsHTTPS(c),
+ })
+}
+
+func requestIsHTTPS(c *gin.Context) bool {
+ if c == nil || c.Request == nil {
+ return false
+ }
+ if c.Request.TLS != nil {
+ return true
+ }
+ return strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
+}
+
func scopedTokenAllowedForRequest(c *gin.Context, claims *Claims) bool {
if claims == nil || strings.TrimSpace(claims.Purpose) == "" {
return true
@@ -320,5 +368,8 @@ func extractToken(c *gin.Context) string {
return value
}
}
+ if cookie, err := c.Cookie(AccessTokenCookieName); err == nil {
+ return strings.TrimSpace(cookie)
+ }
return ""
}
diff --git a/internal/middleware/middleware_test.go b/internal/middleware/middleware_test.go
index 737a47c..33c99a1 100644
--- a/internal/middleware/middleware_test.go
+++ b/internal/middleware/middleware_test.go
@@ -4,8 +4,10 @@ import (
"net/http"
"net/http/httptest"
"testing"
+ "time"
"github.com/gin-gonic/gin"
+ "github.com/golang-jwt/jwt/v5"
)
func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
@@ -29,3 +31,161 @@ func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
t.Fatalf("Access-Control-Allow-Origin = %q, want *", got)
}
}
+
+func TestAuthRequiredAcceptsAccessTokenCookie(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ secret := "cookie-secret"
+ token := signedMiddlewareTestToken(t, secret, Claims{
+ UserID: "user-1",
+ Role: "admin",
+ RegisteredClaims: jwt.RegisteredClaims{
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
+ },
+ })
+
+ router := gin.New()
+ router.Use(AuthRequired(secret))
+ router.GET("/api/img", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"user_id": c.GetString(CtxUserID)})
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/api/img?url=https%3A%2F%2Fexample.test%2Fposter.jpg", nil)
+ req.AddCookie(&http.Cookie{Name: AccessTokenCookieName, Value: token})
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+}
+
+func TestAuthRequiredKeepsExplicitQueryTokenPriority(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ secret := "cookie-secret"
+ accountToken := signedMiddlewareTestToken(t, secret, Claims{
+ UserID: "user-1",
+ Role: "admin",
+ RegisteredClaims: jwt.RegisteredClaims{
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
+ },
+ })
+ scopedToken := signedMiddlewareTestToken(t, secret, Claims{
+ UserID: "user-1",
+ Role: "admin",
+ Purpose: "external_play",
+ MediaID: "media-1",
+ RegisteredClaims: jwt.RegisteredClaims{
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
+ },
+ })
+
+ router := gin.New()
+ router.Use(AuthRequired(secret))
+ router.GET("/api/me", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/api/me?token="+scopedToken, nil)
+ req.AddCookie(&http.Cookie{Name: AccessTokenCookieName, Value: accountToken})
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusForbidden {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+}
+
+func TestAuthRequiredSyncsAccessTokenCookieFromBearer(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ secret := "cookie-secret"
+ token := signedMiddlewareTestToken(t, secret, Claims{
+ UserID: "user-1",
+ Role: "admin",
+ RegisteredClaims: jwt.RegisteredClaims{
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
+ },
+ })
+
+ router := gin.New()
+ router.Use(AuthRequired(secret))
+ router.GET("/api/discover/feed", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/api/discover/feed", nil)
+ req.Header.Set("Authorization", "Bearer "+token)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ cookie := middlewareTestResponseCookie(t, w, AccessTokenCookieName)
+ if cookie.Value != token {
+ t.Fatal("synced cookie should contain the bearer token")
+ }
+ if cookie.Path != AccessTokenCookiePath {
+ t.Fatalf("cookie path = %q, want %q", cookie.Path, AccessTokenCookiePath)
+ }
+ if !cookie.HttpOnly || cookie.SameSite != http.SameSiteLaxMode {
+ t.Fatalf("cookie flags not suitable: httpOnly=%v sameSite=%v", cookie.HttpOnly, cookie.SameSite)
+ }
+}
+
+func TestAuthRequiredDoesNotSyncScopedPlaybackTokenCookie(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ secret := "cookie-secret"
+ token := signedMiddlewareTestToken(t, secret, Claims{
+ UserID: "user-1",
+ Role: "admin",
+ Purpose: "external_play",
+ MediaID: "media-1",
+ RegisteredClaims: jwt.RegisteredClaims{
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
+ },
+ })
+
+ router := gin.New()
+ router.Use(AuthRequired(secret))
+ router.GET("/api/stream/media-1", func(c *gin.Context) {
+ c.Status(http.StatusNoContent)
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/api/stream/media-1?token="+token, nil)
+ w := httptest.NewRecorder()
+ router.ServeHTTP(w, req)
+
+ if w.Code != http.StatusNoContent {
+ t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
+ }
+ if cookie := optionalMiddlewareTestResponseCookie(w, AccessTokenCookieName); cookie != nil {
+ t.Fatalf("scoped playback token should not be synced as web cookie: %#v", cookie)
+ }
+}
+
+func signedMiddlewareTestToken(t *testing.T, secret string, claims Claims) string {
+ t.Helper()
+ token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
+ if err != nil {
+ t.Fatalf("sign token: %v", err)
+ }
+ return token
+}
+
+func middlewareTestResponseCookie(t *testing.T, w *httptest.ResponseRecorder, name string) *http.Cookie {
+ t.Helper()
+ cookie := optionalMiddlewareTestResponseCookie(w, name)
+ if cookie == nil {
+ t.Fatalf("missing response cookie %q", name)
+ }
+ return cookie
+}
+
+func optionalMiddlewareTestResponseCookie(w *httptest.ResponseRecorder, name string) *http.Cookie {
+ for _, cookie := range w.Result().Cookies() {
+ if cookie.Name == name {
+ return cookie
+ }
+ }
+ return nil
+}
diff --git a/internal/model/bot.go b/internal/model/bot.go
index 9bf345f..4a43ebd 100644
--- a/internal/model/bot.go
+++ b/internal/model/bot.go
@@ -97,6 +97,9 @@ type UserDevice struct {
LastPlayAt *time.Time `gorm:"index" json:"last_play_at,omitempty"`
Warnings int `gorm:"default:0" json:"warnings"` // 指纹不匹配累计告警次数
Kicked bool `gorm:"default:false" json:"kicked"` // 被一键踢下线(强制重新登录)
+ Realtime bool `gorm:"-" json:"realtime,omitempty"`
+ Online bool `gorm:"-" json:"online,omitempty"`
+ Playing bool `gorm:"-" json:"playing,omitempty"`
}
// BeforeCreate 生成 UUID。
diff --git a/internal/model/model.go b/internal/model/model.go
index 1c483d5..f34a1c5 100644
--- a/internal/model/model.go
+++ b/internal/model/model.go
@@ -51,10 +51,12 @@ type User struct {
// ShareWarnings counts anti-account-sharing warnings, mainly device
// fingerprint mismatches. Once it exceeds the configured threshold a
// re-offence disables the account until an admin re-enables it.
- ShareWarnings int `gorm:"default:0" json:"share_warnings"`
- LastShareWarnAt *time.Time `json:"last_share_warn_at,omitempty"`
- IsDefaultAdmin bool `gorm:"-" json:"is_default_admin,omitempty"`
- IsProtected bool `gorm:"-" json:"is_protected,omitempty"`
+ ShareWarnings int `gorm:"default:0" json:"share_warnings"`
+ LastShareWarnAt *time.Time `json:"last_share_warn_at,omitempty"`
+ IsDefaultAdmin bool `gorm:"-" json:"is_default_admin,omitempty"`
+ IsProtected bool `gorm:"-" json:"is_protected,omitempty"`
+ RealtimeOnline bool `gorm:"-" json:"realtime_online,omitempty"`
+ RealtimeDeviceCount int `gorm:"-" json:"realtime_device_count,omitempty"`
}
// Library 表示用户定义的媒体根目录。
@@ -73,6 +75,7 @@ type Media struct {
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
Title string `gorm:"size:255;not null" json:"title"`
OriginalName string `gorm:"size:255" json:"original_name,omitempty"`
+ EpisodeTitle string `gorm:"size:255" json:"episode_title,omitempty"`
Path string `gorm:"uniqueIndex;size:1024;not null" json:"path"`
SizeBytes int64 `json:"size_bytes"`
DurationSec int `json:"duration_sec"`
@@ -200,17 +203,17 @@ type PlaylistItem struct {
// DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。
type DownloadTask struct {
Base
- UserID string `gorm:"index;size:36" json:"user_id"`
- SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"`
- Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / http
- URL string `gorm:"size:2048;not null" json:"-"`
- Title string `gorm:"size:512" json:"title,omitempty"`
- PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
- BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
- Overview string `gorm:"type:text" json:"overview,omitempty"`
- SavePath string `gorm:"size:1024" json:"save_path"`
- MediaType string `gorm:"size:16" json:"media_type,omitempty"`
- MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
+ UserID string `gorm:"index;size:36" json:"user_id"`
+ SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"`
+ Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / http
+ URL string `gorm:"size:2048;not null" json:"-"`
+ Title string `gorm:"size:512" json:"title,omitempty"`
+ PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
+ BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
+ Overview string `gorm:"type:text" json:"overview,omitempty"`
+ SavePath string `gorm:"size:1024" json:"save_path"`
+ MediaType string `gorm:"size:16" json:"media_type,omitempty"`
+ MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
@@ -228,38 +231,38 @@ type DownloadTask struct {
// Subscription 是自动化规则,轮询 RSS 源并将匹配种子排队到配置的下载客户端。
type Subscription struct {
Base
- UserID string `gorm:"index;size:36" json:"user_id"`
- Name string `gorm:"size:128;not null" json:"name"`
- FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
- Filter string `gorm:"size:512" json:"filter"`
- MediaType string `gorm:"size:16" json:"media_type,omitempty"`
- MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
- SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
- SearchMode string `gorm:"size:16;default:keyword" json:"search_mode,omitempty"` // keyword / imdb
- IMDBID string `gorm:"size:32" json:"imdb_id,omitempty"`
- Source string `gorm:"size:32" json:"source,omitempty"`
- PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
- BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
- Overview string `gorm:"type:text" json:"overview,omitempty"`
+ UserID string `gorm:"index;size:36" json:"user_id"`
+ Name string `gorm:"size:128;not null" json:"name"`
+ FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
+ Filter string `gorm:"size:512" json:"filter"`
+ MediaType string `gorm:"size:16" json:"media_type,omitempty"`
+ MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
+ SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
+ SearchMode string `gorm:"size:16;default:keyword" json:"search_mode,omitempty"` // keyword / imdb
+ IMDBID string `gorm:"size:32" json:"imdb_id,omitempty"`
+ Source string `gorm:"size:32" json:"source,omitempty"`
+ PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
+ BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
+ Overview string `gorm:"type:text" json:"overview,omitempty"`
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
- OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
- OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
- Year int `json:"year,omitempty"`
- Rating float32 `json:"rating,omitempty"`
- Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
- Resolution string `gorm:"size:32" json:"resolution,omitempty"` // 2160p / 1080p / 720p / best
- Quality string `gorm:"size:64" json:"quality,omitempty"` // remux / bluray / web-dl / hdtv
- Effects string `gorm:"size:128" json:"effects,omitempty"` // hdr,dolby-vision,atmos
- ReleaseGroups string `gorm:"size:255" json:"release_groups,omitempty"` // comma separated
- ExcludeWords string `gorm:"size:255" json:"exclude_words,omitempty"` // comma separated
- WashEnabled bool `gorm:"default:false" json:"wash_enabled"`
- WashPriority string `gorm:"size:32" json:"wash_priority,omitempty"` // balanced / resolution / quality / effects / seeders
- TotalEpisodes int `gorm:"default:0" json:"total_episodes,omitempty"`
- Priority int `gorm:"default:50" json:"priority,omitempty"` // lower is earlier when schedulers sort later
- Enabled bool `gorm:"default:true" json:"enabled"`
- LastRunAt *time.Time `json:"last_run_at,omitempty"`
- ArchivedAt *time.Time `gorm:"index" json:"archived_at,omitempty"`
- ArchiveReason string `gorm:"size:255" json:"archive_reason,omitempty"`
+ OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
+ OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
+ Year int `json:"year,omitempty"`
+ Rating float32 `json:"rating,omitempty"`
+ Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
+ Resolution string `gorm:"size:32" json:"resolution,omitempty"` // 2160p / 1080p / 720p / best
+ Quality string `gorm:"size:64" json:"quality,omitempty"` // remux / bluray / web-dl / hdtv
+ Effects string `gorm:"size:128" json:"effects,omitempty"` // hdr,dolby-vision,atmos
+ ReleaseGroups string `gorm:"size:255" json:"release_groups,omitempty"` // comma separated
+ ExcludeWords string `gorm:"size:255" json:"exclude_words,omitempty"` // comma separated
+ WashEnabled bool `gorm:"default:false" json:"wash_enabled"`
+ WashPriority string `gorm:"size:32" json:"wash_priority,omitempty"` // balanced / resolution / quality / effects / seeders
+ TotalEpisodes int `gorm:"default:0" json:"total_episodes,omitempty"`
+ Priority int `gorm:"default:50" json:"priority,omitempty"` // lower is earlier when schedulers sort later
+ Enabled bool `gorm:"default:true" json:"enabled"`
+ LastRunAt *time.Time `json:"last_run_at,omitempty"`
+ ArchivedAt *time.Time `gorm:"index" json:"archived_at,omitempty"`
+ ArchiveReason string `gorm:"size:255" json:"archive_reason,omitempty"`
DownloadedEpisodes int `gorm:"-" json:"downloaded_episodes,omitempty"`
LocalMediaCount int `gorm:"-" json:"local_media_count,omitempty"`
diff --git a/internal/repository/access_log_repository.go b/internal/repository/access_log_repository.go
new file mode 100644
index 0000000..6d29886
--- /dev/null
+++ b/internal/repository/access_log_repository.go
@@ -0,0 +1,24 @@
+package repository
+
+import (
+ "context"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// AccessLogRepository persists model.AccessLog records.
+type AccessLogRepository struct{ db *gorm.DB }
+
+// Create inserts one structured audit-trail entry.
+func (r *AccessLogRepository) Create(ctx context.Context, l *model.AccessLog) error {
+ return r.db.WithContext(ctx).Create(l).Error
+}
+
+// Recent returns the latest access-log entries (admin Activity panel).
+func (r *AccessLogRepository) Recent(ctx context.Context, limit int) ([]model.AccessLog, error) {
+ var rows []model.AccessLog
+ err := r.db.WithContext(ctx).Order("created_at desc").Limit(limit).Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/api_config_repository.go b/internal/repository/api_config_repository.go
new file mode 100644
index 0000000..e5fc2cb
--- /dev/null
+++ b/internal/repository/api_config_repository.go
@@ -0,0 +1,78 @@
+package repository
+
+import (
+ "context"
+ "errors"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// ApiConfigRepository persists model.ApiConfig records.
+type ApiConfigRepository struct{ db *gorm.DB }
+
+// Create inserts a new API config record.
+func (r *ApiConfigRepository) Create(ctx context.Context, c *model.ApiConfig) error {
+ return r.db.WithContext(ctx).Create(c).Error
+}
+
+// FindByProvider returns the API config for a provider, or (nil, nil).
+func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
+ var c model.ApiConfig
+ err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &c, nil
+}
+
+// List returns all API configs.
+func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) {
+ var rows []model.ApiConfig
+ err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error
+ return rows, err
+}
+
+// Upsert creates or updates an API config.
+func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.ApiConfig) error {
+ return r.db.WithContext(ctx).Where("provider = ?", c.Provider).
+ Assign(model.ApiConfig{
+ Base: model.Base{UpdatedAt: time.Now()},
+ APIKey: c.APIKey,
+ BaseURL: c.BaseURL,
+ Extra: c.Extra,
+ Enabled: c.Enabled,
+ }).FirstOrCreate(c).Error
+}
+
+// Update updates an API config.
+func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) error {
+ return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
+ Where("provider = ?", c.Provider).Updates(map[string]any{
+ "api_key": c.APIKey,
+ "base_url": c.BaseURL,
+ "extra": c.Extra,
+ "enabled": c.Enabled,
+ "updated_at": time.Now(),
+ }).Error
+}
+
+// Delete removes an API config.
+func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error {
+ return r.db.WithContext(ctx).Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
+}
+
+// UpdateTestResult 更新测试结果。
+func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error {
+ now := time.Now()
+ return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
+ Where("provider = ?", provider).Updates(map[string]any{
+ "test_result": result,
+ "last_tested_at": &now,
+ }).Error
+}
diff --git a/internal/repository/download_repository.go b/internal/repository/download_repository.go
new file mode 100644
index 0000000..041a557
--- /dev/null
+++ b/internal/repository/download_repository.go
@@ -0,0 +1,24 @@
+package repository
+
+import (
+ "context"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// DownloadRepository persists model.DownloadTask records.
+type DownloadRepository struct{ db *gorm.DB }
+
+// Create inserts a new download task.
+func (r *DownloadRepository) Create(ctx context.Context, t *model.DownloadTask) error {
+ return r.db.WithContext(ctx).Create(t).Error
+}
+
+// List returns all download tasks (admin view).
+func (r *DownloadRepository) List(ctx context.Context) ([]model.DownloadTask, error) {
+ var rows []model.DownloadTask
+ err := r.db.WithContext(ctx).Order("created_at desc").Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/favorite_repository.go b/internal/repository/favorite_repository.go
new file mode 100644
index 0000000..33150ea
--- /dev/null
+++ b/internal/repository/favorite_repository.go
@@ -0,0 +1,34 @@
+package repository
+
+import (
+ "context"
+ "errors"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// FavoriteRepository persists model.Favorite records.
+type FavoriteRepository struct{ db *gorm.DB }
+
+// Toggle flips the favourite flag for (user, media). Returns the new state.
+func (r *FavoriteRepository) Toggle(ctx context.Context, userID, mediaID string) (bool, error) {
+ var f model.Favorite
+ err := r.db.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ fav := model.Favorite{UserID: userID, MediaID: mediaID}
+ return true, r.db.WithContext(ctx).Create(&fav).Error
+ }
+ if err != nil {
+ return false, err
+ }
+ return false, r.db.WithContext(ctx).Delete(&f).Error
+}
+
+// ListByUser returns all favourite media IDs for a user.
+func (r *FavoriteRepository) ListByUser(ctx context.Context, userID string) ([]model.Favorite, error) {
+ var rows []model.Favorite
+ err := r.db.WithContext(ctx).Where("user_id = ?", userID).Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/history_repository.go b/internal/repository/history_repository.go
new file mode 100644
index 0000000..acee023
--- /dev/null
+++ b/internal/repository/history_repository.go
@@ -0,0 +1,41 @@
+package repository
+
+import (
+ "context"
+ "errors"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// HistoryRepository persists model.PlaybackHistory entries. The application
+// upserts on (UserID, MediaID) so resume always reads the latest position.
+type HistoryRepository struct{ db *gorm.DB }
+
+// Upsert atomically inserts/updates the resume position.
+func (r *HistoryRepository) Upsert(ctx context.Context, h *model.PlaybackHistory) error {
+ var existing model.PlaybackHistory
+ err := r.db.WithContext(ctx).
+ Where("user_id = ? AND media_id = ?", h.UserID, h.MediaID).
+ First(&existing).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return r.db.WithContext(ctx).Create(h).Error
+ }
+ if err != nil {
+ return err
+ }
+ existing.PositionMs = h.PositionMs
+ existing.DurationMs = h.DurationMs
+ existing.WatchedAt = h.WatchedAt
+ existing.Completed = h.Completed
+ return r.db.WithContext(ctx).Save(&existing).Error
+}
+
+// ListByUser returns the most recent history rows for the user.
+func (r *HistoryRepository) ListByUser(ctx context.Context, userID string, limit int) ([]model.PlaybackHistory, error) {
+ var rows []model.PlaybackHistory
+ err := r.db.WithContext(ctx).Where("user_id = ?", userID).
+ Order("watched_at desc").Limit(limit).Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/library_repository.go b/internal/repository/library_repository.go
new file mode 100644
index 0000000..e2a3bd3
--- /dev/null
+++ b/internal/repository/library_repository.go
@@ -0,0 +1,44 @@
+package repository
+
+import (
+ "context"
+ "errors"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// LibraryRepository persists model.Library records.
+type LibraryRepository struct{ db *gorm.DB }
+
+// Create persists a new library row.
+func (r *LibraryRepository) Create(ctx context.Context, l *model.Library) error {
+ return r.db.WithContext(ctx).Create(l).Error
+}
+
+// List returns all enabled+disabled libraries.
+func (r *LibraryRepository) List(ctx context.Context) ([]model.Library, error) {
+ var ls []model.Library
+ err := r.db.WithContext(ctx).Order("created_at asc").Find(&ls).Error
+ return ls, err
+}
+
+// FindByID returns the library, or (nil, nil) when missing.
+func (r *LibraryRepository) FindByID(ctx context.Context, id string) (*model.Library, error) {
+ var l model.Library
+ err := r.db.WithContext(ctx).Where("id = ?", id).First(&l).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &l, nil
+}
+
+// Delete removes a library and (soft) cascades to its media via repository
+// callers; we do not run CASCADE here to keep this method narrow.
+func (r *LibraryRepository) Delete(ctx context.Context, id string) error {
+ return r.db.WithContext(ctx).Delete(&model.Library{}, "id = ?", id).Error
+}
diff --git a/internal/repository/media_repository.go b/internal/repository/media_repository.go
new file mode 100644
index 0000000..ae6d41b
--- /dev/null
+++ b/internal/repository/media_repository.go
@@ -0,0 +1,313 @@
+package repository
+
+import (
+ "context"
+ "errors"
+ "strings"
+ "sync"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// MediaRepository persists model.Media records.
+type MediaRepository struct {
+ db *gorm.DB
+
+ searchIndexOnce sync.Once
+ searchIndexAvailable bool
+ searchBackend MediaSearchBackend
+}
+
+type MediaSearchBackend interface {
+ SearchMediaIDs(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]string, int64, error)
+}
+
+type MediaSearchSyncBackend interface {
+ MediaSearchBackend
+ EnsureIndex(ctx context.Context) error
+ IndexMedia(ctx context.Context, rows []model.Media) error
+}
+
+func (r *MediaRepository) SetSearchBackend(backend MediaSearchBackend) {
+ if r != nil {
+ r.searchBackend = backend
+ }
+}
+
+// MediaQueryFilter is applied to user-facing media queries so NSFW items and
+// profile-restricted libraries are filtered in SQL instead of only in React.
+type MediaQueryFilter struct {
+ IncludeNSFW bool
+ AllowedLibraryIDs []string
+ HiddenLibraryIDs []string
+}
+
+func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
+ if !filter.IncludeNSFW {
+ q = q.Where("nsfw = ?", false)
+ }
+ if len(filter.HiddenLibraryIDs) > 0 {
+ q = q.Where("library_id NOT IN ?", filter.HiddenLibraryIDs)
+ }
+ if len(filter.AllowedLibraryIDs) > 0 {
+ q = q.Where("library_id IN ?", filter.AllowedLibraryIDs)
+ }
+ return q
+}
+
+// Upsert inserts or updates a media row keyed by Path (unique index).
+//
+// 重要:当一条行已经存在时,scanner 重扫只应该刷新文件级元数据
+// (时长、宽高、编码、容器、大小),不能把刮削器维护的字段(标题改写、
+// 海报、TMDb/Bangumi ID、scrape_status 等)覆盖回零值。
+//
+// 之前用 Assign(*m).FirstOrCreate(m) 会把整张零值结构体写回,导致:
+// 1. scrape_status 从 'matched' / 'no_match' 被清空成 ”;
+// 2. 新建行使 GORM `default:pending` 也得不到应用(因为 zero value 被
+// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
+// 永远捞不到数据。
+func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.upsert(ctx, m)
+ })
+}
+
+func (r *MediaRepository) upsert(ctx context.Context, m *model.Media) error {
+ var existing model.Media
+ err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ // 新行:保证 scrape_status 走 GORM default:pending(即留空让数据库填)。
+ if m.ScrapeStatus == "" {
+ m.ScrapeStatus = "pending"
+ }
+ if createErr := r.db.WithContext(ctx).Create(m).Error; createErr == nil {
+ r.indexMediaBestEffort(ctx, *m)
+ return nil
+ } else if retryErr := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil {
+ return createErr
+ }
+ }
+ if err != nil {
+ return err
+ }
+
+ // 已存在:仅刷新文件层面的字段。
+ updates := map[string]any{}
+ setIfChanged(updates, "size_bytes", existing.SizeBytes, m.SizeBytes)
+ setIfChanged(updates, "duration_sec", existing.DurationSec, m.DurationSec)
+ setIfChanged(updates, "width", existing.Width, m.Width)
+ setIfChanged(updates, "height", existing.Height, m.Height)
+ setIfChanged(updates, "video_codec", existing.VideoCodec, m.VideoCodec)
+ setIfChanged(updates, "audio_codec", existing.AudioCodec, m.AudioCodec)
+ setIfChanged(updates, "container", existing.Container, m.Container)
+ if existing.DeletedAt.Valid {
+ updates["deleted_at"] = nil
+ }
+ // 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。
+ if m.FileID != "" && m.FileID != existing.FileID {
+ updates["file_id"] = m.FileID
+ }
+ if m.Title != "" {
+ // scanner 给出的标题只是从路径推导,刮削后 title 已被替换为
+ // 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题,
+ // 避免覆盖刮削结果。
+ if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" {
+ setIfChanged(updates, "title", existing.Title, m.Title)
+ if m.Year > 0 {
+ setIfChanged(updates, "year", existing.Year, m.Year)
+ }
+ }
+ }
+ status := strings.TrimSpace(existing.ScrapeStatus)
+ canRefreshExternalIDs := status == "pending" || status == "" || status == "no_match" ||
+ m.ScrapeStatus == "matched" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://")
+ if canRefreshExternalIDs {
+ changedExternalID := false
+ if m.TMDbID > 0 && existing.TMDbID != m.TMDbID {
+ updates["tm_db_id"] = m.TMDbID
+ changedExternalID = true
+ }
+ if m.BangumiID > 0 && existing.BangumiID != m.BangumiID {
+ updates["bangumi_id"] = m.BangumiID
+ changedExternalID = true
+ }
+ if m.DoubanID != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(m.DoubanID) {
+ updates["douban_id"] = m.DoubanID
+ changedExternalID = true
+ }
+ if m.TheTVDBID != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(m.TheTVDBID) {
+ updates["thetvdb_id"] = m.TheTVDBID
+ changedExternalID = true
+ }
+ if m.Year > 0 && existing.Year <= 0 {
+ updates["year"] = m.Year
+ }
+ if changedExternalID && (status == "no_match" || status == "matched") && m.ScrapeStatus != "matched" {
+ updates["scrape_status"] = "pending"
+ }
+ }
+ if m.ScrapeStatus == "matched" {
+ setIfChanged(updates, "scrape_status", existing.ScrapeStatus, m.ScrapeStatus)
+ if m.OriginalName != "" {
+ setIfChanged(updates, "original_name", existing.OriginalName, m.OriginalName)
+ }
+ if m.EpisodeTitle != "" {
+ setIfChanged(updates, "episode_title", existing.EpisodeTitle, m.EpisodeTitle)
+ }
+ if m.PosterURL != "" {
+ setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
+ }
+ if m.BackdropURL != "" {
+ setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
+ }
+ if m.Overview != "" {
+ setIfChanged(updates, "overview", existing.Overview, m.Overview)
+ }
+ if m.Rating > 0 {
+ setIfChanged(updates, "rating", existing.Rating, m.Rating)
+ }
+ if m.Year > 0 {
+ setIfChanged(updates, "year", existing.Year, m.Year)
+ }
+ if m.TMDbID > 0 {
+ setIfChanged(updates, "tm_db_id", existing.TMDbID, m.TMDbID)
+ }
+ if m.BangumiID > 0 {
+ setIfChanged(updates, "bangumi_id", existing.BangumiID, m.BangumiID)
+ }
+ if m.DoubanID != "" {
+ setIfChanged(updates, "douban_id", existing.DoubanID, m.DoubanID)
+ }
+ if m.TheTVDBID != "" {
+ setIfChanged(updates, "thetvdb_id", existing.TheTVDBID, m.TheTVDBID)
+ }
+ if m.Languages != "" {
+ setIfChanged(updates, "languages", existing.Languages, m.Languages)
+ }
+ if m.Countries != "" {
+ setIfChanged(updates, "countries", existing.Countries, m.Countries)
+ }
+ if m.Genres != "" {
+ setIfChanged(updates, "genres", existing.Genres, m.Genres)
+ }
+ if m.NSFW && !existing.NSFW {
+ updates["nsfw"] = true
+ }
+ }
+ if m.PosterURL != "" {
+ setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
+ }
+ if m.BackdropURL != "" {
+ setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
+ }
+ // 云盘媒体:同一 cloud:// 文件可能先被父目录库扫描入库,之后用户按二级
+ // 分类重新挂载/扫描到更精确的分类库。此时让 library_id 迁移到当前扫描库,
+ // 否则媒体被钉死在旧库、新分类库里看不到(表现为"媒体部分消失")。
+ // 本地媒体物理位置固定:仅在原 library_id 为空时回填,不迁移。
+ if isCloudMediaPath := strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://"); m.LibraryID != "" && m.LibraryID != existing.LibraryID {
+ if isCloudMediaPath || existing.LibraryID == "" {
+ updates["library_id"] = m.LibraryID
+ }
+ }
+ if (m.SeasonNum > 0 || m.EpisodeNum > 0) && existing.SeasonNum != m.SeasonNum {
+ updates["season_num"] = m.SeasonNum
+ }
+ if m.EpisodeNum > 0 && existing.EpisodeNum != m.EpisodeNum {
+ updates["episode_num"] = m.EpisodeNum
+ }
+ if m.STRMURL != "" {
+ setIfChanged(updates, "strm_url", existing.STRMURL, m.STRMURL)
+ }
+
+ if len(updates) == 0 {
+ *m = existing
+ return nil
+ }
+ if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}).
+ Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
+ return err
+ }
+ // 回写 ID / 不可变字段,让 caller 拿到完整的现有行。
+ *m = existing
+ if fresh, err := r.FindByID(ctx, existing.ID); err == nil && fresh != nil {
+ r.indexMediaBestEffort(ctx, *fresh)
+ }
+ return nil
+}
+
+func (r *MediaRepository) indexMediaBestEffort(ctx context.Context, media model.Media) {
+ backend, ok := r.searchBackend.(MediaSearchSyncBackend)
+ if !ok {
+ return
+ }
+ _ = backend.IndexMedia(ctx, []model.Media{media})
+}
+
+func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
+ if current != next {
+ updates[key] = next
+ }
+}
+
+// FindByID returns the media row or (nil, nil).
+func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) {
+ var m model.Media
+ err := r.db.WithContext(ctx).Where("id = ?", id).First(&m).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &m, nil
+}
+
+// ListByLibrary returns paginated media items for a library.
+func (r *MediaRepository) ListByLibrary(ctx context.Context, libraryID string, offset, limit int) ([]model.Media, int64, error) {
+ return r.ListByLibraryFiltered(ctx, libraryID, offset, limit, MediaQueryFilter{IncludeNSFW: true})
+}
+
+func (r *MediaRepository) ListByLibraryFiltered(ctx context.Context, libraryID string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
+ return r.ListByLibrariesFiltered(ctx, []string{libraryID}, offset, limit, filter)
+}
+
+func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
+ var items []model.Media
+ var total int64
+ if len(libraryIDs) == 0 {
+ return items, 0, nil
+ }
+ q := r.db.WithContext(ctx).Model(&model.Media{})
+ if len(libraryIDs) == 1 {
+ q = q.Where("library_id = ?", libraryIDs[0])
+ } else {
+ q = q.Where("library_id IN ?", libraryIDs)
+ }
+ q = applyMediaQueryFilter(q, filter)
+ if err := q.Count(&total).Error; err != nil {
+ return nil, 0, err
+ }
+ // 多级排序消除"随机"观感:
+ // 1. year desc — 上映年份新→旧(用户期望的上映时间维度)
+ // 2. updated_at desc — 同年按最近更新(刮削/补集会刷新)
+ // 3. created_at desc — 再按入库时间
+ // 4. id desc — 稳定 tie-breaker:云盘批量扫描同批 created_at 相同时,
+ // 没有它 DB 返回顺序不确定,正是"随机排序"的根因。
+ err := q.Order("year DESC, updated_at DESC, created_at DESC, id DESC").
+ Offset(offset).Limit(limit).Find(&items).Error
+ return items, total, err
+}
+
+// DeleteByLibrary purges all media tied to a library.
+func (r *MediaRepository) DeleteByLibrary(ctx context.Context, libraryID string) error {
+ // FTS 行由 media 表上的触发器同步清理(软删/硬删都覆盖)。
+ return r.db.WithContext(ctx).Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
+}
+
+// PurgeByLibrary permanently removes media tied to a library. Used for virtual
+// cloud mounts where "remove mount" must not populate the recycle bin.
+func (r *MediaRepository) PurgeByLibrary(ctx context.Context, libraryID string) error {
+ return r.db.WithContext(ctx).Unscoped().Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
+}
diff --git a/internal/repository/media_search_repository.go b/internal/repository/media_search_repository.go
new file mode 100644
index 0000000..0437d62
--- /dev/null
+++ b/internal/repository/media_search_repository.go
@@ -0,0 +1,268 @@
+package repository
+
+import (
+ "context"
+ "strings"
+ "unicode"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// Search runs a LIKE search against the title field. Empty query returns the
+// most recently added items.
+func (r *MediaRepository) Search(ctx context.Context, query string, limit int) ([]model.Media, error) {
+ return r.SearchFiltered(ctx, query, limit, MediaQueryFilter{IncludeNSFW: true})
+}
+
+func (r *MediaRepository) SearchFiltered(ctx context.Context, query string, limit int, filter MediaQueryFilter) ([]model.Media, error) {
+ items, _, err := r.SearchFilteredPage(ctx, query, 0, limit, filter)
+ return items, err
+}
+
+func (r *MediaRepository) SearchFilteredPage(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
+ query = strings.TrimSpace(query)
+ if limit <= 0 {
+ limit = 50
+ }
+ if query != "" && r.searchBackend != nil {
+ if items, total, ok := r.searchFilteredBackend(ctx, query, offset, limit, filter); ok {
+ return items, total, nil
+ }
+ }
+ if query != "" {
+ if items, total, ok := r.searchFilteredFTS(ctx, query, offset, limit, filter); ok {
+ if total > 0 {
+ return items, total, nil
+ }
+ }
+ }
+ return r.searchFilteredLIKE(ctx, query, offset, limit, filter)
+}
+
+func (r *MediaRepository) searchFilteredBackend(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
+ ids, total, err := r.searchBackend.SearchMediaIDs(ctx, query, offset, limit, filter)
+ if err != nil {
+ return nil, 0, false
+ }
+ if len(ids) == 0 {
+ return []model.Media{}, total, true
+ }
+ var rows []model.Media
+ q := r.db.WithContext(ctx).Model(&model.Media{}).Where("id IN ?", ids)
+ q = applyMediaQueryFilter(q, filter)
+ if err := q.Find(&rows).Error; err != nil {
+ return nil, 0, false
+ }
+ byID := make(map[string]model.Media, len(rows))
+ for _, row := range rows {
+ byID[row.ID] = row
+ }
+ items := make([]model.Media, 0, len(ids))
+ for _, id := range ids {
+ if row, ok := byID[id]; ok {
+ items = append(items, row)
+ }
+ }
+ if len(items) == 0 && total > 0 {
+ return nil, 0, false
+ }
+ return items, total, true
+}
+
+func (r *MediaRepository) searchFilteredFTS(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
+ if !r.searchIndexEnabled(ctx) {
+ return nil, 0, false
+ }
+ ftsQuery := mediaFTSQuery(query)
+ if ftsQuery == "" {
+ return nil, 0, false
+ }
+ var total int64
+ var items []model.Media
+ q := r.db.WithContext(ctx).
+ Table("media").
+ Joins("JOIN media_search_fts ON media_search_fts.rowid = media.rowid").
+ Where("media.deleted_at IS NULL").
+ Where("media_search_fts MATCH ?", ftsQuery)
+ q = applyQualifiedMediaQueryFilter(q, filter)
+ if err := q.Count(&total).Error; err != nil {
+ return nil, 0, false
+ }
+ if total == 0 {
+ return items, 0, true
+ }
+ err := q.Select("media.*").Order("bm25(media_search_fts), media.created_at DESC").Offset(offset).Limit(limit).Find(&items).Error
+ if err != nil {
+ return nil, 0, false
+ }
+ return items, total, true
+}
+
+func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
+ var items []model.Media
+ var total int64
+ q := r.db.WithContext(ctx).Model(&model.Media{})
+ q = applyMediaQueryFilter(q, filter)
+ terms := mediaSearchTerms(query)
+ for _, term := range terms {
+ like := "%" + escapeLike(term) + "%"
+ q = q.Where(
+ "(title LIKE ? ESCAPE '\\' OR original_name LIKE ? ESCAPE '\\' OR path LIKE ? ESCAPE '\\' OR genres LIKE ? ESCAPE '\\')",
+ like, like, like, like,
+ )
+ }
+ if err := q.Count(&total).Error; err != nil {
+ return nil, 0, err
+ }
+ if query != "" {
+ prefix := escapeLike(query) + "%"
+ exact := query
+ q = q.Order(gorm.Expr(
+ "CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title LIKE ? ESCAPE '\\' THEN 2 WHEN original_name LIKE ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc",
+ exact, exact, prefix, prefix,
+ ))
+ } else {
+ q = q.Order("created_at desc")
+ }
+ err := q.Offset(offset).Limit(limit).Find(&items).Error
+ return items, total, err
+}
+
+func applyQualifiedMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
+ if !filter.IncludeNSFW {
+ q = q.Where("media.nsfw = ?", false)
+ }
+ if len(filter.HiddenLibraryIDs) > 0 {
+ q = q.Where("media.library_id NOT IN ?", filter.HiddenLibraryIDs)
+ }
+ if len(filter.AllowedLibraryIDs) > 0 {
+ q = q.Where("media.library_id IN ?", filter.AllowedLibraryIDs)
+ }
+ return q
+}
+
+func mediaFTSQuery(query string) string {
+ terms := mediaSearchTerms(query)
+ if len(terms) == 0 {
+ return ""
+ }
+ quoted := make([]string, 0, len(terms))
+ for _, term := range terms {
+ term = strings.ReplaceAll(term, `"`, `""`)
+ if term != "" {
+ quoted = append(quoted, `"`+term+`"`)
+ }
+ }
+ return strings.Join(quoted, " AND ")
+}
+
+func mediaSearchTerms(query string) []string {
+ query = strings.TrimSpace(query)
+ if query == "" {
+ return nil
+ }
+ fields := strings.FieldsFunc(query, func(r rune) bool {
+ return unicode.IsSpace(r) || unicode.IsPunct(r) || unicode.IsSymbol(r)
+ })
+ out := make([]string, 0, len(fields))
+ seen := map[string]struct{}{}
+ for _, field := range fields {
+ field = strings.TrimSpace(field)
+ if field == "" {
+ continue
+ }
+ lower := strings.ToLower(field)
+ if _, ok := seen[lower]; ok {
+ continue
+ }
+ seen[lower] = struct{}{}
+ out = append(out, field)
+ }
+ return out
+}
+
+func escapeLike(value string) string {
+ value = strings.ReplaceAll(value, `\`, `\\`)
+ value = strings.ReplaceAll(value, `%`, `\%`)
+ value = strings.ReplaceAll(value, `_`, `\_`)
+ return value
+}
+
+func (r *MediaRepository) BackfillSearchIndex(ctx context.Context, batchLimit int) (int64, error) {
+ if backend, ok := r.searchBackend.(MediaSearchSyncBackend); ok {
+ return r.backfillExternalSearchIndex(ctx, backend, batchLimit)
+ }
+ if batchLimit <= 0 {
+ batchLimit = 1000
+ }
+ if !r.searchIndexEnabled(ctx) {
+ return 0, nil
+ }
+ // 关键性能点:FTS5 普通列(含 UNINDEXED)不支持索引查找,按
+ // media_id 做 NOT EXISTS 是对 FTS 表的整表扫描,再叠加 ORDER BY
+ // 后每个批次都要对全部 media 行探测一遍——大库一次启动回填等于
+ // 上百亿次行访问,曾把 CPU 钉满数小时。v2 布局下 FTS 行 rowid 与
+ // media.rowid 对齐,NOT EXISTS 走 rowid 点查,且无需排序。
+ res := r.db.WithContext(ctx).Exec(`
+INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
+SELECT m.rowid, m.id, COALESCE(m.title, ''), COALESCE(m.original_name, ''), COALESCE(m.path, ''), COALESCE(m.genres, '')
+FROM media AS m
+WHERE m.deleted_at IS NULL
+ AND NOT EXISTS (
+ SELECT 1 FROM media_search_fts AS f WHERE f.rowid = m.rowid
+ )
+LIMIT ?
+`, batchLimit)
+ return res.RowsAffected, res.Error
+}
+
+func (r *MediaRepository) backfillExternalSearchIndex(ctx context.Context, backend MediaSearchSyncBackend, batchLimit int) (int64, error) {
+ if batchLimit <= 0 {
+ batchLimit = 1000
+ }
+ if err := backend.EnsureIndex(ctx); err != nil {
+ return 0, err
+ }
+ var lastID string
+ for {
+ var rows []model.Media
+ q := r.db.WithContext(ctx).
+ Model(&model.Media{}).
+ Where("deleted_at IS NULL")
+ if lastID != "" {
+ q = q.Where("id > ?", lastID)
+ }
+ if err := q.Order("id ASC").Limit(batchLimit).Find(&rows).Error; err != nil {
+ return 0, err
+ }
+ if len(rows) == 0 {
+ return 0, nil
+ }
+ if err := backend.IndexMedia(ctx, rows); err != nil {
+ return 0, err
+ }
+ lastID = rows[len(rows)-1].ID
+ if len(rows) < batchLimit {
+ return 0, nil
+ }
+ }
+}
+
+func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool {
+ if r == nil || r.db == nil {
+ return false
+ }
+ if r.db.Dialector == nil || r.db.Dialector.Name() != "sqlite" {
+ return false
+ }
+ r.searchIndexOnce.Do(func() {
+ var count int64
+ err := r.db.WithContext(ctx).
+ Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`).
+ Scan(&count).Error
+ r.searchIndexAvailable = err == nil && count > 0
+ })
+ return r.searchIndexAvailable
+}
diff --git a/internal/repository/permission_repository.go b/internal/repository/permission_repository.go
new file mode 100644
index 0000000..ae86172
--- /dev/null
+++ b/internal/repository/permission_repository.go
@@ -0,0 +1,59 @@
+package repository
+
+import (
+ "context"
+ "errors"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// PermissionRepository persists model.UserPermission records.
+type PermissionRepository struct{ db *gorm.DB }
+
+// Create inserts a new permission record.
+func (r *PermissionRepository) Create(ctx context.Context, p *model.UserPermission) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Create(p).Error
+ })
+}
+
+// FindByUserID returns the permission record for a user, or (nil, nil) when absent.
+func (r *PermissionRepository) FindByUserID(ctx context.Context, userID string) (*model.UserPermission, error) {
+ var p model.UserPermission
+ err := withSQLiteBusyRetry(ctx, func() error {
+ p = model.UserPermission{}
+ return r.db.WithContext(ctx).Where("user_id = ?", userID).First(&p).Error
+ })
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &p, nil
+}
+
+// Update updates permission fields for a user.
+func (r *PermissionRepository) Update(ctx context.Context, userID string, updates map[string]bool) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Model(&model.UserPermission{}).
+ Where("user_id = ?", userID).Updates(updates).Error
+ })
+}
+
+// Upsert creates or updates a permission record.
+func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
+ Assign(*p).FirstOrCreate(p).Error
+ })
+}
+
+// Delete removes a permission record.
+func (r *PermissionRepository) Delete(ctx context.Context, userID string) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
+ })
+}
diff --git a/internal/repository/playlist_repository.go b/internal/repository/playlist_repository.go
new file mode 100644
index 0000000..e61c1b5
--- /dev/null
+++ b/internal/repository/playlist_repository.go
@@ -0,0 +1,25 @@
+package repository
+
+import (
+ "context"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// PlaylistRepository persists model.Playlist + PlaylistItem.
+type PlaylistRepository struct{ db *gorm.DB }
+
+// Create inserts a new playlist.
+func (r *PlaylistRepository) Create(ctx context.Context, p *model.Playlist) error {
+ return r.db.WithContext(ctx).Create(p).Error
+}
+
+// ListByUser returns playlists owned by a user.
+func (r *PlaylistRepository) ListByUser(ctx context.Context, userID string) ([]model.Playlist, error) {
+ var rows []model.Playlist
+ err := r.db.WithContext(ctx).Where("user_id = ?", userID).
+ Order("created_at desc").Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/refresh_token_repository.go b/internal/repository/refresh_token_repository.go
new file mode 100644
index 0000000..d17bf53
--- /dev/null
+++ b/internal/repository/refresh_token_repository.go
@@ -0,0 +1,94 @@
+package repository
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/hex"
+ "errors"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// RefreshTokenRepository persists model.RefreshToken records.
+type RefreshTokenRepository struct{ db *gorm.DB }
+
+// Create inserts a new refresh token record.
+func (r *RefreshTokenRepository) Create(ctx context.Context, t *model.RefreshToken) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Create(t).Error
+ })
+}
+
+// FindByHash returns the refresh token matching the hash, or (nil, nil).
+func (r *RefreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) {
+ var t model.RefreshToken
+ err := withSQLiteBusyRetry(ctx, func() error {
+ t = model.RefreshToken{}
+ return r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&t).Error
+ })
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &t, nil
+}
+
+// RevokeByUserID revokes all refresh tokens for a user.
+func (r *RefreshTokenRepository) RevokeByUserID(ctx context.Context, userID string) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
+ Where("user_id = ?", userID).Update("revoked", true).Error
+ })
+}
+
+// RevokeOldestActiveByUserID keeps at most limit active refresh tokens for a
+// user by revoking the oldest non-expired, non-revoked tokens.
+func (r *RefreshTokenRepository) RevokeOldestActiveByUserID(ctx context.Context, userID string, limit int) error {
+ if limit < 1 {
+ limit = 1
+ }
+ return withSQLiteBusyRetry(ctx, func() error {
+ var tokens []model.RefreshToken
+ if err := r.db.WithContext(ctx).
+ Where("user_id = ? AND revoked = ? AND expires_at > ?", userID, false, time.Now()).
+ Order("created_at desc, id desc").
+ Find(&tokens).Error; err != nil {
+ return err
+ }
+ if len(tokens) <= limit {
+ return nil
+ }
+ ids := make([]string, 0, len(tokens)-limit)
+ for _, token := range tokens[limit:] {
+ ids = append(ids, token.ID)
+ }
+ return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
+ Where("id IN ?", ids).Update("revoked", true).Error
+ })
+}
+
+// DeleteExpired removes all expired refresh tokens.
+func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error
+ })
+}
+
+// Revoke revokes a specific refresh token.
+func (r *RefreshTokenRepository) Revoke(ctx context.Context, hash string) error {
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
+ Where("token_hash = ?", hash).Update("revoked", true).Error
+ })
+}
+
+// HashToken returns the SHA256 hash of a token.
+func HashToken(token string) string {
+ h := sha256.Sum256([]byte(token))
+ return hex.EncodeToString(h[:])
+}
diff --git a/internal/repository/repository.go b/internal/repository/repository.go
index ef6c93f..e94f42c 100644
--- a/internal/repository/repository.go
+++ b/internal/repository/repository.go
@@ -5,20 +5,7 @@
// 业务逻辑位于 internal/service。
package repository
-import (
- "context"
- "crypto/sha256"
- "encoding/hex"
- "errors"
- "strings"
- "sync"
- "time"
- "unicode"
-
- "gorm.io/gorm"
-
- "github.com/ShukeBta/MediaStationGo/internal/model"
-)
+import "gorm.io/gorm"
// Container 是所有 repositories 的注册表,注入到 services 中。
type Container struct {
@@ -79,1155 +66,3 @@ func New(db *gorm.DB) *Container {
UserDevice: &UserDeviceRepository{db: db},
}
}
-
-// ─── User ────────────────────────────────────────────────────────────────────
-
-// UserRepository persists model.User records.
-type UserRepository struct{ db *gorm.DB }
-
-// Create inserts a new user. Caller must pre-hash the password.
-func (r *UserRepository) Create(ctx context.Context, u *model.User) error {
- return r.db.WithContext(ctx).Create(u).Error
-}
-
-// ReleaseDeletedUsername renames soft-deleted rows that still hold a unique
-// username so the same account name can be created again.
-func (r *UserRepository) ReleaseDeletedUsername(ctx context.Context, username string) error {
- if username == "" {
- return nil
- }
- released := username + "__deleted__" + time.Now().Format("20060102150405.000000000")
- if len(released) > 64 {
- sum := sha256.Sum256([]byte(released))
- released = username
- if len(released) > 43 {
- released = released[:43]
- }
- released += "__deleted__" + hex.EncodeToString(sum[:])[:10]
- }
- return r.db.WithContext(ctx).Unscoped().
- Model(&model.User{}).
- Where("username = ? AND deleted_at IS NOT NULL", username).
- Update("username", released).Error
-}
-
-// FindByUsername returns the user matching username, or (nil, nil) when absent.
-func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) {
- var u model.User
- err := withSQLiteBusyRetry(ctx, func() error {
- u = model.User{}
- err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error
- if errors.Is(err, gorm.ErrRecordNotFound) && username != "" {
- err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error
- }
- return err
- })
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &u, nil
-}
-
-// FindByID returns the user with the matching primary key, or (nil, nil).
-func (r *UserRepository) FindByID(ctx context.Context, id string) (*model.User, error) {
- var u model.User
- err := withSQLiteBusyRetry(ctx, func() error {
- u = model.User{}
- return r.db.WithContext(ctx).Where("id = ?", id).First(&u).Error
- })
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &u, nil
-}
-
-// Count returns the total number of non-deleted users.
-func (r *UserRepository) Count(ctx context.Context) (int64, error) {
- var n int64
- err := r.db.WithContext(ctx).Model(&model.User{}).Count(&n).Error
- return n, err
-}
-
-// CountAdmins returns the number of users that hold the admin role.
-func (r *UserRepository) CountAdmins(ctx context.Context) (int64, error) {
- var n int64
- err := r.db.WithContext(ctx).Model(&model.User{}).
- Where("role = ?", "admin").Count(&n).Error
- return n, err
-}
-
-// FirstAdmin returns the earliest admin user. This row represents the protected
-// built-in/default administrator even if its username is later changed.
-func (r *UserRepository) FirstAdmin(ctx context.Context) (*model.User, error) {
- var u model.User
- err := r.db.WithContext(ctx).Where("role = ?", "admin").Order("created_at asc").First(&u).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &u, nil
-}
-
-// List returns all users ordered by creation time desc.
-func (r *UserRepository) List(ctx context.Context) ([]model.User, error) {
- var users []model.User
- err := r.db.WithContext(ctx).Order("created_at desc").Find(&users).Error
- return users, err
-}
-
-// UpdateFields applies a narrow set of user field updates.
-func (r *UserRepository) UpdateFields(ctx context.Context, id string, updates map[string]any) error {
- return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error
-}
-
-// UpdatePassword sets a new password hash and clears ForcePasswordReset.
-func (r *UserRepository) UpdatePassword(ctx context.Context, id, hash string) error {
- return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
- Updates(map[string]any{"password_hash": hash, "force_password_reset": false}).Error
-}
-
-// TouchLogin updates the last login timestamp.
-func (r *UserRepository) TouchLogin(ctx context.Context, id string) error {
- now := time.Now()
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
- Update("last_login_at", &now).Error
- })
-}
-
-// Delete removes a user (soft-delete via gorm.DeletedAt), releases the unique
-// username, and drops Telegram bindings so future re-created users bind cleanly.
-func (r *UserRepository) Delete(ctx context.Context, id string) error {
- return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
- var user model.User
- if err := tx.Where("id = ?", id).First(&user).Error; err != nil {
- return err
- }
- if err := tx.Unscoped().Where("user_id = ?", id).Delete(&model.TelegramBinding{}).Error; err != nil {
- return err
- }
- released := user.Username + "__deleted__" + time.Now().Format("20060102150405.000000000")
- if len(released) > 64 {
- sum := sha256.Sum256([]byte(user.ID + user.Username))
- base := user.Username
- if len(base) > 43 {
- base = base[:43]
- }
- released = base + "__deleted__" + hex.EncodeToString(sum[:])[:10]
- }
- if err := tx.Model(&model.User{}).Where("id = ?", id).Update("username", released).Error; err != nil {
- return err
- }
- return tx.Delete(&model.User{}, "id = ?", id).Error
- })
-}
-
-// ─── Library ─────────────────────────────────────────────────────────────────
-
-// LibraryRepository persists model.Library records.
-type LibraryRepository struct{ db *gorm.DB }
-
-// Create persists a new library row.
-func (r *LibraryRepository) Create(ctx context.Context, l *model.Library) error {
- return r.db.WithContext(ctx).Create(l).Error
-}
-
-// List returns all enabled+disabled libraries.
-func (r *LibraryRepository) List(ctx context.Context) ([]model.Library, error) {
- var ls []model.Library
- err := r.db.WithContext(ctx).Order("created_at asc").Find(&ls).Error
- return ls, err
-}
-
-// FindByID returns the library, or (nil, nil) when missing.
-func (r *LibraryRepository) FindByID(ctx context.Context, id string) (*model.Library, error) {
- var l model.Library
- err := r.db.WithContext(ctx).Where("id = ?", id).First(&l).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &l, nil
-}
-
-// Delete removes a library and (soft) cascades to its media via repository
-// callers — we do not run CASCADE here to keep this method narrow.
-func (r *LibraryRepository) Delete(ctx context.Context, id string) error {
- return r.db.WithContext(ctx).Delete(&model.Library{}, "id = ?", id).Error
-}
-
-// ─── Media ───────────────────────────────────────────────────────────────────
-
-// MediaRepository persists model.Media records.
-type MediaRepository struct {
- db *gorm.DB
-
- searchIndexOnce sync.Once
- searchIndexAvailable bool
- searchBackend MediaSearchBackend
-}
-
-type MediaSearchBackend interface {
- SearchMediaIDs(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]string, int64, error)
-}
-
-type MediaSearchSyncBackend interface {
- MediaSearchBackend
- EnsureIndex(ctx context.Context) error
- IndexMedia(ctx context.Context, rows []model.Media) error
-}
-
-func (r *MediaRepository) SetSearchBackend(backend MediaSearchBackend) {
- if r != nil {
- r.searchBackend = backend
- }
-}
-
-// MediaQueryFilter is applied to user-facing media queries so NSFW items and
-// profile-restricted libraries are filtered in SQL instead of only in React.
-type MediaQueryFilter struct {
- IncludeNSFW bool
- AllowedLibraryIDs []string
- HiddenLibraryIDs []string
-}
-
-func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
- if !filter.IncludeNSFW {
- q = q.Where("nsfw = ?", false)
- }
- if len(filter.HiddenLibraryIDs) > 0 {
- q = q.Where("library_id NOT IN ?", filter.HiddenLibraryIDs)
- }
- if len(filter.AllowedLibraryIDs) > 0 {
- q = q.Where("library_id IN ?", filter.AllowedLibraryIDs)
- }
- return q
-}
-
-// Upsert inserts or updates a media row keyed by Path (unique index).
-//
-// 重要:当一条行已经存在时,scanner 重扫只应该刷新文件级元数据
-// (时长、宽高、编码、容器、大小),不能把刮削器维护的字段(标题改写、
-// 海报、TMDb/Bangumi ID、scrape_status 等)覆盖回零值。
-//
-// 之前用 Assign(*m).FirstOrCreate(m) 会把整张零值结构体写回,导致:
-// 1. scrape_status 从 'matched' / 'no_match' 被清空成 ”;
-// 2. 新建行使 GORM `default:pending` 也得不到应用(因为 zero value 被
-// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
-// 永远捞不到数据。
-func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.upsert(ctx, m)
- })
-}
-
-func (r *MediaRepository) upsert(ctx context.Context, m *model.Media) error {
- var existing model.Media
- err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- // 新行:保证 scrape_status 走 GORM default:pending(即留空让数据库填)。
- if m.ScrapeStatus == "" {
- m.ScrapeStatus = "pending"
- }
- if createErr := r.db.WithContext(ctx).Create(m).Error; createErr == nil {
- r.indexMediaBestEffort(ctx, *m)
- return nil
- } else if retryErr := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil {
- return createErr
- }
- }
- if err != nil {
- return err
- }
-
- // 已存在:仅刷新文件层面的字段。
- updates := map[string]any{}
- setIfChanged(updates, "size_bytes", existing.SizeBytes, m.SizeBytes)
- setIfChanged(updates, "duration_sec", existing.DurationSec, m.DurationSec)
- setIfChanged(updates, "width", existing.Width, m.Width)
- setIfChanged(updates, "height", existing.Height, m.Height)
- setIfChanged(updates, "video_codec", existing.VideoCodec, m.VideoCodec)
- setIfChanged(updates, "audio_codec", existing.AudioCodec, m.AudioCodec)
- setIfChanged(updates, "container", existing.Container, m.Container)
- if existing.DeletedAt.Valid {
- updates["deleted_at"] = nil
- }
- // 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。
- if m.FileID != "" && m.FileID != existing.FileID {
- updates["file_id"] = m.FileID
- }
- if m.Title != "" {
- // scanner 给出的标题只是从路径推导,刮削后 title 已被替换为
- // 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题,
- // 避免覆盖刮削结果。
- if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" {
- setIfChanged(updates, "title", existing.Title, m.Title)
- if m.Year > 0 {
- setIfChanged(updates, "year", existing.Year, m.Year)
- }
- }
- }
- status := strings.TrimSpace(existing.ScrapeStatus)
- canRefreshExternalIDs := status == "pending" || status == "" || status == "no_match" ||
- m.ScrapeStatus == "matched" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://")
- if canRefreshExternalIDs {
- changedExternalID := false
- if m.TMDbID > 0 && existing.TMDbID != m.TMDbID {
- updates["tm_db_id"] = m.TMDbID
- changedExternalID = true
- }
- if m.BangumiID > 0 && existing.BangumiID != m.BangumiID {
- updates["bangumi_id"] = m.BangumiID
- changedExternalID = true
- }
- if m.DoubanID != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(m.DoubanID) {
- updates["douban_id"] = m.DoubanID
- changedExternalID = true
- }
- if m.TheTVDBID != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(m.TheTVDBID) {
- updates["thetvdb_id"] = m.TheTVDBID
- changedExternalID = true
- }
- if m.Year > 0 && existing.Year <= 0 {
- updates["year"] = m.Year
- }
- if changedExternalID && (status == "no_match" || status == "matched") && m.ScrapeStatus != "matched" {
- updates["scrape_status"] = "pending"
- }
- }
- if m.ScrapeStatus == "matched" {
- setIfChanged(updates, "scrape_status", existing.ScrapeStatus, m.ScrapeStatus)
- if m.OriginalName != "" {
- setIfChanged(updates, "original_name", existing.OriginalName, m.OriginalName)
- }
- if m.PosterURL != "" {
- setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
- }
- if m.BackdropURL != "" {
- setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
- }
- if m.Overview != "" {
- setIfChanged(updates, "overview", existing.Overview, m.Overview)
- }
- if m.Rating > 0 {
- setIfChanged(updates, "rating", existing.Rating, m.Rating)
- }
- if m.Year > 0 {
- setIfChanged(updates, "year", existing.Year, m.Year)
- }
- if m.TMDbID > 0 {
- setIfChanged(updates, "tm_db_id", existing.TMDbID, m.TMDbID)
- }
- if m.BangumiID > 0 {
- setIfChanged(updates, "bangumi_id", existing.BangumiID, m.BangumiID)
- }
- if m.DoubanID != "" {
- setIfChanged(updates, "douban_id", existing.DoubanID, m.DoubanID)
- }
- if m.TheTVDBID != "" {
- setIfChanged(updates, "thetvdb_id", existing.TheTVDBID, m.TheTVDBID)
- }
- if m.Languages != "" {
- setIfChanged(updates, "languages", existing.Languages, m.Languages)
- }
- if m.Countries != "" {
- setIfChanged(updates, "countries", existing.Countries, m.Countries)
- }
- if m.Genres != "" {
- setIfChanged(updates, "genres", existing.Genres, m.Genres)
- }
- if m.NSFW && !existing.NSFW {
- updates["nsfw"] = true
- }
- }
- if m.PosterURL != "" {
- setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
- }
- if m.BackdropURL != "" {
- setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
- }
- // 云盘媒体:同一 cloud:// 文件可能先被父目录库扫描入库,之后用户按二级
- // 分类重新挂载/扫描到更精确的分类库。此时让 library_id 迁移到当前扫描库,
- // 否则媒体被钉死在旧库、新分类库里看不到(表现为"媒体部分消失")。
- // 本地媒体物理位置固定:仅在原 library_id 为空时回填,不迁移。
- if isCloudMediaPath := strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://"); m.LibraryID != "" && m.LibraryID != existing.LibraryID {
- if isCloudMediaPath || existing.LibraryID == "" {
- updates["library_id"] = m.LibraryID
- }
- }
- if (m.SeasonNum > 0 || m.EpisodeNum > 0) && existing.SeasonNum != m.SeasonNum {
- updates["season_num"] = m.SeasonNum
- }
- if m.EpisodeNum > 0 && existing.EpisodeNum != m.EpisodeNum {
- updates["episode_num"] = m.EpisodeNum
- }
- if m.STRMURL != "" {
- setIfChanged(updates, "strm_url", existing.STRMURL, m.STRMURL)
- }
-
- if len(updates) == 0 {
- *m = existing
- return nil
- }
- if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}).
- Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
- return err
- }
- // 回写 ID / 不可变字段,让 caller 拿到完整的现有行。
- *m = existing
- if fresh, err := r.FindByID(ctx, existing.ID); err == nil && fresh != nil {
- r.indexMediaBestEffort(ctx, *fresh)
- }
- return nil
-}
-
-func (r *MediaRepository) indexMediaBestEffort(ctx context.Context, media model.Media) {
- backend, ok := r.searchBackend.(MediaSearchSyncBackend)
- if !ok {
- return
- }
- _ = backend.IndexMedia(ctx, []model.Media{media})
-}
-
-func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
- if current != next {
- updates[key] = next
- }
-}
-
-// FindByID returns the media row or (nil, nil).
-func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) {
- var m model.Media
- err := r.db.WithContext(ctx).Where("id = ?", id).First(&m).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &m, nil
-}
-
-// ListByLibrary returns paginated media items for a library.
-func (r *MediaRepository) ListByLibrary(ctx context.Context, libraryID string, offset, limit int) ([]model.Media, int64, error) {
- return r.ListByLibraryFiltered(ctx, libraryID, offset, limit, MediaQueryFilter{IncludeNSFW: true})
-}
-
-func (r *MediaRepository) ListByLibraryFiltered(ctx context.Context, libraryID string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
- return r.ListByLibrariesFiltered(ctx, []string{libraryID}, offset, limit, filter)
-}
-
-func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
- var items []model.Media
- var total int64
- if len(libraryIDs) == 0 {
- return items, 0, nil
- }
- q := r.db.WithContext(ctx).Model(&model.Media{})
- if len(libraryIDs) == 1 {
- q = q.Where("library_id = ?", libraryIDs[0])
- } else {
- q = q.Where("library_id IN ?", libraryIDs)
- }
- q = applyMediaQueryFilter(q, filter)
- if err := q.Count(&total).Error; err != nil {
- return nil, 0, err
- }
- // 多级排序消除"随机"观感:
- // 1. year desc — 上映年份新→旧(用户期望的上映时间维度)
- // 2. updated_at desc — 同年按最近更新(刮削/补集会刷新)
- // 3. created_at desc — 再按入库时间
- // 4. id desc — 稳定 tie-breaker:云盘批量扫描同批 created_at 相同时,
- // 没有它 DB 返回顺序不确定,正是"随机排序"的根因。
- err := q.Order("year DESC, updated_at DESC, created_at DESC, id DESC").
- Offset(offset).Limit(limit).Find(&items).Error
- return items, total, err
-}
-
-// Search runs a LIKE search against the title field. Empty query returns the
-// most recently added items.
-func (r *MediaRepository) Search(ctx context.Context, query string, limit int) ([]model.Media, error) {
- return r.SearchFiltered(ctx, query, limit, MediaQueryFilter{IncludeNSFW: true})
-}
-
-func (r *MediaRepository) SearchFiltered(ctx context.Context, query string, limit int, filter MediaQueryFilter) ([]model.Media, error) {
- items, _, err := r.SearchFilteredPage(ctx, query, 0, limit, filter)
- return items, err
-}
-
-func (r *MediaRepository) SearchFilteredPage(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
- query = strings.TrimSpace(query)
- if limit <= 0 {
- limit = 50
- }
- if query != "" && r.searchBackend != nil {
- if items, total, ok := r.searchFilteredBackend(ctx, query, offset, limit, filter); ok {
- return items, total, nil
- }
- }
- if query != "" {
- if items, total, ok := r.searchFilteredFTS(ctx, query, offset, limit, filter); ok {
- if total > 0 {
- return items, total, nil
- }
- }
- }
- return r.searchFilteredLIKE(ctx, query, offset, limit, filter)
-}
-
-func (r *MediaRepository) searchFilteredBackend(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
- ids, total, err := r.searchBackend.SearchMediaIDs(ctx, query, offset, limit, filter)
- if err != nil {
- return nil, 0, false
- }
- if len(ids) == 0 {
- return []model.Media{}, total, true
- }
- var rows []model.Media
- q := r.db.WithContext(ctx).Model(&model.Media{}).Where("id IN ?", ids)
- q = applyMediaQueryFilter(q, filter)
- if err := q.Find(&rows).Error; err != nil {
- return nil, 0, false
- }
- byID := make(map[string]model.Media, len(rows))
- for _, row := range rows {
- byID[row.ID] = row
- }
- items := make([]model.Media, 0, len(ids))
- for _, id := range ids {
- if row, ok := byID[id]; ok {
- items = append(items, row)
- }
- }
- if len(items) == 0 && total > 0 {
- return nil, 0, false
- }
- return items, total, true
-}
-
-func (r *MediaRepository) searchFilteredFTS(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
- if !r.searchIndexEnabled(ctx) {
- return nil, 0, false
- }
- ftsQuery := mediaFTSQuery(query)
- if ftsQuery == "" {
- return nil, 0, false
- }
- var total int64
- var items []model.Media
- q := r.db.WithContext(ctx).
- Table("media").
- Joins("JOIN media_search_fts ON media_search_fts.rowid = media.rowid").
- Where("media.deleted_at IS NULL").
- Where("media_search_fts MATCH ?", ftsQuery)
- q = applyQualifiedMediaQueryFilter(q, filter)
- if err := q.Count(&total).Error; err != nil {
- return nil, 0, false
- }
- if total == 0 {
- return items, 0, true
- }
- err := q.Select("media.*").Order("bm25(media_search_fts), media.created_at DESC").Offset(offset).Limit(limit).Find(&items).Error
- if err != nil {
- return nil, 0, false
- }
- return items, total, true
-}
-
-func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
- var items []model.Media
- var total int64
- q := r.db.WithContext(ctx).Model(&model.Media{})
- q = applyMediaQueryFilter(q, filter)
- terms := mediaSearchTerms(query)
- for _, term := range terms {
- like := "%" + escapeLike(term) + "%"
- q = q.Where(
- "(title LIKE ? ESCAPE '\\' OR original_name LIKE ? ESCAPE '\\' OR path LIKE ? ESCAPE '\\' OR genres LIKE ? ESCAPE '\\')",
- like, like, like, like,
- )
- }
- if err := q.Count(&total).Error; err != nil {
- return nil, 0, err
- }
- if query != "" {
- prefix := escapeLike(query) + "%"
- exact := query
- q = q.Order(gorm.Expr(
- "CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title LIKE ? ESCAPE '\\' THEN 2 WHEN original_name LIKE ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc",
- exact, exact, prefix, prefix,
- ))
- } else {
- q = q.Order("created_at desc")
- }
- err := q.Offset(offset).Limit(limit).Find(&items).Error
- return items, total, err
-}
-
-func applyQualifiedMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
- if !filter.IncludeNSFW {
- q = q.Where("media.nsfw = ?", false)
- }
- if len(filter.HiddenLibraryIDs) > 0 {
- q = q.Where("media.library_id NOT IN ?", filter.HiddenLibraryIDs)
- }
- if len(filter.AllowedLibraryIDs) > 0 {
- q = q.Where("media.library_id IN ?", filter.AllowedLibraryIDs)
- }
- return q
-}
-
-func mediaFTSQuery(query string) string {
- terms := mediaSearchTerms(query)
- if len(terms) == 0 {
- return ""
- }
- quoted := make([]string, 0, len(terms))
- for _, term := range terms {
- term = strings.ReplaceAll(term, `"`, `""`)
- if term != "" {
- quoted = append(quoted, `"`+term+`"`)
- }
- }
- return strings.Join(quoted, " AND ")
-}
-
-func mediaSearchTerms(query string) []string {
- query = strings.TrimSpace(query)
- if query == "" {
- return nil
- }
- fields := strings.FieldsFunc(query, func(r rune) bool {
- return unicode.IsSpace(r) || unicode.IsPunct(r) || unicode.IsSymbol(r)
- })
- out := make([]string, 0, len(fields))
- seen := map[string]struct{}{}
- for _, field := range fields {
- field = strings.TrimSpace(field)
- if field == "" {
- continue
- }
- lower := strings.ToLower(field)
- if _, ok := seen[lower]; ok {
- continue
- }
- seen[lower] = struct{}{}
- out = append(out, field)
- }
- return out
-}
-
-func escapeLike(value string) string {
- value = strings.ReplaceAll(value, `\`, `\\`)
- value = strings.ReplaceAll(value, `%`, `\%`)
- value = strings.ReplaceAll(value, `_`, `\_`)
- return value
-}
-
-func (r *MediaRepository) BackfillSearchIndex(ctx context.Context, batchLimit int) (int64, error) {
- if backend, ok := r.searchBackend.(MediaSearchSyncBackend); ok {
- return r.backfillExternalSearchIndex(ctx, backend, batchLimit)
- }
- if batchLimit <= 0 {
- batchLimit = 1000
- }
- if !r.searchIndexEnabled(ctx) {
- return 0, nil
- }
- // 关键性能点:FTS5 普通列(含 UNINDEXED)不支持索引查找,按
- // media_id 做 NOT EXISTS 是对 FTS 表的整表扫描,再叠加 ORDER BY
- // 后每个批次都要对全部 media 行探测一遍——大库一次启动回填等于
- // 上百亿次行访问,曾把 CPU 钉满数小时。v2 布局下 FTS 行 rowid 与
- // media.rowid 对齐,NOT EXISTS 走 rowid 点查,且无需排序。
- res := r.db.WithContext(ctx).Exec(`
-INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
-SELECT m.rowid, m.id, COALESCE(m.title, ''), COALESCE(m.original_name, ''), COALESCE(m.path, ''), COALESCE(m.genres, '')
-FROM media AS m
-WHERE m.deleted_at IS NULL
- AND NOT EXISTS (
- SELECT 1 FROM media_search_fts AS f WHERE f.rowid = m.rowid
- )
-LIMIT ?
-`, batchLimit)
- return res.RowsAffected, res.Error
-}
-
-func (r *MediaRepository) backfillExternalSearchIndex(ctx context.Context, backend MediaSearchSyncBackend, batchLimit int) (int64, error) {
- if batchLimit <= 0 {
- batchLimit = 1000
- }
- if err := backend.EnsureIndex(ctx); err != nil {
- return 0, err
- }
- var lastID string
- for {
- var rows []model.Media
- q := r.db.WithContext(ctx).
- Model(&model.Media{}).
- Where("deleted_at IS NULL")
- if lastID != "" {
- q = q.Where("id > ?", lastID)
- }
- if err := q.Order("id ASC").Limit(batchLimit).Find(&rows).Error; err != nil {
- return 0, err
- }
- if len(rows) == 0 {
- return 0, nil
- }
- if err := backend.IndexMedia(ctx, rows); err != nil {
- return 0, err
- }
- lastID = rows[len(rows)-1].ID
- if len(rows) < batchLimit {
- return 0, nil
- }
- }
-}
-
-func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool {
- if r == nil || r.db == nil {
- return false
- }
- if r.db.Dialector == nil || r.db.Dialector.Name() != "sqlite" {
- return false
- }
- r.searchIndexOnce.Do(func() {
- var count int64
- err := r.db.WithContext(ctx).
- Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`).
- Scan(&count).Error
- r.searchIndexAvailable = err == nil && count > 0
- })
- return r.searchIndexAvailable
-}
-
-// DeleteByLibrary purges all media tied to a library.
-func (r *MediaRepository) DeleteByLibrary(ctx context.Context, libraryID string) error {
- // FTS 行由 media 表上的触发器同步清理(软删/硬删都覆盖)。
- return r.db.WithContext(ctx).Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
-}
-
-// PurgeByLibrary permanently removes media tied to a library. Used for virtual
-// cloud mounts where "remove mount" must not populate the recycle bin.
-func (r *MediaRepository) PurgeByLibrary(ctx context.Context, libraryID string) error {
- return r.db.WithContext(ctx).Unscoped().Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
-}
-
-// ─── Series ──────────────────────────────────────────────────────────────────
-
-// SeriesRepository persists model.Series records.
-type SeriesRepository struct{ db *gorm.DB }
-
-// FindByID returns the series or (nil, nil).
-func (r *SeriesRepository) FindByID(ctx context.Context, id string) (*model.Series, error) {
- var s model.Series
- err := r.db.WithContext(ctx).Where("id = ?", id).First(&s).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &s, nil
-}
-
-// List returns all series (ordered by title).
-func (r *SeriesRepository) List(ctx context.Context) ([]model.Series, error) {
- var s []model.Series
- err := r.db.WithContext(ctx).Order("title asc").Find(&s).Error
- return s, err
-}
-
-// ─── Playback History ────────────────────────────────────────────────────────
-
-// HistoryRepository persists model.PlaybackHistory entries. The application
-// upserts on (UserID, MediaID) so resume always reads the latest position.
-type HistoryRepository struct{ db *gorm.DB }
-
-// Upsert atomically inserts/updates the resume position.
-func (r *HistoryRepository) Upsert(ctx context.Context, h *model.PlaybackHistory) error {
- var existing model.PlaybackHistory
- err := r.db.WithContext(ctx).
- Where("user_id = ? AND media_id = ?", h.UserID, h.MediaID).
- First(&existing).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return r.db.WithContext(ctx).Create(h).Error
- }
- if err != nil {
- return err
- }
- existing.PositionMs = h.PositionMs
- existing.DurationMs = h.DurationMs
- existing.WatchedAt = h.WatchedAt
- existing.Completed = h.Completed
- return r.db.WithContext(ctx).Save(&existing).Error
-}
-
-// ListByUser returns the most recent history rows for the user.
-func (r *HistoryRepository) ListByUser(ctx context.Context, userID string, limit int) ([]model.PlaybackHistory, error) {
- var rows []model.PlaybackHistory
- err := r.db.WithContext(ctx).Where("user_id = ?", userID).
- Order("watched_at desc").Limit(limit).Find(&rows).Error
- return rows, err
-}
-
-// ─── Favorite ───────────────────────────────────────────────────────────────
-
-// FavoriteRepository persists model.Favorite records.
-type FavoriteRepository struct{ db *gorm.DB }
-
-// Toggle flips the favourite flag for (user, media). Returns the new state.
-func (r *FavoriteRepository) Toggle(ctx context.Context, userID, mediaID string) (bool, error) {
- var f model.Favorite
- err := r.db.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- fav := model.Favorite{UserID: userID, MediaID: mediaID}
- return true, r.db.WithContext(ctx).Create(&fav).Error
- }
- if err != nil {
- return false, err
- }
- return false, r.db.WithContext(ctx).Delete(&f).Error
-}
-
-// ListByUser returns all favourite media IDs for a user.
-func (r *FavoriteRepository) ListByUser(ctx context.Context, userID string) ([]model.Favorite, error) {
- var rows []model.Favorite
- err := r.db.WithContext(ctx).Where("user_id = ?", userID).Find(&rows).Error
- return rows, err
-}
-
-// ─── Playlist ────────────────────────────────────────────────────────────────
-
-// PlaylistRepository persists model.Playlist + PlaylistItem.
-type PlaylistRepository struct{ db *gorm.DB }
-
-// Create inserts a new playlist.
-func (r *PlaylistRepository) Create(ctx context.Context, p *model.Playlist) error {
- return r.db.WithContext(ctx).Create(p).Error
-}
-
-// ListByUser returns playlists owned by a user.
-func (r *PlaylistRepository) ListByUser(ctx context.Context, userID string) ([]model.Playlist, error) {
- var rows []model.Playlist
- err := r.db.WithContext(ctx).Where("user_id = ?", userID).
- Order("created_at desc").Find(&rows).Error
- return rows, err
-}
-
-// ─── Download ───────────────────────────────────────────────────────────────
-
-// DownloadRepository persists model.DownloadTask records.
-type DownloadRepository struct{ db *gorm.DB }
-
-// Create inserts a new download task.
-func (r *DownloadRepository) Create(ctx context.Context, t *model.DownloadTask) error {
- return r.db.WithContext(ctx).Create(t).Error
-}
-
-// List returns all download tasks (admin view).
-func (r *DownloadRepository) List(ctx context.Context) ([]model.DownloadTask, error) {
- var rows []model.DownloadTask
- err := r.db.WithContext(ctx).Order("created_at desc").Find(&rows).Error
- return rows, err
-}
-
-// ─── Subscription ───────────────────────────────────────────────────────────
-
-// SubscriptionRepository persists model.Subscription records.
-type SubscriptionRepository struct{ db *gorm.DB }
-
-// Create inserts a new subscription rule.
-func (r *SubscriptionRepository) Create(ctx context.Context, s *model.Subscription) error {
- return r.db.WithContext(ctx).Select("*").Omit("DeletedAt").Create(s).Error
-}
-
-// List returns active subscription rules. Archived rows live in history and are
-// intentionally excluded from scheduler polling and the active management list.
-func (r *SubscriptionRepository) List(ctx context.Context) ([]model.Subscription, error) {
- var rows []model.Subscription
- err := r.db.WithContext(ctx).Where("archived_at IS NULL").Order("created_at desc").Find(&rows).Error
- return rows, err
-}
-
-// History returns archived subscription rules.
-func (r *SubscriptionRepository) History(ctx context.Context) ([]model.Subscription, error) {
- var rows []model.Subscription
- err := r.db.WithContext(ctx).Where("archived_at IS NOT NULL").Order("archived_at desc, updated_at desc").Find(&rows).Error
- return rows, err
-}
-
-// Archive moves a completed subscription out of the active list without
-// deleting its rule details, so users can audit completed subscriptions later.
-func (r *SubscriptionRepository) Archive(ctx context.Context, id, reason string, archivedAt time.Time) error {
- return r.db.WithContext(ctx).Model(&model.Subscription{}).
- Where("id = ? AND archived_at IS NULL", id).
- Updates(map[string]any{
- "enabled": false,
- "archived_at": &archivedAt,
- "archive_reason": reason,
- }).Error
-}
-
-// ─── Setting ─────────────────────────────────────────────────────────────────
-
-// SettingRepository persists key/value preferences.
-type SettingRepository struct{ db *gorm.DB }
-
-// Get returns the value or empty string when absent.
-func (r *SettingRepository) Get(ctx context.Context, key string) (string, error) {
- var s model.Setting
- err := r.db.WithContext(ctx).Where("key = ?", key).First(&s).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return "", nil
- }
- return s.Value, err
-}
-
-// Set upserts a setting value.
-func (r *SettingRepository) Set(ctx context.Context, key, value string) error {
- s := model.Setting{Key: key, Value: value, UpdatedAt: time.Now()}
- return r.db.WithContext(ctx).Save(&s).Error
-}
-
-// Delete removes a setting key.
-func (r *SettingRepository) Delete(ctx context.Context, key string) error {
- return r.db.WithContext(ctx).Where("key = ?", key).Delete(&model.Setting{}).Error
-}
-
-// All returns every key/value pair (used by the admin UI).
-func (r *SettingRepository) All(ctx context.Context) ([]model.Setting, error) {
- var rows []model.Setting
- err := r.db.WithContext(ctx).Find(&rows).Error
- return rows, err
-}
-
-// ─── Access Log ──────────────────────────────────────────────────────────────
-
-// AccessLogRepository persists model.AccessLog records.
-type AccessLogRepository struct{ db *gorm.DB }
-
-// Create inserts one structured audit-trail entry.
-func (r *AccessLogRepository) Create(ctx context.Context, l *model.AccessLog) error {
- return r.db.WithContext(ctx).Create(l).Error
-}
-
-// Recent returns the latest access-log entries (admin Activity panel).
-func (r *AccessLogRepository) Recent(ctx context.Context, limit int) ([]model.AccessLog, error) {
- var rows []model.AccessLog
- err := r.db.WithContext(ctx).Order("created_at desc").Limit(limit).Find(&rows).Error
- return rows, err
-}
-
-// ─── Permission ──────────────────────────────────────────────────────────────
-
-// PermissionRepository persists model.UserPermission records.
-type PermissionRepository struct{ db *gorm.DB }
-
-// Create inserts a new permission record.
-func (r *PermissionRepository) Create(ctx context.Context, p *model.UserPermission) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Create(p).Error
- })
-}
-
-// FindByUserID returns the permission record for a user, or (nil, nil) when absent.
-func (r *PermissionRepository) FindByUserID(ctx context.Context, userID string) (*model.UserPermission, error) {
- var p model.UserPermission
- err := withSQLiteBusyRetry(ctx, func() error {
- p = model.UserPermission{}
- return r.db.WithContext(ctx).Where("user_id = ?", userID).First(&p).Error
- })
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &p, nil
-}
-
-// Update updates permission fields for a user.
-func (r *PermissionRepository) Update(ctx context.Context, userID string, updates map[string]bool) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Model(&model.UserPermission{}).
- Where("user_id = ?", userID).Updates(updates).Error
- })
-}
-
-// Upsert creates or updates a permission record.
-func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
- Assign(*p).FirstOrCreate(p).Error
- })
-}
-
-// Delete removes a permission record.
-func (r *PermissionRepository) Delete(ctx context.Context, userID string) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
- })
-}
-
-// ─── Refresh Token ───────────────────────────────────────────────────────────
-
-// RefreshTokenRepository persists model.RefreshToken records.
-type RefreshTokenRepository struct{ db *gorm.DB }
-
-// Create inserts a new refresh token record.
-func (r *RefreshTokenRepository) Create(ctx context.Context, t *model.RefreshToken) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Create(t).Error
- })
-}
-
-// FindByHash returns the refresh token matching the hash, or (nil, nil).
-func (r *RefreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) {
- var t model.RefreshToken
- err := withSQLiteBusyRetry(ctx, func() error {
- t = model.RefreshToken{}
- return r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&t).Error
- })
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &t, nil
-}
-
-// RevokeByUserID revokes all refresh tokens for a user.
-func (r *RefreshTokenRepository) RevokeByUserID(ctx context.Context, userID string) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
- Where("user_id = ?", userID).Update("revoked", true).Error
- })
-}
-
-// RevokeOldestActiveByUserID keeps at most limit active refresh tokens for a
-// user by revoking the oldest non-expired, non-revoked tokens.
-func (r *RefreshTokenRepository) RevokeOldestActiveByUserID(ctx context.Context, userID string, limit int) error {
- if limit < 1 {
- limit = 1
- }
- return withSQLiteBusyRetry(ctx, func() error {
- var tokens []model.RefreshToken
- if err := r.db.WithContext(ctx).
- Where("user_id = ? AND revoked = ? AND expires_at > ?", userID, false, time.Now()).
- Order("created_at desc, id desc").
- Find(&tokens).Error; err != nil {
- return err
- }
- if len(tokens) <= limit {
- return nil
- }
- ids := make([]string, 0, len(tokens)-limit)
- for _, token := range tokens[limit:] {
- ids = append(ids, token.ID)
- }
- return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
- Where("id IN ?", ids).Update("revoked", true).Error
- })
-}
-
-// DeleteExpired removes all expired refresh tokens.
-func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error
- })
-}
-
-// Revoke revokes a specific refresh token.
-func (r *RefreshTokenRepository) Revoke(ctx context.Context, hash string) error {
- return withSQLiteBusyRetry(ctx, func() error {
- return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
- Where("token_hash = ?", hash).Update("revoked", true).Error
- })
-}
-
-// HashToken returns the SHA256 hash of a token.
-func HashToken(token string) string {
- h := sha256.Sum256([]byte(token))
- return hex.EncodeToString(h[:])
-}
-
-// ─── API Config ──────────────────────────────────────────────────────────────
-
-// ApiConfigRepository persists model.ApiConfig records.
-type ApiConfigRepository struct{ db *gorm.DB }
-
-// Create inserts a new API config record.
-func (r *ApiConfigRepository) Create(ctx context.Context, c *model.ApiConfig) error {
- return r.db.WithContext(ctx).Create(c).Error
-}
-
-// FindByProvider returns the API config for a provider, or (nil, nil).
-func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
- var c model.ApiConfig
- err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- if err != nil {
- return nil, err
- }
- return &c, nil
-}
-
-// List returns all API configs.
-func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) {
- var rows []model.ApiConfig
- err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error
- return rows, err
-}
-
-// Upsert creates or updates an API config.
-func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.ApiConfig) error {
- return r.db.WithContext(ctx).Where("provider = ?", c.Provider).
- Assign(model.ApiConfig{
- Base: model.Base{UpdatedAt: time.Now()},
- APIKey: c.APIKey,
- BaseURL: c.BaseURL,
- Extra: c.Extra,
- Enabled: c.Enabled,
- }).FirstOrCreate(c).Error
-}
-
-// Update updates an API config.
-func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) error {
- return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
- Where("provider = ?", c.Provider).Updates(map[string]any{
- "api_key": c.APIKey,
- "base_url": c.BaseURL,
- "extra": c.Extra,
- "enabled": c.Enabled,
- "updated_at": time.Now(),
- }).Error
-}
-
-// Delete removes an API config.
-func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error {
- return r.db.WithContext(ctx).Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
-}
-
-// UpdateTestResult 更新测试结果。
-func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error {
- now := time.Now()
- return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
- Where("provider = ?", provider).Updates(map[string]any{
- "test_result": result,
- "last_tested_at": &now,
- }).Error
-}
diff --git a/internal/repository/series_repository.go b/internal/repository/series_repository.go
new file mode 100644
index 0000000..06da499
--- /dev/null
+++ b/internal/repository/series_repository.go
@@ -0,0 +1,33 @@
+package repository
+
+import (
+ "context"
+ "errors"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// SeriesRepository persists model.Series records.
+type SeriesRepository struct{ db *gorm.DB }
+
+// FindByID returns the series or (nil, nil).
+func (r *SeriesRepository) FindByID(ctx context.Context, id string) (*model.Series, error) {
+ var s model.Series
+ err := r.db.WithContext(ctx).Where("id = ?", id).First(&s).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &s, nil
+}
+
+// List returns all series (ordered by title).
+func (r *SeriesRepository) List(ctx context.Context) ([]model.Series, error) {
+ var s []model.Series
+ err := r.db.WithContext(ctx).Order("title asc").Find(&s).Error
+ return s, err
+}
diff --git a/internal/repository/setting_repository.go b/internal/repository/setting_repository.go
new file mode 100644
index 0000000..80ade8c
--- /dev/null
+++ b/internal/repository/setting_repository.go
@@ -0,0 +1,42 @@
+package repository
+
+import (
+ "context"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// SettingRepository persists key/value preferences.
+type SettingRepository struct{ db *gorm.DB }
+
+// Get returns the value or empty string when absent.
+func (r *SettingRepository) Get(ctx context.Context, key string) (string, error) {
+ var value string
+ err := r.db.WithContext(ctx).
+ Model(&model.Setting{}).
+ Select("value").
+ Where("key = ?", key).
+ Scan(&value).Error
+ return value, err
+}
+
+// Set upserts a setting value.
+func (r *SettingRepository) Set(ctx context.Context, key, value string) error {
+ s := model.Setting{Key: key, Value: value, UpdatedAt: time.Now()}
+ return r.db.WithContext(ctx).Save(&s).Error
+}
+
+// Delete removes a setting key.
+func (r *SettingRepository) Delete(ctx context.Context, key string) error {
+ return r.db.WithContext(ctx).Where("key = ?", key).Delete(&model.Setting{}).Error
+}
+
+// All returns every key/value pair (used by the admin UI).
+func (r *SettingRepository) All(ctx context.Context) ([]model.Setting, error) {
+ var rows []model.Setting
+ err := r.db.WithContext(ctx).Find(&rows).Error
+ return rows, err
+}
diff --git a/internal/repository/subscription_repository.go b/internal/repository/subscription_repository.go
new file mode 100644
index 0000000..0276cd7
--- /dev/null
+++ b/internal/repository/subscription_repository.go
@@ -0,0 +1,45 @@
+package repository
+
+import (
+ "context"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// SubscriptionRepository persists model.Subscription records.
+type SubscriptionRepository struct{ db *gorm.DB }
+
+// Create inserts a new subscription rule.
+func (r *SubscriptionRepository) Create(ctx context.Context, s *model.Subscription) error {
+ return r.db.WithContext(ctx).Select("*").Omit("DeletedAt").Create(s).Error
+}
+
+// List returns active subscription rules. Archived rows live in history and are
+// intentionally excluded from scheduler polling and the active management list.
+func (r *SubscriptionRepository) List(ctx context.Context) ([]model.Subscription, error) {
+ var rows []model.Subscription
+ err := r.db.WithContext(ctx).Where("archived_at IS NULL").Order("created_at desc").Find(&rows).Error
+ return rows, err
+}
+
+// History returns archived subscription rules.
+func (r *SubscriptionRepository) History(ctx context.Context) ([]model.Subscription, error) {
+ var rows []model.Subscription
+ err := r.db.WithContext(ctx).Where("archived_at IS NOT NULL").Order("archived_at desc, updated_at desc").Find(&rows).Error
+ return rows, err
+}
+
+// Archive moves a completed subscription out of the active list without
+// deleting its rule details, so users can audit completed subscriptions later.
+func (r *SubscriptionRepository) Archive(ctx context.Context, id, reason string, archivedAt time.Time) error {
+ return r.db.WithContext(ctx).Model(&model.Subscription{}).
+ Where("id = ? AND archived_at IS NULL", id).
+ Updates(map[string]any{
+ "enabled": false,
+ "archived_at": &archivedAt,
+ "archive_reason": reason,
+ }).Error
+}
diff --git a/internal/repository/user_repository.go b/internal/repository/user_repository.go
new file mode 100644
index 0000000..0e51362
--- /dev/null
+++ b/internal/repository/user_repository.go
@@ -0,0 +1,161 @@
+package repository
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/hex"
+ "errors"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// UserRepository persists model.User records.
+type UserRepository struct{ db *gorm.DB }
+
+// Create inserts a new user. Caller must pre-hash the password.
+func (r *UserRepository) Create(ctx context.Context, u *model.User) error {
+ return r.db.WithContext(ctx).Create(u).Error
+}
+
+// ReleaseDeletedUsername renames soft-deleted rows that still hold a unique
+// username so the same account name can be created again.
+func (r *UserRepository) ReleaseDeletedUsername(ctx context.Context, username string) error {
+ if username == "" {
+ return nil
+ }
+ released := username + "__deleted__" + time.Now().Format("20060102150405.000000000")
+ if len(released) > 64 {
+ sum := sha256.Sum256([]byte(released))
+ released = username
+ if len(released) > 43 {
+ released = released[:43]
+ }
+ released += "__deleted__" + hex.EncodeToString(sum[:])[:10]
+ }
+ return r.db.WithContext(ctx).Unscoped().
+ Model(&model.User{}).
+ Where("username = ? AND deleted_at IS NOT NULL", username).
+ Update("username", released).Error
+}
+
+// FindByUsername returns the user matching username, or (nil, nil) when absent.
+func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) {
+ var u model.User
+ err := withSQLiteBusyRetry(ctx, func() error {
+ u = model.User{}
+ err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) && username != "" {
+ err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error
+ }
+ return err
+ })
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &u, nil
+}
+
+// FindByID returns the user with the matching primary key, or (nil, nil).
+func (r *UserRepository) FindByID(ctx context.Context, id string) (*model.User, error) {
+ var u model.User
+ err := withSQLiteBusyRetry(ctx, func() error {
+ u = model.User{}
+ return r.db.WithContext(ctx).Where("id = ?", id).First(&u).Error
+ })
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &u, nil
+}
+
+// Count returns the total number of non-deleted users.
+func (r *UserRepository) Count(ctx context.Context) (int64, error) {
+ var n int64
+ err := r.db.WithContext(ctx).Model(&model.User{}).Count(&n).Error
+ return n, err
+}
+
+// CountAdmins returns the number of users that hold the admin role.
+func (r *UserRepository) CountAdmins(ctx context.Context) (int64, error) {
+ var n int64
+ err := r.db.WithContext(ctx).Model(&model.User{}).
+ Where("role = ?", "admin").Count(&n).Error
+ return n, err
+}
+
+// FirstAdmin returns the earliest admin user. This row represents the protected
+// built-in/default administrator even if its username is later changed.
+func (r *UserRepository) FirstAdmin(ctx context.Context) (*model.User, error) {
+ var u model.User
+ err := r.db.WithContext(ctx).Where("role = ?", "admin").Order("created_at asc").First(&u).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &u, nil
+}
+
+// List returns all users ordered by creation time desc.
+func (r *UserRepository) List(ctx context.Context) ([]model.User, error) {
+ var users []model.User
+ err := r.db.WithContext(ctx).Order("created_at desc").Find(&users).Error
+ return users, err
+}
+
+// UpdateFields applies a narrow set of user field updates.
+func (r *UserRepository) UpdateFields(ctx context.Context, id string, updates map[string]any) error {
+ return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error
+}
+
+// UpdatePassword sets a new password hash and clears ForcePasswordReset.
+func (r *UserRepository) UpdatePassword(ctx context.Context, id, hash string) error {
+ return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
+ Updates(map[string]any{"password_hash": hash, "force_password_reset": false}).Error
+}
+
+// TouchLogin updates the last login timestamp.
+func (r *UserRepository) TouchLogin(ctx context.Context, id string) error {
+ now := time.Now()
+ return withSQLiteBusyRetry(ctx, func() error {
+ return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
+ Update("last_login_at", &now).Error
+ })
+}
+
+// Delete removes a user (soft-delete via gorm.DeletedAt), releases the unique
+// username, and drops Telegram bindings so future re-created users bind cleanly.
+func (r *UserRepository) Delete(ctx context.Context, id string) error {
+ return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
+ var user model.User
+ if err := tx.Where("id = ?", id).First(&user).Error; err != nil {
+ return err
+ }
+ if err := tx.Unscoped().Where("user_id = ?", id).Delete(&model.TelegramBinding{}).Error; err != nil {
+ return err
+ }
+ released := user.Username + "__deleted__" + time.Now().Format("20060102150405.000000000")
+ if len(released) > 64 {
+ sum := sha256.Sum256([]byte(user.ID + user.Username))
+ base := user.Username
+ if len(base) > 43 {
+ base = base[:43]
+ }
+ released = base + "__deleted__" + hex.EncodeToString(sum[:])[:10]
+ }
+ if err := tx.Model(&model.User{}).Where("id = ?", id).Update("username", released).Error; err != nil {
+ return err
+ }
+ return tx.Delete(&model.User{}, "id = ?", id).Error
+ })
+}
diff --git a/internal/service/adult_scraper_test.go b/internal/service/adult_scraper_test.go
index 8e9ed41..d8c700e 100644
--- a/internal/service/adult_scraper_test.go
+++ b/internal/service/adult_scraper_test.go
@@ -6,9 +6,7 @@ import (
"net/http/httptest"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -88,13 +86,7 @@ func TestAdultProviderUsesConfiguredMultipleSources(t *testing.T) {
}))
defer good.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.APIConfig{})
apiConfig := NewAPIConfigService(zap.NewNop(), repository.New(db), NewCryptoService("", zap.NewNop()))
baseURL := bad.URL + "\n" + good.URL
if _, err := apiConfig.Update(context.Background(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
diff --git a/internal/service/ai_test.go b/internal/service/ai_test.go
index b025568..c9e27a0 100644
--- a/internal/service/ai_test.go
+++ b/internal/service/ai_test.go
@@ -4,9 +4,7 @@ import (
"context"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -14,13 +12,7 @@ import (
)
func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
crypto := NewCryptoService("test-secret", zap.NewNop())
apiConfig := NewAPIConfigService(zap.NewNop(), repo, crypto)
@@ -52,13 +44,7 @@ func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
}
func TestAIStatusHonorsDisabledDatabaseOpenAIConfig(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
apiConfig := NewAPIConfigService(zap.NewNop(), repo, NewCryptoService("test-secret", zap.NewNop()))
key := "sk-test"
diff --git a/internal/service/auth_user_limits_test.go b/internal/service/auth_user_limits_test.go
index 6ceb9ae..fe10110 100644
--- a/internal/service/auth_user_limits_test.go
+++ b/internal/service/auth_user_limits_test.go
@@ -9,11 +9,9 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"github.com/golang-jwt/jwt/v5"
"go.uber.org/zap"
"golang.org/x/crypto/bcrypt"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/database"
@@ -23,13 +21,7 @@ import (
func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *ProfileService, *PermissionService) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{})
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
diff --git a/internal/service/boot_cloud_health.go b/internal/service/boot_cloud_health.go
index 221db7c..8da6560 100644
--- a/internal/service/boot_cloud_health.go
+++ b/internal/service/boot_cloud_health.go
@@ -24,7 +24,7 @@ func (c *Container) BootCloudStorageHealthCheck(ctx context.Context) {
cloudConfigs := make([]StorageView, 0)
for _, cfg := range configs {
- if cfg.Enabled && (cfg.Type == "quark" || cfg.Type == "cloud115" || cfg.Type == "clouddrive2" || cfg.Type == "openlist") {
+ if cfg.Enabled && IsAdminCloudConfigurable(cfg.Type) {
cloudConfigs = append(cloudConfigs, cfg)
}
}
@@ -92,7 +92,7 @@ func cloudStorageMissingConfigReason(err error) string {
}
msg := strings.ToLower(strings.TrimSpace(err.Error()))
switch {
- case strings.Contains(msg, "missing cookie"):
+ 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"
diff --git a/internal/service/boot_cloud_health_test.go b/internal/service/boot_cloud_health_test.go
index eb7f806..09245cc 100644
--- a/internal/service/boot_cloud_health_test.go
+++ b/internal/service/boot_cloud_health_test.go
@@ -5,10 +5,8 @@ import (
"errors"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -20,8 +18,9 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
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("quark: token expired"), ""},
+ {errors.New("openlist: token expired"), ""},
}
for _, tc := range cases {
if got := cloudStorageMissingConfigReason(tc.err); got != tc.want {
@@ -31,19 +30,13 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
}
func TestWarnMissingCloudStorageConfigOncePersistsMarker(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ 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")
+ err := errors.New("115: missing cookie")
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
t.Fatal("missing config should be handled")
diff --git a/internal/service/bot_cleanup_test.go b/internal/service/bot_cleanup_test.go
new file mode 100644
index 0000000..71629b7
--- /dev/null
+++ b/internal/service/bot_cleanup_test.go
@@ -0,0 +1,246 @@
+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, "已清理 1") {
+ 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 · login_7d", "new_7d · 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)
+ }
+ }
+}
diff --git a/internal/service/bot_features_test.go b/internal/service/bot_features_test.go
index 1fc6b2a..8708f89 100644
--- a/internal/service/bot_features_test.go
+++ b/internal/service/bot_features_test.go
@@ -2,14 +2,11 @@ package service
import (
"context"
- "encoding/json"
"strings"
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -18,13 +15,7 @@ import (
func newBotTestService(t *testing.T) (*repository.Container, *TelegramBotService) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(model.AllModels()...); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, model.AllModels()...)
repos := repository.New(db)
cfg := &config.Config{}
cfg.Secrets.JWTSecret = "test-secret"
@@ -503,242 +494,6 @@ func TestBotAdminCommandsManageDevicePolicy(t *testing.T) {
}
}
-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, "已清理 1") {
- 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 · login_7d", "new_7d · 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)
- }
- }
-}
-
func TestBotRegistrationCommandUsesOpenRegQuota(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
@@ -828,137 +583,6 @@ func TestBotUserCommandsAndAdminGate(t *testing.T) {
}
}
-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")
- }
-}
-
func TestBotAdminCodeAndUserCommands(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
@@ -1032,122 +656,3 @@ func TestBotGroupMenuShowsAdminActionsOnlyForAdmins(t *testing.T) {
t.Fatalf("non-admin group menu must not expose management actions, got %#v", reply)
}
}
-
-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, "已解绑:2") || !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}
- recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime}
- for _, user := range []*model.User{admin, oldUser, 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: 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"}}
-
- reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30")
- if err != nil {
- t.Fatal(err)
- }
- if !strings.Contains(reply.Text, "已解绑:1") || !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")
- }
-
- reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates")
- if err != nil {
- t.Fatal(err)
- }
- if !strings.Contains(reply.Text, "已解绑:1") || !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)
- }
-}
\ No newline at end of file
diff --git a/internal/service/bot_registration_code_test.go b/internal/service/bot_registration_code_test.go
new file mode 100644
index 0000000..ab98e21
--- /dev/null
+++ b/internal/service/bot_registration_code_test.go
@@ -0,0 +1,141 @@
+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")
+ }
+}
diff --git a/internal/service/bot_unbind_test.go b/internal/service/bot_unbind_test.go
new file mode 100644
index 0000000..aa471de
--- /dev/null
+++ b/internal/service/bot_unbind_test.go
@@ -0,0 +1,129 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "testing"
+ "time"
+
+ "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, "已解绑:2") || !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}
+ recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime}
+ for _, user := range []*model.User{admin, oldUser, 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: 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"}}
+
+ reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !strings.Contains(reply.Text, "已解绑:1") || !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")
+ }
+
+ reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !strings.Contains(reply.Text, "已解绑:1") || !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)
+ }
+}
diff --git a/internal/service/cloud/cloud.go b/internal/service/cloud/cloud.go
index 34dee4c..985bef7 100644
--- a/internal/service/cloud/cloud.go
+++ b/internal/service/cloud/cloud.go
@@ -27,7 +27,6 @@ var timeNow = time.Now
// Provider types recognised by the registry.
const (
- TypeQuark = "quark" // 夸克网盘
Type115 = "cloud115" // 115 网盘
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
@@ -42,7 +41,7 @@ type FileEntry struct {
Name string `json:"name"`
IsDir bool `json:"is_dir"`
Size int64 `json:"size"`
- // PickCode is 115-specific; quark uses ID directly.
+ // PickCode is 115-specific; other providers use ID directly.
PickCode string `json:"pick_code,omitempty"`
}
@@ -59,7 +58,7 @@ type DirectLink struct {
// Provider is the common cloud-disk interface.
type Provider interface {
- // Type returns the provider key (TypeQuark / Type115).
+ // Type returns the provider key.
Type() string
// Ping validates the stored credentials (cookie). Cheap, used by the
// storage-config Test() probe.
@@ -71,6 +70,21 @@ type Provider interface {
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.
@@ -79,8 +93,6 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
client = http.DefaultClient
}
switch typ {
- case TypeQuark:
- return newQuark(cfg, client), nil
case Type115:
return new115(cfg, client), nil
case TypeCloudDrive2:
@@ -94,7 +106,7 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
// IsCloudType reports whether typ is a cloud-disk provider.
func IsCloudType(typ string) bool {
- return typ == TypeQuark || typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
+ return typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
}
// str coerces a config value to a trimmed string.
@@ -121,5 +133,5 @@ func boolish(v any) bool {
}
}
-// defaultUA is a desktop browser UA accepted by both 115 and quark.
+// 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"
diff --git a/internal/service/cloud/cloud_115_test.go b/internal/service/cloud/cloud_115_test.go
new file mode 100644
index 0000000..631ebf7
--- /dev/null
+++ b/internal/service/cloud/cloud_115_test.go
@@ -0,0 +1,212 @@
+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
+}
diff --git a/internal/service/cloud/cloud_test.go b/internal/service/cloud/cloud_test.go
index 8df5440..41e200f 100644
--- a/internal/service/cloud/cloud_test.go
+++ b/internal/service/cloud/cloud_test.go
@@ -2,115 +2,15 @@ package cloud
import (
"context"
- "encoding/base64"
"encoding/json"
- "fmt"
"net/http"
"net/http/httptest"
- "strconv"
"strings"
"testing"
"time"
)
-func TestQuarkListAndResolve(t *testing.T) {
- var gotCookie string
- srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- gotCookie = r.Header.Get("Cookie")
- switch {
- case r.URL.Path == "/file/sort":
- if r.URL.Query().Get("pdir_fid") != "0" {
- t.Errorf("unexpected pdir_fid %q", r.URL.Query().Get("pdir_fid"))
- }
- w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"d1","file_name":"Movies","dir":true,"size":0},
- {"fid":"f1","file_name":"Inception.mkv","dir":false,"size":123}]}}`))
- case r.URL.Path == "/file/download":
- if r.Method != http.MethodPost {
- t.Errorf("download must be POST, got %s", r.Method)
- }
- w.Write([]byte(`{"status":200,"code":0,"data":[{"fid":"f1","download_url":"https://cdn.quark/x.mkv?sign=1"}]}`))
- default:
- t.Errorf("unexpected path %s", r.URL.Path)
- }
- }))
- defer srv.Close()
-
- p, err := New(TypeQuark, map[string]any{"cookie": "kps=abc", "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) != 2 || !entries[0].IsDir || entries[1].Name != "Inception.mkv" || entries[1].Size != 123 {
- t.Fatalf("unexpected entries: %#v", entries)
- }
- if gotCookie != "kps=abc" {
- t.Fatalf("cookie not forwarded: %q", gotCookie)
- }
- link, err := p.Resolve(context.Background(), "f1")
- if err != nil {
- t.Fatalf("resolve: %v", err)
- }
- if link.URL != "https://cdn.quark/x.mkv?sign=1" {
- t.Fatalf("bad url: %s", link.URL)
- }
- if !link.Proxy {
- t.Fatalf("quark should default to proxy mode")
- }
- if link.Headers["Cookie"] != "kps=abc" {
- t.Fatalf("resolve must carry cookie header: %#v", link.Headers)
- }
-}
-
-func TestQuarkListPaginates(t *testing.T) {
- srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/sort" {
- t.Fatalf("unexpected path %s", r.URL.Path)
- }
- page, _ := strconv.Atoi(r.URL.Query().Get("_page"))
- w.Write([]byte(`{"status":200,"code":0,"data":{"list":[` + quarkPagePayload(page) + `]}}`))
- }))
- defer srv.Close()
-
- p, err := New(TypeQuark, map[string]any{"cookie": "kps=abc", "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 != "f100" || entries[100].Name != "Movie.100.mkv" {
- t.Fatalf("last entry wrong: %#v", entries[100])
- }
-}
-
-func quarkPagePayload(page int) string {
- count := 100
- offset := 0
- if page > 1 {
- count = 1
- offset = 100
- }
- items := make([]string, 0, count)
- for i := 0; i < count; i++ {
- n := offset + i
- items = append(items, fmt.Sprintf(`{"fid":"f%d","file_name":"Movie.%03d.mkv","dir":false,"size":%d}`, n, n, n))
- }
- return strings.Join(items, ",")
-}
-
func TestDeprecatedProviderPlaybackOverrideKeysAreIgnored(t *testing.T) {
- quark := newQuark(map[string]any{"cookie": "c", "force_302": "true"}, http.DefaultClient)
- if !quark.proxy {
- t.Fatalf("quark should keep safe proxy mode; force_302 is deprecated")
- }
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")
@@ -121,197 +21,6 @@ func TestDeprecatedProviderPlaybackOverrideKeysAreIgnored(t *testing.T) {
}
}
-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 TestCloudDrive2WebDAVListAndResolve(t *testing.T) {
var gotAuth, gotDepth, gotRange string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -523,6 +232,144 @@ func TestOpenListListUsesAPIUsernamePasswordWithoutWebDAVFallback(t *testing.T)
}
}
+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 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/%E5%A4%96%E8%AF%AD%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)
+ }
+}
+
func TestOpenListListAPIFailureDoesNotFallbackToWebDAV(t *testing.T) {
var davSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -804,22 +651,12 @@ func TestUnsupportedProvider(t *testing.T) {
if _, err := New("dropbox", nil, nil); err != ErrUnsupported {
t.Fatalf("want ErrUnsupported, got %v", err)
}
-}
-
-func containsAll(s string, subs ...string) bool {
- for _, sub := range subs {
- found := false
- for i := 0; i+len(sub) <= len(s); i++ {
- if s[i:i+len(sub)] == sub {
- found = true
- break
- }
- }
- if !found {
- return false
- }
+ 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")
}
- return true
}
var _ = time.Second
diff --git a/internal/service/cloud/clouddrive2.go b/internal/service/cloud/clouddrive2.go
index 2c0cd3b..64ddfa6 100644
--- a/internal/service/cloud/clouddrive2.go
+++ b/internal/service/cloud/clouddrive2.go
@@ -1,25 +1,19 @@
package cloud
import (
- "bytes"
"context"
"encoding/base64"
- "encoding/json"
- "encoding/xml"
"fmt"
- "io"
"net/http"
"net/url"
"path"
- "sort"
- "strconv"
"strings"
)
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
//
-// CloudDrive2 already integrates many cloud disks (115 / 123 / Aliyun / Quark
-// and more). Treating it as a WebDAV-backed cloud provider lets MediaStationGo
+// 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 {
@@ -76,133 +70,6 @@ func (p *cloudDrive2Provider) Ping(ctx context.Context) error {
return err
}
-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) 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, fmt.Errorf("%s: api list %s returned http %d", p.name, 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) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
if err := p.validate(); err != nil {
return nil, err
@@ -239,323 +106,6 @@ func (p *cloudDrive2Provider) Resolve(ctx context.Context, fileRef string) (*Dir
return &DirectLink{URL: p.urlFor(ref), Headers: headers, Proxy: p.proxy}, 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, fmt.Errorf("%s: api get %s returned http %d", p.name, 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) 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 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 (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
-}
-
-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 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 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
- }
-}
-
func (p *cloudDrive2Provider) validate() error {
if p.base == nil || p.base.Scheme == "" || p.base.Host == "" {
return fmt.Errorf("%s: missing WebDAV URL", p.name)
@@ -563,17 +113,6 @@ func (p *cloudDrive2Provider) validate() error {
return 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)
- }
-}
-
func webDAVURLFromConfig(cfg map[string]any, defaultDAVPath string) string {
rawURL := str(cfg["url"])
if rawURL == "" {
@@ -669,153 +208,6 @@ func ensureDefaultDAVPath(rawURL, defaultDAVPath string) string {
return rawURL
}
-func (p *cloudDrive2Provider) decorateDAVStatusError(resp *http.Response, target string) error {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
- detail := compactDAVErrorBody(string(body))
- if detail == "" {
- 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
-}
-
-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 = `
-
-
-
-
-
-
-`
-
-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"`
-}
-
-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 normalizeCloudDAVPath(p string) string {
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
if p == "" || p == "." {
@@ -835,20 +227,6 @@ func sameCloudDAVPath(a, b string) bool {
return strings.TrimRight(normalizeCloudDAVPath(a), "/") == strings.TrimRight(normalizeCloudDAVPath(b), "/")
}
-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)
-}
-
-func parseDAVSize(raw string) int64 {
- n, _ := strconv.ParseInt(strings.TrimSpace(raw), 10, 64)
- return n
-}
-
func firstNonEmpty(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
diff --git a/internal/service/cloud/clouddrive2_dav.go b/internal/service/cloud/clouddrive2_dav.go
new file mode 100644
index 0000000..f817f39
--- /dev/null
+++ b/internal/service/cloud/clouddrive2_dav.go
@@ -0,0 +1,292 @@
+package cloud
+
+import (
+ "context"
+ "encoding/base64"
+ "encoding/xml"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "path"
+ "strconv"
+ "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)
+ }
+}
+
+func (p *cloudDrive2Provider) decorateDAVStatusError(resp *http.Response, target string) error {
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
+ detail := compactDAVErrorBody(string(body))
+ if detail == "" {
+ 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
+}
+
+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 = `
+
+
+
+
+
+
+`
+
+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
+}
diff --git a/internal/service/cloud/clouddrive2_mutation.go b/internal/service/cloud/clouddrive2_mutation.go
new file mode 100644
index 0000000..11a97f4
--- /dev/null
+++ b/internal/service/cloud/clouddrive2_mutation.go
@@ -0,0 +1,233 @@
+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)
+}
diff --git a/internal/service/cloud/clouddrive2_openlist.go b/internal/service/cloud/clouddrive2_openlist.go
new file mode 100644
index 0000000..b3a8cf0
--- /dev/null
+++ b/internal/service/cloud/clouddrive2_openlist.go
@@ -0,0 +1,352 @@
+package cloud
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "path"
+ "sort"
+ "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, fmt.Errorf("%s: api list %s returned http %d", p.name, 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, fmt.Errorf("%s: api get %s returned http %d", p.name, 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 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 (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
+}
+
+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)
+}
diff --git a/internal/service/cloud/quark.go b/internal/service/cloud/quark.go
deleted file mode 100644
index fbb3bf0..0000000
--- a/internal/service/cloud/quark.go
+++ /dev/null
@@ -1,176 +0,0 @@
-package cloud
-
-import (
- "bytes"
- "context"
- "encoding/json"
- "fmt"
- "io"
- "net/http"
- "net/url"
- "strings"
-)
-
-// quarkProvider implements the 夸克网盘 cloud disk using cookie auth.
-//
-// Quark's web API is plain JSON over HTTPS keyed by a session cookie; no
-// request-body encryption is required (unlike 115). The resolved download_url
-// is tied to the session, so playback runs in proxy mode by default.
-type quarkProvider struct {
- cookie string
- ua string
- base string // override for tests; defaults to quarkBase
- client *http.Client
- proxy bool
-}
-
-const quarkBase = "https://drive-pc.quark.cn/1/clouddrive"
-
-func newQuark(cfg map[string]any, client *http.Client) *quarkProvider {
- base := str(cfg["base"])
- if base == "" {
- base = quarkBase
- }
- ua := str(cfg["ua"])
- if ua == "" {
- ua = defaultUA
- }
- // Quark download links require the session cookie + UA, so the host must
- // reverse-proxy. The global cloud playback setting decides whether clients
- // receive a STRMURL entry or a /Videos stream entry; this provider only
- // reports whether the resolved upstream URL itself is safe for raw 302.
- proxy := true
- return &quarkProvider{
- cookie: str(cfg["cookie"]),
- ua: ua,
- base: strings.TrimRight(base, "/"),
- client: client,
- proxy: proxy,
- }
-}
-
-func (q *quarkProvider) Type() string { return TypeQuark }
-
-func (q *quarkProvider) do(ctx context.Context, method, path string, body io.Reader) (*http.Response, error) {
- req, err := http.NewRequestWithContext(ctx, method, q.base+path, body)
- if err != nil {
- return nil, err
- }
- req.Header.Set("Cookie", q.cookie)
- req.Header.Set("User-Agent", q.ua)
- req.Header.Set("Accept", "application/json, text/plain, */*")
- req.Header.Set("Referer", "https://pan.quark.cn/")
- if body != nil {
- req.Header.Set("Content-Type", "application/json")
- }
- return q.client.Do(req)
-}
-
-type quarkResp struct {
- Status int `json:"status"`
- Code int `json:"code"`
- Message string `json:"message"`
- Data json.RawMessage `json:"data"`
-}
-
-func (q *quarkProvider) Ping(ctx context.Context) error {
- if q.cookie == "" {
- return fmt.Errorf("quark: missing cookie")
- }
- _, err := q.List(ctx, "0")
- return err
-}
-
-func (q *quarkProvider) List(ctx context.Context, dirID string) ([]FileEntry, error) {
- if dirID == "" {
- dirID = "0"
- }
- const pageSize = 100
- out := make([]FileEntry, 0, pageSize)
- for page := 1; ; page++ {
- query := url.Values{}
- query.Set("pr", "ucpro")
- query.Set("fr", "pc")
- query.Set("uc_param_str", "")
- query.Set("pdir_fid", dirID)
- query.Set("_page", fmt.Sprint(page))
- query.Set("_size", fmt.Sprint(pageSize))
- query.Set("_fetch_total", "1")
- query.Set("_sort", "file_type:asc,updated_at:desc")
- path := "/file/sort?" + query.Encode()
- resp, err := q.do(ctx, http.MethodGet, path, nil)
- if err != nil {
- return nil, err
- }
- var r quarkResp
- err = json.NewDecoder(resp.Body).Decode(&r)
- _ = resp.Body.Close()
- if err != nil {
- return nil, fmt.Errorf("quark: decode list: %w", err)
- }
- if r.Code != 0 && r.Status != 200 {
- return nil, fmt.Errorf("quark: list failed: %s", r.Message)
- }
- var data struct {
- List []struct {
- Fid string `json:"fid"`
- FileName string `json:"file_name"`
- Dir bool `json:"dir"`
- Size int64 `json:"size"`
- } `json:"list"`
- }
- if err := json.Unmarshal(r.Data, &data); err != nil {
- return nil, fmt.Errorf("quark: decode list data: %w", err)
- }
- for _, it := range data.List {
- out = append(out, FileEntry{
- ID: it.Fid,
- Name: it.FileName,
- IsDir: it.Dir,
- Size: it.Size,
- })
- }
- if len(data.List) < pageSize {
- break
- }
- }
- return out, nil
-}
-
-func (q *quarkProvider) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
- if fileRef == "" {
- return nil, fmt.Errorf("quark: empty file id")
- }
- payload, _ := json.Marshal(map[string]any{"fids": []string{fileRef}})
- resp, err := q.do(ctx, http.MethodPost, "/file/download?pr=ucpro&fr=pc&uc_param_str=", bytes.NewReader(payload))
- if err != nil {
- return nil, err
- }
- defer resp.Body.Close()
- var r quarkResp
- if err := json.NewDecoder(resp.Body).Decode(&r); err != nil {
- return nil, fmt.Errorf("quark: decode download: %w", err)
- }
- if r.Code != 0 && r.Status != 200 {
- return nil, fmt.Errorf("quark: download failed: %s", r.Message)
- }
- var data []struct {
- DownloadURL string `json:"download_url"`
- Fid string `json:"fid"`
- }
- if err := json.Unmarshal(r.Data, &data); err != nil {
- return nil, fmt.Errorf("quark: decode download data: %w", err)
- }
- if len(data) == 0 || data[0].DownloadURL == "" {
- return nil, fmt.Errorf("quark: no download url returned")
- }
- return &DirectLink{
- URL: data[0].DownloadURL,
- Headers: map[string]string{
- "Cookie": q.cookie,
- "User-Agent": q.ua,
- "Referer": "https://pan.quark.cn/",
- },
- Proxy: q.proxy,
- }, nil
-}
diff --git a/internal/service/cloud_auto_category.go b/internal/service/cloud_auto_category.go
new file mode 100644
index 0000000..4d330cc
--- /dev/null
+++ b/internal/service/cloud_auto_category.go
@@ -0,0 +1,180 @@
+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 {
+ base := BuildCloudLibraryPath(provider, "", 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 {
+ info, ok := ParseCloudLibraryMount(path)
+ if !ok {
+ return ""
+ }
+ parts := strmSlashParts(info.DisplayDir)
+ if len(parts) <= 1 {
+ return ""
+ }
+ parts = parts[:len(parts)-1]
+ categoryParts := cloudAutoCategoryParts(parts)
+ if len(categoryParts) == 0 {
+ return ""
+ }
+ return strings.Join(categoryParts, "/")
+}
+
+func cloudAutoCategoryParts(parts []string) []string {
+ for i, part := range parts {
+ root := strmCanonicalRoot(part)
+ if root != "" {
+ if i+1 >= len(parts) {
+ return nil
+ }
+ category := strings.TrimSpace(parts[i+1])
+ if cloudAutoCategoryRootMatches(root, category) {
+ return []string{root, category}
+ }
+ return nil
+ }
+ if root := strmCategoryRoot(part); root != "" {
+ return []string{root, strings.TrimSpace(part)}
+ }
+ }
+ return 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
+}
+
+func (s *ScannerService) ensureCloudAutoCategoryLibrary(ctx context.Context, rootLib *model.Library, provider, displayDir string) (*model.Library, error) {
+ displayDir = normalizeCloudMountDir(provider, displayDir)
+ if s == nil || s.repo == nil || s.repo.DB == nil || rootLib == nil || provider == "" || displayDir == "" {
+ return rootLib, nil
+ }
+ if existing := s.findCloudLibraryByDisplayDir(ctx, provider, displayDir); existing != nil {
+ return existing, nil
+ }
+ path := BuildCloudAutoCategoryLibraryPath(provider, displayDir)
+ if path == "" {
+ return rootLib, nil
+ }
+ name := cloudMountDirBase(displayDir)
+ if name == "" {
+ name = displayDir
+ }
+ lib := &model.Library{
+ Name: name,
+ Path: path,
+ Type: InferCloudMountMediaType(displayDir, name),
+ Enabled: true,
+ }
+ if err := s.repo.Library.Create(ctx, lib); err != nil {
+ if existing := s.findCloudLibraryByDisplayDir(ctx, provider, displayDir); existing != nil {
+ return existing, nil
+ }
+ return nil, 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))
+ }
+ return lib, nil
+}
+
+func (s *ScannerService) findCloudLibraryByDisplayDir(ctx context.Context, provider, displayDir string) *model.Library {
+ if s == nil || s.repo == nil || s.repo.Library == nil {
+ return 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
+ }
+ displayDir = normalizeCloudMountDir(provider, displayDir)
+ for _, lib := range libs {
+ info, ok := ParseCloudLibraryMount(lib.Path)
+ if !ok || info.Provider != provider || normalizeCloudMountDir(provider, info.DisplayDir) != displayDir {
+ continue
+ }
+ return &lib
+ }
+ return nil
+}
+
+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
+}
diff --git a/internal/service/cloud_metadata.go b/internal/service/cloud_metadata.go
index 5a5a637..a9acdd4 100644
--- a/internal/service/cloud_metadata.go
+++ b/internal/service/cloud_metadata.go
@@ -3,8 +3,8 @@ package service
import (
"context"
"encoding/xml"
- "net/url"
"path/filepath"
+ "strconv"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
@@ -13,6 +13,9 @@ import (
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
}
@@ -20,6 +23,9 @@ 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 {
@@ -40,7 +46,11 @@ func newCloudSidecarSet(typ string, entries []cloud.FileEntry) cloudSidecarSet {
case ".nfo":
set.nfoByName[strings.ToLower(name)] = ref
set.nfoByBase[base] = ref
- case ".jpg", ".jpeg", ".png", ".webp":
+ 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
}
}
@@ -60,12 +70,23 @@ func (s *ScannerService) cloudDirectoryMetadata(ctx context.Context, typ, displa
if ref == "" {
continue
}
- if local, _, err := s.readCloudNFO(ctx, typ, ref, true); err == nil && local != nil {
+ 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
}
}
- meta = applyCloudDirectoryArtwork(typ, sidecars, meta)
+ 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
}
@@ -77,11 +98,12 @@ func (s *ScannerService) cloudFileMetadata(ctx context.Context, typ, displayPath
seriesLike = seriesLike || season > 0 || episode > 0
meta := cloneLocalMetadata(inherited)
if hinted, _ := pathHintMetadata(displayPath, seriesLike); hinted != nil {
- meta = mergeCloudMetadata(meta, hinted)
+ 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{}
@@ -93,13 +115,39 @@ func (s *ScannerService) cloudFileMetadata(ctx context.Context, typ, displayPath
}
}
}
- meta = applyCloudFileArtwork(typ, sidecars, base, meta)
+ 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 {
+ 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
@@ -116,18 +164,18 @@ func (s *ScannerService) readCloudNFO(ctx context.Context, typ, ref string, seri
return meta, &doc, nil
}
-func applyCloudDirectoryArtwork(typ string, sidecars cloudSidecarSet, meta *LocalMetadata) *LocalMetadata {
+func applyCloudDirectoryArtwork(typ, displayDir string, sidecars cloudSidecarSet, meta *LocalMetadata) *LocalMetadata {
if meta == nil {
meta = &LocalMetadata{}
}
if meta.PosterURL == "" {
- if ref := firstCloudImageRef(sidecars, "poster", "folder", "cover", "show", "tvshow"); ref != "" {
+ 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, "fanart", "backdrop", "background", "landscape"); ref != "" {
+ if ref := firstCloudImageRef(sidecars, cloudBackdropNameCandidates(cloudDirectoryArtworkBases(displayDir), "fanart", "backdrop", "background", "landscape")...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
@@ -135,24 +183,19 @@ func applyCloudDirectoryArtwork(typ string, sidecars cloudSidecarSet, meta *Loca
return meta
}
-func applyCloudFileArtwork(typ string, sidecars cloudSidecarSet, base string, meta *LocalMetadata) *LocalMetadata {
+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,
- base+"-poster", base+".poster", base+"-cover", base+".cover", base+"-thumb", base+".thumb",
- "poster", "folder", "cover", "movie", "show", "thumb",
- ); ref != "" {
+ 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,
- base+"-fanart", base+".fanart", base+"-backdrop", base+".backdrop", base+"-background", base+".background",
- "fanart", "backdrop", "background", "landscape",
- ); ref != "" {
+ if ref := firstCloudImageRef(sidecars, cloudBackdropNameCandidates(bases, "fanart", "backdrop", "background", "landscape")...); ref != "" {
meta.BackdropURL = cloudPlaybackURL(typ, ref)
meta.HasArtwork = true
}
@@ -160,17 +203,8 @@ func applyCloudFileArtwork(typ string, sidecars cloudSidecarSet, base string, me
return meta
}
-func firstCloudImageRef(sidecars cloudSidecarSet, names ...string) string {
- for _, name := range names {
- if ref := sidecars.imageByBase[strings.ToLower(strings.TrimSpace(name))]; ref != "" {
- return ref
- }
- }
- return ""
-}
-
func cloudShowNFOCandidates(displayDir string) []string {
- names := []string{"tvshow.nfo", "series.nfo", "show.nfo"}
+ names := []string{"tvshow.nfo", "series.nfo", "show.nfo", "movie.nfo"}
base := strings.TrimSpace(pathBaseSlash(displayDir))
if base != "" {
names = append(names, base+".nfo")
@@ -178,6 +212,112 @@ func cloudShowNFOCandidates(displayDir string) []string {
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
+}
+
func mergeCloudMetadata(dst, src *LocalMetadata) *LocalMetadata {
if src == nil {
return dst
@@ -191,6 +331,9 @@ func mergeCloudMetadata(dst, src *LocalMetadata) *LocalMetadata {
if src.OriginalName != "" {
dst.OriginalName = src.OriginalName
}
+ if src.EpisodeTitle != "" {
+ dst.EpisodeTitle = src.EpisodeTitle
+ }
if src.AdultCode != "" {
dst.AdultCode = src.AdultCode
}
@@ -243,6 +386,38 @@ func mergeCloudMetadata(dst, src *LocalMetadata) *LocalMetadata {
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.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
@@ -256,7 +431,7 @@ func cloudMetadataUseful(meta *LocalMetadata) bool {
}
func cloudPlaybackURL(typ, ref string) string {
- return "/api/cloud/play/" + typ + "?ref=" + url.QueryEscape(ref)
+ return CloudArtworkURL(typ, ref)
}
func joinCloudDisplayPath(parent, child string) string {
@@ -280,3 +455,15 @@ func pathBaseSlash(value string) string {
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]
+}
diff --git a/internal/service/cloud_metadata_artwork.go b/internal/service/cloud_metadata_artwork.go
new file mode 100644
index 0000000..87a477b
--- /dev/null
+++ b/internal/service/cloud_metadata_artwork.go
@@ -0,0 +1,80 @@
+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 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)
+}
diff --git a/internal/service/cloud_metadata_json.go b/internal/service/cloud_metadata_json.go
new file mode 100644
index 0000000..2c62dd3
--- /dev/null
+++ b/internal/service/cloud_metadata_json.go
@@ -0,0 +1,264 @@
+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"),
+ 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(firstJSONString(obj, "release_date", "releasedate", "premiered", "aired", "date"))
+ }
+ 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 ""
+}
diff --git a/internal/service/cloud_mount.go b/internal/service/cloud_mount.go
index b36c18e..cf74226 100644
--- a/internal/service/cloud_mount.go
+++ b/internal/service/cloud_mount.go
@@ -89,6 +89,9 @@ func FindCloudMountConflict(libs []model.Library, provider, scanDir, displayDir
ScanDir: normalizeCloudMountDir(provider, scanDir),
}
for _, lib := range libs {
+ if CloudLibraryAutoCategory(lib) {
+ continue
+ }
existing, ok := ParseCloudLibraryMount(lib.Path)
if !ok || existing.Provider != candidate.Provider {
continue
@@ -115,6 +118,9 @@ func CloudLibraryShadowed(libs []model.Library, lib model.Library) *CloudMountCo
if existing.ID == lib.ID || !existing.Enabled {
continue
}
+ if CloudLibraryAutoCategory(existing) {
+ continue
+ }
info, ok := ParseCloudLibraryMount(existing.Path)
if !ok || info.Provider != current.Provider {
continue
@@ -143,6 +149,7 @@ func FilterDisplayCloudLibraries(ctx context.Context, repo *repository.Container
if len(libs) == 0 {
return libs
}
+ libs = FilterDeprecatedNativeCloudLibraries(libs)
counts := cloudLibraryMediaCounts(ctx, repo, libs)
collapsed := make([]model.Library, 0, len(libs))
byKey := make(map[string]int, len(libs))
@@ -173,6 +180,12 @@ func FilterScannableCloudLibraries(ctx context.Context, repo *repository.Contain
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)
@@ -190,6 +203,21 @@ func FilterScannableCloudLibraries(ctx context.Context, repo *repository.Contain
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
+}
+
func NormalizeCloudLibraryDisplayNames(libs []model.Library) []model.Library {
out := make([]model.Library, 0, len(libs))
for _, lib := range libs {
@@ -388,12 +416,17 @@ func ExpandMediaVisibilityForMergedCloudLibraries(ctx context.Context, repo *rep
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 = expandMergedLibraryIDs(ctx, repo, visibility.AllowedLibraryIDs)
+ visibility.AllowedLibraryIDs = expandMergedLibraryIDsFromLibraries(libs, visibility.AllowedLibraryIDs)
}
if len(visibility.HiddenLibraryIDs) > 0 {
- visibility.HiddenLibraryIDs = expandMergedLibraryIDs(ctx, repo, visibility.HiddenLibraryIDs)
+ visibility.HiddenLibraryIDs = expandMergedLibraryIDsFromLibraries(libs, visibility.HiddenLibraryIDs)
}
+ visibility.HiddenLibraryIDs = appendUniqueLibraryIDs(visibility.HiddenLibraryIDs, DeprecatedNativeCloudLibraryIDs(libs)...)
return visibility
}
@@ -405,40 +438,63 @@ func expandMergedLibraryIDs(ctx context.Context, repo *repository.Container, ids
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))
- seen := map[string]struct{}{}
- add := func(id string) {
- id = strings.TrimSpace(id)
- if id == "" {
- return
- }
- if _, ok := seen[id]; ok {
- return
- }
- seen[id] = struct{}{}
- out = append(out, id)
- }
for _, id := range ids {
lib, ok := byID[id]
if !ok {
- add(id)
+ out = appendUniqueLibraryIDs(out, id)
continue
}
for _, mergedID := range MergedLibraryIDs(libs, lib) {
- add(mergedID)
+ out = appendUniqueLibraryIDs(out, mergedID)
}
}
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, more ...string) []string {
+ for _, id := range more {
+ id = strings.TrimSpace(id)
+ if id == "" {
+ continue
+ }
+ found := false
+ for _, existing := range ids {
+ if existing == id {
+ found = true
+ break
+ }
+ }
+ if !found {
+ ids = append(ids, id)
+ }
+ }
+ return ids
+}
+
func CloudMountProviderLabel(provider string) string {
switch strings.TrimSpace(provider) {
- case cloud.TypeQuark:
- return "夸克网盘"
+ case LegacyQuarkProvider:
+ return "已停用网盘"
case cloud.Type115:
return "115 网盘"
case cloud.TypeCloudDrive2:
@@ -556,7 +612,7 @@ func normalizeCloudMountDir(provider, value string) string {
}
value = strings.ReplaceAll(value, "\\", "/")
value = strings.Trim(strings.TrimSpace(value), "/")
- if value == "." || ((provider == cloud.Type115 || provider == cloud.TypeQuark) && value == "0") {
+ if value == "." || ((provider == cloud.Type115 || provider == LegacyQuarkProvider) && value == "0") {
return ""
}
return value
diff --git a/internal/service/cloud_mount_filter_test.go b/internal/service/cloud_mount_filter_test.go
index ff02466..bc4b582 100644
--- a/internal/service/cloud_mount_filter_test.go
+++ b/internal/service/cloud_mount_filter_test.go
@@ -6,22 +6,14 @@ import (
"time"
"github.com/ShukeBta/MediaStationGo/internal/config"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestFilterDisplayCloudLibrariesPrefersPopulatedCanonicalDuplicate(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
now := time.Now()
oldEmpty := model.Library{
@@ -64,13 +56,7 @@ func TestFilterDisplayCloudLibrariesPrefersPopulatedCanonicalDuplicate(t *testin
}
func TestFilterDisplayCloudLibrariesMergesCloudMountIntoExistingLibrary(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ 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}
@@ -98,14 +84,38 @@ func TestFilterDisplayCloudLibrariesMergesCloudMountIntoExistingLibrary(t *testi
}
}
+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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ 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}
@@ -165,13 +175,7 @@ func TestListMediaVisibleIncludesMergedCloudLibraryItems(t *testing.T) {
}
func TestListMediaVisibleUsesSpecificCloudChildLibraryAsDisplayTarget(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ 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}
@@ -205,13 +209,7 @@ func TestListMediaVisibleUsesSpecificCloudChildLibraryAsDisplayTarget(t *testing
}
func TestStartAllCloudLibraryScansIncludesMergedCloudMounts(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ 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}
@@ -231,6 +229,64 @@ func TestStartAllCloudLibraryScansIncludesMergedCloudMounts(t *testing.T) {
}
}
+func TestAutoCategoryCloudLibrariesDoNotShadowRootOrScan(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)
+ }
+ }
+
+ 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 len(display) != 2 {
+ t.Fatalf("display libraries = %#v, want root plus auto category", display)
+ }
+ scannable := FilterScannableCloudLibraries(t.Context(), repos, libs)
+ if len(scannable) != 1 || scannable[0].ID != root.ID {
+ t.Fatalf("scannable libraries = %#v, want only root", scannable)
+ }
+
+ 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 root queued", statuses)
+ }
+}
+
+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 {
diff --git a/internal/service/cloud_path_repair.go b/internal/service/cloud_path_repair.go
index 6e0e554..da9ca98 100644
--- a/internal/service/cloud_path_repair.go
+++ b/internal/service/cloud_path_repair.go
@@ -14,11 +14,12 @@ import (
// "Movie (2025) {tmdb-123}" so existing placeholder rows can be scraped
// without requiring another successful filesystem or cloud provider traversal.
//
-// 传入 libraryID 时只修复该媒体库的行;为空则修复全库。
+// 传入 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).
@@ -35,8 +36,8 @@ func (c *Container) RepairCloudPathMetadata(ctx context.Context, libraryID ...st
"LOWER(path) LIKE ?",
}, " OR ")+")",
"%tmdb%", "%tmdbid%", "%douban%", "%db%", "%bangumi%", "%bgm%", "%thetvdb%", "%tvdb%")
- if len(libraryID) > 0 && strings.TrimSpace(libraryID[0]) != "" {
- query = query.Where("library_id = ?", strings.TrimSpace(libraryID[0]))
+ if len(libraryIDs) > 0 {
+ query = query.Where("library_id IN ?", libraryIDs)
}
err := query.FindInBatches(&rows, 500, func(_ *gorm.DB, _ int) error {
@@ -112,13 +113,15 @@ func cloudPathRepairShouldReplaceTitle(current, hinted string) bool {
return len([]rune(current)) > len([]rune(hinted))*2
}
-
// RepairAndRescrapeResult 汇总一次「全库修复+重刮」的结果。
type RepairAndRescrapeResult struct {
- Repaired int `json:"repaired"` // 从路径占位符回填外部 ID 的媒体数
- Libraries int `json:"libraries"` // 参与重刮的媒体库数
- Matched int `json:"matched"` // 重刮后成功匹配的媒体数
- Reset int `json:"reset"` // 被重置为 pending 以便重刮的剧集行数
+ 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 的行重置为
@@ -130,16 +133,17 @@ type RepairAndRescrapeResult struct {
// 重刮」会跳过,导致「无法修复」。源头已在 local_metadata.go 修正,这里把脏的
// matched 剧集行放回 pending,借重刮写回正确的整剧 ID / 原名。
//
-// libraryID 为空时处理全库;非空时仅该库。返回被重置的行数。
-func (c *Container) resetEpisodicMatchedForRescrape(ctx context.Context, libraryID string) (int, error) {
+// 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 id := strings.TrimSpace(libraryID); id != "" {
- q = q.Where("library_id = ?", id)
+ if len(ids) > 0 {
+ q = q.Where("library_id IN ?", ids)
}
res := q.Update("scrape_status", "pending")
if res.Error != nil {
@@ -148,21 +152,44 @@ func (c *Container) resetEpisodicMatchedForRescrape(ctx context.Context, library
reset := int(res.RowsAffected)
if reset > 0 && c.Log != nil {
c.Log.Info("episodic matched rows reset to pending for rescrape",
- zap.String("library", strings.TrimSpace(libraryID)),
+ zap.String("libraries", strings.Join(ids, ",")),
zap.Int("reset", reset))
}
return reset, nil
}
+func compactLibraryIDs(ids ...string) []string {
+ out := make([]string, 0, len(ids))
+ for _, id := range ids {
+ out = appendUniqueLibraryIDs(out, id)
+ }
+ return out
+}
+
// RepairAndRescrapeAllLibraries 修复并重刮所有媒体库:先从媒体路径中的
// {tmdb-123}/{bangumi-456} 等占位符回填缺失或错误的外部 ID(回填后会把相关
// 行的 scrape_status 重置为 pending),随后逐个媒体库重刮(含 no_match 重试),
// 让此前因空 ID / 脏 ID 无法刮削的媒体重新匹配到正确数据。
-func (c *Container) RepairAndRescrapeAllLibraries(ctx context.Context) (RepairAndRescrapeResult, error) {
+func repairRescrapeOptions(values ...ScrapeOptions) ScrapeOptions {
+ options := ScrapeOptions{RetryNoMatch: true, IncludeMatched: true}
+ if len(values) > 0 {
+ options = values[0]
+ options.RetryNoMatch = true
+ options.IncludeMatched = 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
@@ -170,7 +197,7 @@ func (c *Container) RepairAndRescrapeAllLibraries(ctx context.Context) (RepairAn
result.Repaired = repaired
// 重置全库脏的 matched 剧集行(单集 id 污染整剧字段),让其下方重刮一并修正。
- if reset, err := c.resetEpisodicMatchedForRescrape(ctx, ""); err != nil {
+ if reset, err := c.resetEpisodicMatchedForRescrape(ctx); err != nil {
return result, err
} else {
result.Reset = reset
@@ -195,20 +222,36 @@ func (c *Container) RepairAndRescrapeAllLibraries(ctx context.Context) (RepairAn
}
result.Libraries++
// retryNoMatch=true:连之前匹配失败的也再试一次,因为这次可能已回填到正确 ID。
- matched, err := c.Scraper.EnrichLibrary(ctx, lib.ID, true)
+ 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 += matched
+ 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)
+ }
}
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("matched", result.Matched),
+ zap.Int("processed", result.Processed),
+ zap.Int("errors", result.Errors))
}
return result, nil
}
@@ -216,20 +259,25 @@ func (c *Container) RepairAndRescrapeAllLibraries(ctx context.Context) (RepairAn
// RepairAndRescrapeLibrary 修复并重刮单个媒体库:先从该库媒体路径中的占位符
// 回填缺失/错误的外部 ID(重置相关行 scrape_status=pending),再对该库重刮
// (含 no_match 重试)。用于「按媒体库」单独触发修复,不影响其它库。
-func (c *Container) RepairAndRescrapeLibrary(ctx context.Context, libraryID string) (RepairAndRescrapeResult, error) {
+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
}
- repaired, err := c.RepairCloudPathMetadata(ctx, libraryID)
+ 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, libraryID); err != nil {
+ if reset, err := c.resetEpisodicMatchedForRescrape(ctx, libraryIDs...); err != nil {
return result, err
} else {
result.Reset = reset
@@ -240,16 +288,31 @@ func (c *Container) RepairAndRescrapeLibrary(ctx context.Context, libraryID stri
}
result.Libraries = 1
// retryNoMatch=true:连之前匹配失败的也再试一次,因为这次可能已回填到正确 ID。
- matched, err := c.Scraper.EnrichLibrary(ctx, libraryID, true)
+ scrapeResult, err := c.Scraper.EnrichLibraryDetailedWithOptions(ctx, libraryID, scrapeOptions)
if err != nil {
return result, err
}
- result.Matched = matched
+ 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)
+ }
+ }
if c.Log != nil {
c.Log.Info("repair and rescrape library done",
zap.String("library", libraryID),
zap.Int("repaired", result.Repaired),
- zap.Int("matched", result.Matched))
+ zap.Int("reclassified", result.Reclassified),
+ zap.Int("matched", result.Matched),
+ zap.Int("processed", result.Processed),
+ zap.Int("errors", result.Errors))
}
return result, nil
}
diff --git a/internal/service/cloud_path_repair_test.go b/internal/service/cloud_path_repair_test.go
index 51c2ec1..7ede847 100644
--- a/internal/service/cloud_path_repair_test.go
+++ b/internal/service/cloud_path_repair_test.go
@@ -3,24 +3,63 @@ package service
import (
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
+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.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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
container := &Container{Repo: repository.New(db), Log: zap.NewNop()}
rows := []model.Media{
@@ -65,3 +104,66 @@ func TestResetEpisodicMatchedForRescrape(t *testing.T) {
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)
+ }
+}
diff --git a/internal/service/device_service.go b/internal/service/device_service.go
index 6ade617..6d4f5fc 100644
--- a/internal/service/device_service.go
+++ b/internal/service/device_service.go
@@ -5,6 +5,7 @@ import (
"crypto/sha256"
"encoding/hex"
"fmt"
+ "sort"
"strings"
"time"
@@ -27,8 +28,9 @@ import (
// a Telegram notification is sent before a destructive action; every policy
// defaults to OFF.
type DeviceService struct {
- log *zap.Logger
- repo *repository.Container
+ log *zap.Logger
+ repo *repository.Container
+ sessions *SessionTrackerService
// notifyUser sends a Telegram message to the local user (resolved to their
// Telegram binding). Wired by the bot service; nil disables notifications.
@@ -45,6 +47,10 @@ func (s *DeviceService) SetNotifier(fn func(ctx context.Context, userID, text st
s.notifyUser = fn
}
+func (s *DeviceService) SetSessionTracker(tracker *SessionTrackerService) {
+ s.sessions = tracker
+}
+
// fingerprint derives a stable terminal hash from the device name. Client/app
// names are deliberately ignored so one phone/TV/PC using multiple apps is
// still counted as one terminal device; Client remains a login channel label.
@@ -70,12 +76,19 @@ func (s *DeviceService) RecordLogin(ctx context.Context, userID, deviceID, devic
if userID == "" {
return
}
+ username := ""
+ if u, _ := s.repo.User.FindByID(ctx, userID); u != nil {
+ username = u.Username
+ }
if deviceID == "" {
// Fall back to a fingerprint-derived id so headless clients still count.
deviceID = "fp-" + fingerprint(client, deviceName)
}
+ if s.sessions != nil {
+ s.sessions.RecordLogin(ctx, userID, username, deviceID, deviceName, client, ip)
+ }
fp := fingerprint(client, deviceName)
- now := time.Now()
+ now := s.now()
existing, _ := s.repo.UserDevice.Find(ctx, userID, deviceID)
mismatch := false
@@ -124,10 +137,17 @@ func (s *DeviceService) RecordPlayback(ctx context.Context, userID, deviceID, de
if userID == "" {
return
}
+ username := ""
+ if u, _ := s.repo.User.FindByID(ctx, userID); u != nil {
+ username = u.Username
+ }
if deviceID == "" {
deviceID = "fp-" + fingerprint(client, deviceName)
}
- now := time.Now()
+ if s.sessions != nil {
+ s.sessions.RecordPlayback(ctx, userID, username, deviceID, deviceName, client, "", "", 0, 0, false)
+ }
+ now := s.now()
existing, _ := s.repo.UserDevice.Find(ctx, userID, deviceID)
if existing == nil {
existing = &model.UserDevice{
@@ -172,7 +192,7 @@ func (s *DeviceService) registerFingerprintWarning(ctx context.Context, userID,
s.log.Info("anti-share: skipping protected account", zap.String("user", u.Username), zap.String("reason", reason))
return
}
- now := time.Now()
+ now := s.now()
if u.LastShareWarnAt != nil && now.Sub(*u.LastShareWarnAt) < time.Minute {
return // debounce burst
}
@@ -202,7 +222,7 @@ func (s *DeviceService) disableForPolicy(ctx context.Context, userID, reason str
s.log.Info("device policy: skipping protected account", zap.String("user", u.Username), zap.String("reason", reason))
return
}
- now := time.Now()
+ now := s.now()
_ = s.repo.User.UpdateFields(ctx, userID, map[string]any{
"is_active": false,
"last_share_warn_at": &now,
@@ -300,7 +320,62 @@ func (s *DeviceService) KickAllDevices(ctx context.Context, userID string) error
// ListDevices returns the device sessions for a user.
func (s *DeviceService) ListDevices(ctx context.Context, userID string) ([]model.UserDevice, error) {
- return s.repo.UserDevice.ListByUser(ctx, userID)
+ rows, err := s.repo.UserDevice.ListByUser(ctx, userID)
+ if err != nil {
+ return nil, err
+ }
+ if s.sessions == nil {
+ return rows, nil
+ }
+ now := s.sessions.now()
+ byDevice := make(map[string]int, len(rows))
+ for i := range rows {
+ byDevice[rows[i].DeviceID] = i
+ }
+ for _, sess := range s.sessions.ListByUser(ctx, userID) {
+ online := sess.LastActivityAt.After(now.Add(-realtimeSessionOnlineTTL))
+ if idx, ok := byDevice[sess.DeviceID]; ok {
+ if sess.LastActivityAt.After(rows[idx].LastSeenAt) {
+ rows[idx].LastSeenAt = sess.LastActivityAt
+ }
+ if sess.DeviceName != "" {
+ rows[idx].DeviceName = sess.DeviceName
+ }
+ if sess.Client != "" {
+ rows[idx].Client = sess.Client
+ }
+ if sess.RemoteEndPoint != "" {
+ rows[idx].LastIP = sess.RemoteEndPoint
+ }
+ if sess.LastPlaybackAt != nil {
+ rows[idx].LastPlayAt = sess.LastPlaybackAt
+ }
+ rows[idx].Realtime = true
+ rows[idx].Online = online
+ rows[idx].Playing = sess.IsPlaying && online
+ continue
+ }
+ row := model.UserDevice{
+ UserID: userID,
+ DeviceID: sess.DeviceID,
+ DeviceName: sess.DeviceName,
+ Client: sess.Client,
+ Fingerprint: fingerprint(sess.Client, sess.DeviceName),
+ LastIP: sess.RemoteEndPoint,
+ FirstSeenAt: sess.LastActivityAt,
+ LastSeenAt: sess.LastActivityAt,
+ LastPlayAt: sess.LastPlaybackAt,
+ Realtime: true,
+ Online: online,
+ Playing: sess.IsPlaying && online,
+ }
+ row.ID = "rt:" + sess.ID
+ rows = append(rows, row)
+ }
+ sort.SliceStable(rows, func(i, j int) bool {
+ return rows[i].LastSeenAt.After(rows[j].LastSeenAt)
+ })
+ return rows, nil
}
// IsDeviceKicked reports whether a (user, device) pair was kicked and should be
@@ -313,6 +388,17 @@ func (s *DeviceService) IsDeviceKicked(ctx context.Context, userID, deviceID str
return err == nil && d != nil && d.Kicked
}
+func (s *DeviceService) UserRecentlyActive(ctx context.Context, userID string, within time.Duration) bool {
+ return s.sessions != nil && s.sessions.UserRecentlyActive(ctx, userID, within)
+}
+
+func (s *DeviceService) now() time.Time {
+ if s != nil && s.sessions != nil && s.sessions.now != nil {
+ return s.sessions.now()
+ }
+ return time.Now()
+}
+
func (s *DeviceService) notify(ctx context.Context, userID, text string) {
if s.notifyUser != nil {
s.notifyUser(ctx, userID, text)
@@ -372,7 +458,11 @@ func (s *DeviceService) userMatchesCleanupRule(ctx context.Context, u *model.Use
if days < 1 {
days = r.WindowDaysMin
}
- ok := u.LastLoginAt != nil && u.LastLoginAt.After(time.Now().Add(-time.Duration(days)*24*time.Hour))
+ cutoff := time.Now().Add(-time.Duration(days) * 24 * time.Hour)
+ ok := u.LastLoginAt != nil && u.LastLoginAt.After(cutoff)
+ if !ok && s.sessions != nil {
+ ok = s.sessions.UserRecentlyActive(ctx, u.ID, time.Duration(days)*24*time.Hour)
+ }
return ok, fmt.Sprintf("%s:%d 天内登录", r.Name, days)
case "signin_streak":
rec, _ := s.repo.SignIn.Get(ctx, u.ID)
diff --git a/internal/service/discover.go b/internal/service/discover.go
index 2a868bf..b40213e 100644
--- a/internal/service/discover.go
+++ b/internal/service/discover.go
@@ -24,6 +24,7 @@ type DiscoverService struct {
log *zap.Logger
tmdb *TMDbProvider
client *http.Client
+ images *ImageProxy
}
// NewDiscoverService is the constructor.
@@ -74,6 +75,7 @@ func (d *DiscoverService) TMDbSection(ctx context.Context, key string) ([]Extern
Rating: item.Rating,
TMDbID: item.TMDbID,
SubscribeKeyword: buildSubscribeKeyword(item.Title, item.Year),
+ SubscribeAliases: buildSubscribeAliases(item.Title, item.OriginalName, item.Year),
})
}
return out, nil
@@ -257,6 +259,7 @@ func (d *DoubanProvider) Discover(ctx context.Context, key string) ([]ExternalMe
Rating: float32(rating),
DoubanID: subject.ID,
SubscribeKeyword: subject.Title,
+ SubscribeAliases: buildSubscribeAliases(subject.Title, "", 0),
})
}
return out, nil
@@ -320,6 +323,7 @@ func (b *BangumiProvider) Calendar(ctx context.Context) ([]ExternalMediaResult,
Rating: item.Rating.Score,
BangumiID: item.ID,
SubscribeKeyword: buildSubscribeKeyword(title, year),
+ SubscribeAliases: buildSubscribeAliases(title, item.Name, year),
})
if len(out) >= 24 {
return out, nil
diff --git a/internal/service/discover_artwork.go b/internal/service/discover_artwork.go
new file mode 100644
index 0000000..5af80b6
--- /dev/null
+++ b/internal/service/discover_artwork.go
@@ -0,0 +1,101 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "sync"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+const (
+ discoverArtworkPrefetchLimit = 48
+ discoverArtworkPrefetchConcurrency = 4
+ discoverArtworkPrefetchTimeout = 45 * time.Second
+)
+
+func (d *DiscoverService) SetImageProxy(images *ImageProxy) *DiscoverService {
+ if d != nil {
+ d.images = images
+ }
+ return d
+}
+
+func (d *DiscoverService) WarmMatchArtwork(items []Match) int {
+ urls := make([]string, 0, len(items)*2)
+ for _, item := range items {
+ urls = append(urls, item.PosterURL, item.BackdropURL)
+ }
+ return d.warmArtworkURLs(urls)
+}
+
+func (d *DiscoverService) WarmExternalArtwork(items []ExternalMediaResult) int {
+ urls := make([]string, 0, len(items)*2)
+ for _, item := range items {
+ urls = append(urls, item.PosterURL, item.BackdropURL)
+ }
+ return d.warmArtworkURLs(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()
+}
diff --git a/internal/service/discover_artwork_test.go b/internal/service/discover_artwork_test.go
new file mode 100644
index 0000000..e0ca7d5
--- /dev/null
+++ b/internal/service/discover_artwork_test.go
@@ -0,0 +1,98 @@
+package service
+
+import (
+ "io"
+ "net/http"
+ "os"
+ "path/filepath"
+ "strings"
+ "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 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(strings.NewReader("image:" + req.URL.Path)),
+ 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 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)
+}
diff --git a/internal/service/download_add.go b/internal/service/download_add.go
new file mode 100644
index 0000000..56eacad
--- /dev/null
+++ b/internal/service/download_add.go
@@ -0,0 +1,298 @@
+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.title, strings.TrimSpace(req.meta.SubscriptionID) != ""); ok {
+ return existing, ErrDownloadAlreadyExists
+ }
+ _ = d.ReloadConfig(ctx)
+ if !d.qb.IsConfigured() {
+ return nil, errors.New("no default downloader configured")
+ }
+ if d.torrentExistsByIdentity(ctx, req.title) {
+ task, err := d.createTask(ctx, userID, urlStr, req.savePath, req.meta)
+ if err != nil {
+ return nil, err
+ }
+ return task, ErrDownloadAlreadyExists
+ }
+ if err := d.addPreparedDownloadToClient(ctx, urlStr, &req); err != nil {
+ return nil, err
+ }
+ return d.createTask(ctx, userID, urlStr, req.savePath, req.meta)
+}
+
+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) error {
+ var siteFetchErr error
+ if d.site != nil {
+ if data, name, err := d.site.FetchTorrentFile(ctx, urlStr); err == nil {
+ if err := d.qb.AddTorrentFileWithCategory(ctx, data, name, req.savePath, req.qbitCategory); err != nil {
+ return err
+ }
+ if strings.TrimSpace(req.meta.Title) == "" {
+ req.meta.Title = strings.TrimSuffix(name, path.Ext(name))
+ }
+ return nil
+ } else {
+ siteFetchErr = err
+ }
+ }
+ if err := d.qb.AddTorrentWithCategory(ctx, urlStr, req.savePath, req.qbitCategory); err != nil {
+ if siteFetchErr != nil && !strings.Contains(siteFetchErr.Error(), "no matching PT site") {
+ return errors.Join(err, siteFetchErr)
+ }
+ return err
+ }
+ return nil
+}
+
+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) 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 {
+ wantSeason, wantEpisode := ParseEpisode(title)
+ if wantSeason <= 0 {
+ wantSeason = 1
+ }
+ if wantEpisode <= 0 {
+ return true
+ }
+ for _, row := range rows {
+ rowSeason, rowEpisode := localMediaRowSeasonEpisode(row)
+ if rowEpisode == wantEpisode && rowSeason == wantSeason {
+ return true
+ }
+ if rowEpisode <= 0 && isSeriesPackTitle(row.Title+" "+row.OriginalName+" "+row.Path) {
+ return true
+ }
+ }
+ return false
+}
+
+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
+}
+
+func (d *DownloadService) findExistingDownloadTask(ctx context.Context, title string, allowDeletedReadd bool) (*model.DownloadTask, bool) {
+ key := downloadTaskIdentityKey(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
+ }
+ for i := range rows {
+ if allowDeletedReadd {
+ if !downloadTaskBlocksReadd(rows[i].Status) {
+ continue
+ }
+ } else if !downloadTaskBlocksDuplicate(rows[i].Status) {
+ continue
+ }
+ current := downloadTaskIdentityKey(rows[i].Title)
+ if current == key || strings.Contains(current, key) || strings.Contains(key, current) {
+ return &rows[i], true
+ }
+ }
+ return nil, false
+}
+
+func (d *DownloadService) torrentExistsByIdentity(ctx context.Context, title string) bool {
+ query := downloadTaskIdentityKey(title)
+ if query == "" {
+ return false
+ }
+ live, err := d.qb.List(ctx, "")
+ if err != nil {
+ return false
+ }
+ for _, torrent := range live {
+ current := downloadTaskIdentityKey(torrent.Name)
+ if current == "" {
+ continue
+ }
+ if current == query || strings.Contains(current, query) || strings.Contains(query, current) {
+ return true
+ }
+ }
+ return false
+}
+
+func (d *DownloadService) createTask(ctx context.Context, userID, urlStr, savePath string, meta DownloadTaskMeta) (*model.DownloadTask, error) {
+ title := strings.TrimSpace(meta.Title)
+ if title == "" {
+ title = publicDownloadTitle(urlStr)
+ }
+ t := &model.DownloadTask{
+ UserID: userID,
+ SubscriptionID: strings.TrimSpace(meta.SubscriptionID),
+ Source: "qbittorrent",
+ 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
+}
diff --git a/internal/service/download_add_test.go b/internal/service/download_add_test.go
new file mode 100644
index 0000000..9f31f17
--- /dev/null
+++ b/internal/service/download_add_test.go
@@ -0,0 +1,488 @@
+package service
+
+import (
+ "errors"
+ "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 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 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)
+ }
+}
+
+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 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)
+ }
+}
+
+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 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)
+ }
+}
+
+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))
+ }
+}
diff --git a/internal/service/download_clients_test.go b/internal/service/download_clients_test.go
index 7d70349..cc152bb 100644
--- a/internal/service/download_clients_test.go
+++ b/internal/service/download_clients_test.go
@@ -5,22 +5,14 @@ import (
"net/http/httptest"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestDownloadClientCreateNormalizesHostAndClearsDefault(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
svc := NewDownloadClientService(zap.NewNop(), repos)
@@ -64,13 +56,7 @@ func TestDownloadClientCreateNormalizesHostAndClearsDefault(t *testing.T) {
}
func TestDownloadClientCreateMakesFirstEnabledClientDefault(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
svc := NewDownloadClientService(zap.NewNop(), repos)
@@ -89,13 +75,7 @@ func TestDownloadClientCreateMakesFirstEnabledClientDefault(t *testing.T) {
}
func TestDownloadClientRejectsUnsupportedHostScheme(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
svc := NewDownloadClientService(zap.NewNop(), repository.New(db))
if _, err := svc.Create(t.Context(), DownloadClientInput{
@@ -109,13 +89,7 @@ func TestDownloadClientRejectsUnsupportedHostScheme(t *testing.T) {
}
func TestDownloadClientRejectsUnsafeEndpointParts(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{})
svc := NewDownloadClientService(zap.NewNop(), repository.New(db))
for _, host := range []string{
@@ -198,13 +172,7 @@ func TestAria2AdapterUsesNormalizedRPCURL(t *testing.T) {
}
func TestDownloadClientDeleteClearsLegacyQBitConnectionWhenNoDefault(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ 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",
@@ -243,13 +211,7 @@ func TestDownloadClientDeleteClearsLegacyQBitConnectionWhenNoDefault(t *testing.
}
func TestDownloadClientUpdateClearsLegacyQBitConnectionWhenDefaultDisabled(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ 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)
diff --git a/internal/service/download_completion.go b/internal/service/download_completion.go
new file mode 100644
index 0000000..3201e88
--- /dev/null
+++ b/internal/service/download_completion.go
@@ -0,0 +1,224 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "os"
+ "path/filepath"
+ "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 := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey); ok {
+ return &task, true
+ }
+ if strings.TrimSpace(torrent.ContentPath) != "" {
+ if task, ok := findMatchingTaskByTorrentIdentity(filepath.Base(torrent.ContentPath), 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 {
+ // 常见路径映射:qBittorrent容器路径 -> MediaStationGo容器路径
+ 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
+ }
+ for _, candidate := range []string{
+ torrent.ContentPath,
+ filepath.Join(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
+}
diff --git a/internal/service/download_completion_state.go b/internal/service/download_completion_state.go
new file mode 100644
index 0000000..cf2826a
--- /dev/null
+++ b/internal/service/download_completion_state.go
@@ -0,0 +1,229 @@
+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 (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 != "" {
+ return hash
+ }
+ parts := []string{torrent.Name, torrent.ContentPath, torrent.SavePath}
+ 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 || strings.TrimSpace(torrent.Name) == "" {
+ return
+ }
+ matched, ok := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey)
+ if !ok {
+ return
+ }
+ status := torrent.State
+ if torrent.Progress >= 1 {
+ 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 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))
+ for _, row := range rows {
+ key := normalizeTorrentName(row.Title)
+ if key != "" {
+ out[key] = row
+ }
+ }
+ return out
+}
+
+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.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.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"
+}
diff --git a/internal/service/download_config_test.go b/internal/service/download_config_test.go
new file mode 100644
index 0000000..3e4fcc0
--- /dev/null
+++ b/internal/service/download_config_test.go
@@ -0,0 +1,233 @@
+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 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 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 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 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)
+ }
+}
diff --git a/internal/service/download_identity.go b/internal/service/download_identity.go
new file mode 100644
index 0000000..13712d7
--- /dev/null
+++ b/internal/service/download_identity.go
@@ -0,0 +1,145 @@
+package service
+
+import (
+ "fmt"
+ "net/url"
+ "path"
+ "regexp"
+ "strings"
+ "unicode"
+)
+
+var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`)
+
+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)
+}
+
+func downloadMediaIdentityKey(name string) string {
+ name = strings.ToLower(strings.TrimSpace(name))
+ if name == "" {
+ return ""
+ }
+ title, year := CleanQuery(name)
+ titleKey := normalizeAvailabilityComparable(title)
+ if titleKey == "" {
+ titleKey = normalizeAvailabilityComparable(availabilityQuery(name, ""))
+ }
+ if titleKey == "" {
+ return ""
+ }
+ season, episode := ParseEpisode(name)
+ parts := []string{titleKey}
+ if year > 0 {
+ parts = append(parts, fmt.Sprintf("y%d", year))
+ }
+ if episode > 0 {
+ if season <= 0 {
+ season = 1
+ }
+ parts = append(parts, fmt.Sprintf("s%02de%03d", season, 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()
+}
diff --git a/internal/service/download_live_snapshot.go b/internal/service/download_live_snapshot.go
new file mode 100644
index 0000000..4d5334a
--- /dev/null
+++ b/internal/service/download_live_snapshot.go
@@ -0,0 +1,44 @@
+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 {
+ if d == nil {
+ return nil
+ }
+ now := d.currentTime()
+ d.mu.Lock()
+ defer d.mu.Unlock()
+ if d.liveTorrentsAt.IsZero() {
+ return nil
+ }
+ if maxAge > 0 && now.Sub(d.liveTorrentsAt) > maxAge {
+ return nil
+ }
+ return cloneQBitTorrentSlice(d.liveTorrents)
+}
+
+func cloneQBitTorrentSlice(in []QBitTorrent) []QBitTorrent {
+ if len(in) == 0 {
+ return nil
+ }
+ return append([]QBitTorrent(nil), in...)
+}
diff --git a/internal/service/download_notification.go b/internal/service/download_notification.go
new file mode 100644
index 0000000..2288bdd
--- /dev/null
+++ b/internal/service/download_notification.go
@@ -0,0 +1,85 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "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 := fmt.Sprintf("任务:%s\n保存路径:%s\nHash:%s", name, firstNonEmpty(torrent.ContentPath, torrent.SavePath), torrent.Hash)
+ 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
+ }
+}
diff --git a/internal/service/download_polling.go b/internal/service/download_polling.go
new file mode 100644
index 0000000..e831f9f
--- /dev/null
+++ b/internal/service/download_polling.go
@@ -0,0 +1,201 @@
+package service
+
+import (
+ "context"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+const completedTorrentOrganizeQueueSize = 64
+
+var completedTorrentOrganizeCooldown = 3 * time.Second
+
+// poll fans out qBittorrent /torrents/info 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.qb.List(ctx, "")
+ if err != nil {
+ 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, torrent.Progress >= 1.0, 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 := findMatchingTaskByTorrentIdentity(torrent.Name, 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()
+}
diff --git a/internal/service/download_views.go b/internal/service/download_views.go
new file mode 100644
index 0000000..0966eb8
--- /dev/null
+++ b/internal/service/download_views.go
@@ -0,0 +1,197 @@
+package service
+
+import (
+ "math"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type DownloadTaskView struct {
+ ID string `json:"id"`
+ Source string `json:"source"`
+ Title string `json:"title"`
+ PosterURL string `json:"poster_url,omitempty"`
+ BackdropURL string `json:"backdrop_url,omitempty"`
+ Overview string `json:"overview,omitempty"`
+ SavePath string `json:"save_path"`
+ MediaType string `json:"media_type,omitempty"`
+ MediaCategory string `json:"media_category,omitempty"`
+ Status string `json:"status"`
+ Progress float32 `json:"progress"`
+ State string `json:"state,omitempty"`
+ DLSpeed int64 `json:"dlspeed,omitempty"`
+ UpSpeed int64 `json:"upspeed,omitempty"`
+ Size int64 `json:"size,omitempty"`
+ Downloaded int64 `json:"downloaded,omitempty"`
+ NumSeeds int `json:"num_seeds,omitempty"`
+ NumLeechs int `json:"num_leechs,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
+ UpdatedAt time.Time `json:"updated_at"`
+}
+
+type DownloadTorrentView struct {
+ Hash string `json:"hash"`
+ 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 != "" {
+ liveByKey[key] = torrent
+ }
+ }
+ taskByKey := map[string]model.DownloadTask{}
+ for _, row := range rows {
+ key := normalizeTorrentName(row.Title)
+ if key != "" {
+ taskByKey[key] = row
+ }
+ }
+
+ taskViews := make([]DownloadTaskView, 0, len(rows))
+ for _, row := range rows {
+ view := downloadTaskView(row, QBitTorrent{})
+ if torrent, ok := findMatchingTorrent(row.Title, 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 := findMatchingTask(torrent.Name, 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,
+ 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,
+ 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 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.Contains(currentKey, key) || strings.Contains(key, currentKey) {
+ return torrent, true
+ }
+ }
+ return QBitTorrent{}, false
+}
+
+func findMatchingTask(title string, taskByKey map[string]model.DownloadTask) (model.DownloadTask, bool) {
+ key := normalizeTorrentName(title)
+ if key == "" {
+ return model.DownloadTask{}, false
+ }
+ if row, ok := taskByKey[key]; ok {
+ return row, true
+ }
+ for currentKey, row := range taskByKey {
+ if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) {
+ return row, true
+ }
+ }
+ return model.DownloadTask{}, false
+}
+
+func downloadedBytes(size int64, progress float32) int64 {
+ if size <= 0 || progress <= 0 {
+ return 0
+ }
+ 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 ""
+}
diff --git a/internal/service/downloads.go b/internal/service/downloads.go
index e891d6f..9b55c51 100644
--- a/internal/service/downloads.go
+++ b/internal/service/downloads.go
@@ -17,19 +17,10 @@ package service
import (
"context"
- "crypto/sha1"
"errors"
- "fmt"
- "math"
- "net/url"
- "os"
- "path"
- "path/filepath"
- "regexp"
"strings"
"sync"
"time"
- "unicode"
"go.uber.org/zap"
@@ -56,6 +47,9 @@ type DownloadService struct {
organizeOnce sync.Once
prevStates map[string]bool // hash -> wasCompleted
pollInitialized bool
+ liveTorrents []QBitTorrent
+ liveTorrentsAt time.Time
+ now func() time.Time
organizeQueue chan QBitTorrent
organizeQueued map[string]struct{}
}
@@ -76,14 +70,8 @@ func (d *DownloadService) SetNotifyChannels(notify *NotifyChannelService) {
d.notify = notify
}
-var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`)
-
const settingDownloadClientsManaged = "download_clients.managed"
-const completedTorrentOrganizeQueueSize = 64
-
-var completedTorrentOrganizeCooldown = 3 * time.Second
-
// ErrDownloadAlreadyExists tells callers that the requested resource is already
// tracked locally or present in qBittorrent. Subscriptions treat this as a
// successful dedup hit, not as a retryable enqueue failure.
@@ -98,69 +86,6 @@ func IsDownloadDedupError(err error) bool {
return errors.Is(err, ErrDownloadAlreadyExists) || errors.Is(err, ErrMediaAlreadyInLibrary)
}
-// 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 DownloadTaskView struct {
- ID string `json:"id"`
- Source string `json:"source"`
- Title string `json:"title"`
- PosterURL string `json:"poster_url,omitempty"`
- BackdropURL string `json:"backdrop_url,omitempty"`
- Overview string `json:"overview,omitempty"`
- SavePath string `json:"save_path"`
- MediaType string `json:"media_type,omitempty"`
- MediaCategory string `json:"media_category,omitempty"`
- Status string `json:"status"`
- Progress float32 `json:"progress"`
- State string `json:"state,omitempty"`
- DLSpeed int64 `json:"dlspeed,omitempty"`
- UpSpeed int64 `json:"upspeed,omitempty"`
- Size int64 `json:"size,omitempty"`
- Downloaded int64 `json:"downloaded,omitempty"`
- NumSeeds int `json:"num_seeds,omitempty"`
- NumLeechs int `json:"num_leechs,omitempty"`
- CreatedAt time.Time `json:"created_at"`
- UpdatedAt time.Time `json:"updated_at"`
-}
-
-type DownloadTorrentView struct {
- Hash string `json:"hash"`
- 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"`
-}
-
// NewDownloadService is the constructor.
func NewDownloadService(log *zap.Logger, repo *repository.Container, hub *Hub, organizer *OrganizerService, site ...*SiteService) *DownloadService {
var siteSvc *SiteService
@@ -175,6 +100,7 @@ func NewDownloadService(log *zap.Logger, repo *repository.Container, hub *Hub, o
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{}),
@@ -270,354 +196,6 @@ func (d *DownloadService) soleEnabledQBitClient(ctx context.Context) (*model.Dow
return selected, nil
}
-// 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) {
- if urlStr == "" {
- return nil, 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
- }
- if !meta.AllowExistingLibrary && d.localMediaAlreadyExists(ctx, title) {
- return nil, ErrMediaAlreadyInLibrary
- }
- if existing, ok := d.findExistingDownloadTask(ctx, title, strings.TrimSpace(meta.SubscriptionID) != ""); ok {
- return existing, ErrDownloadAlreadyExists
- }
- _ = d.ReloadConfig(ctx)
- if !d.qb.IsConfigured() {
- return nil, errors.New("no default downloader configured")
- }
- if d.torrentExistsByIdentity(ctx, title) {
- task, err := d.createTask(ctx, userID, urlStr, savePath, meta)
- if err != nil {
- return nil, err
- }
- return task, ErrDownloadAlreadyExists
- }
- var siteFetchErr error
- qbitCategory := strings.TrimSpace(meta.MediaCategory)
- if d.site != nil {
- if data, name, err := d.site.FetchTorrentFile(ctx, urlStr); err == nil {
- if err := d.qb.AddTorrentFileWithCategory(ctx, data, name, savePath, qbitCategory); err != nil {
- return nil, err
- }
- if strings.TrimSpace(meta.Title) == "" {
- meta.Title = strings.TrimSuffix(name, path.Ext(name))
- }
- return d.createTask(ctx, userID, urlStr, savePath, meta)
- } else {
- siteFetchErr = err
- }
- }
- if err := d.qb.AddTorrentWithCategory(ctx, urlStr, savePath, qbitCategory); err != nil {
- if siteFetchErr != nil && !strings.Contains(siteFetchErr.Error(), "no matching PT site") {
- return nil, errors.Join(err, siteFetchErr)
- }
- return nil, err
- }
- return d.createTask(ctx, userID, urlStr, savePath, meta)
-}
-
-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) localMediaAlreadyExists(ctx context.Context, title string) bool {
- if d == nil || d.repo == nil || d.repo.DB == nil {
- return false
- }
- if !d.repo.DB.Migrator().HasTable(&model.Media{}) {
- return false
- }
- queries := localAvailabilityTitleCandidates(title)
- if len(queries) == 0 {
- return 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 false
- }
-
- wantSeason, wantEpisode := ParseEpisode(title)
- if wantSeason <= 0 {
- wantSeason = 1
- }
- if wantEpisode <= 0 {
- return true
- }
- for _, row := range rows {
- 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
- }
- if rowEpisode == wantEpisode && rowSeason == wantSeason {
- return true
- }
- if rowEpisode <= 0 && isSeriesPackTitle(row.Title+" "+row.OriginalName+" "+row.Path) {
- return true
- }
- }
- return false
-}
-
-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 (d *DownloadService) findExistingDownloadTask(ctx context.Context, title string, allowDeletedReadd bool) (*model.DownloadTask, bool) {
- key := downloadTaskIdentityKey(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
- }
- for i := range rows {
- if allowDeletedReadd {
- if !downloadTaskBlocksReadd(rows[i].Status) {
- continue
- }
- } else if !downloadTaskBlocksDuplicate(rows[i].Status) {
- continue
- }
- current := downloadTaskIdentityKey(rows[i].Title)
- if current == key || strings.Contains(current, key) || strings.Contains(key, current) {
- return &rows[i], true
- }
- }
- return nil, false
-}
-
-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 (d *DownloadService) torrentExistsByIdentity(ctx context.Context, title string) bool {
- query := downloadTaskIdentityKey(title)
- if query == "" {
- return false
- }
- live, err := d.qb.List(ctx, "")
- if err != nil {
- return false
- }
- for _, torrent := range live {
- current := downloadTaskIdentityKey(torrent.Name)
- if current == "" {
- continue
- }
- if current == query || strings.Contains(current, query) || strings.Contains(query, current) {
- return true
- }
- }
- return false
-}
-
-func downloadTaskIdentityKey(name string) string {
- if key := downloadMediaIdentityKey(name); key != "" {
- return key
- }
- return normalizedDownloadTitleKey(name)
-}
-
-func downloadMediaIdentityKey(name string) string {
- name = strings.ToLower(strings.TrimSpace(name))
- if name == "" {
- return ""
- }
- title, year := CleanQuery(name)
- titleKey := normalizeAvailabilityComparable(title)
- if titleKey == "" {
- titleKey = normalizeAvailabilityComparable(availabilityQuery(name, ""))
- }
- if titleKey == "" {
- return ""
- }
- season, episode := ParseEpisode(name)
- parts := []string{titleKey}
- if year > 0 {
- parts = append(parts, fmt.Sprintf("y%d", year))
- }
- if episode > 0 {
- if season <= 0 {
- season = 1
- }
- parts = append(parts, fmt.Sprintf("s%02de%03d", season, 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 (d *DownloadService) createTask(ctx context.Context, userID, urlStr, savePath string, meta DownloadTaskMeta) (*model.DownloadTask, error) {
- title := strings.TrimSpace(meta.Title)
- if title == "" {
- title = publicDownloadTitle(urlStr)
- }
- t := &model.DownloadTask{
- UserID: userID,
- SubscriptionID: strings.TrimSpace(meta.SubscriptionID),
- Source: "qbittorrent",
- 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
-}
-
-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 (d *DownloadService) TorrentExistsByName(ctx context.Context, name string) bool {
query := normalizeTorrentName(name)
if query == "" {
@@ -639,17 +217,6 @@ func (d *DownloadService) TorrentExistsByName(ctx context.Context, name string)
return false
}
-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()
-}
-
// List returns every persisted download task augmented with live data
// from qBittorrent when available.
func (d *DownloadService) List(ctx context.Context) ([]model.DownloadTask, []QBitTorrent, error) {
@@ -667,151 +234,6 @@ func (d *DownloadService) List(ctx context.Context) ([]model.DownloadTask, []QBi
return rows, live, nil
}
-func DownloadViews(rows []model.DownloadTask, live []QBitTorrent) ([]DownloadTaskView, []DownloadTorrentView) {
- liveByKey := map[string]QBitTorrent{}
- for _, torrent := range live {
- key := normalizeTorrentName(torrent.Name)
- if key != "" {
- liveByKey[key] = torrent
- }
- }
- taskByKey := map[string]model.DownloadTask{}
- for _, row := range rows {
- key := normalizeTorrentName(row.Title)
- if key != "" {
- taskByKey[key] = row
- }
- }
-
- taskViews := make([]DownloadTaskView, 0, len(rows))
- for _, row := range rows {
- view := downloadTaskView(row, QBitTorrent{})
- if torrent, ok := findMatchingTorrent(row.Title, 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 := findMatchingTask(torrent.Name, 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,
- 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,
- 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 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.Contains(currentKey, key) || strings.Contains(key, currentKey) {
- return torrent, true
- }
- }
- return QBitTorrent{}, false
-}
-
-func findMatchingTask(title string, taskByKey map[string]model.DownloadTask) (model.DownloadTask, bool) {
- key := normalizeTorrentName(title)
- if key == "" {
- return model.DownloadTask{}, false
- }
- if row, ok := taskByKey[key]; ok {
- return row, true
- }
- for currentKey, row := range taskByKey {
- if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) {
- return row, true
- }
- }
- return model.DownloadTask{}, false
-}
-
-func downloadedBytes(size int64, progress float32) int64 {
- if size <= 0 || progress <= 0 {
- return 0
- }
- 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 ""
-}
-
// Delete removes a torrent (and optionally its files) from qBittorrent.
func (d *DownloadService) Delete(ctx context.Context, hash string, withFiles bool) error {
hash = strings.TrimSpace(hash)
@@ -897,658 +319,3 @@ func (d *DownloadService) RelocateTorrent(ctx context.Context, hash, location st
}
return d.qb.SetLocation(ctx, hash, strings.TrimSpace(location))
}
-
-// poll fans out qBittorrent /torrents/info 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.qb.List(ctx, "")
- if err != nil {
- 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.mu.Lock()
- if d.prevStates == nil {
- d.prevStates = make(map[string]bool)
- }
- firstSnapshot := !d.pollInitialized
- if firstSnapshot {
- d.pollInitialized = true
- }
- d.mu.Unlock()
-
- for _, torrent := range live {
- stateKey := completedTorrentQueueKey(torrent)
- complete := torrent.Progress >= 1.0
- matchedTask, hasTask := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey)
- autoOrganize := d.downloadAutoOrganizeEnabled(ctx)
- catchupRecorded := hasTask && d.completedTorrentCatchupRecorded(ctx, torrent)
- taskNeedsOrganize := hasTask && !catchupRecorded &&
- autoOrganize &&
- (downloadTaskNeedsCompletion(matchedTask) || recentlyCompletedTorrent(torrent, time.Now()))
- d.syncDownloadTaskProgress(ctx, torrent, taskByKey)
- if stateKey == "" {
- continue
- }
-
- shouldQueue := false
- d.mu.Lock()
- wasComplete, wasSeen := d.prevStates[stateKey]
- switch {
- case complete && (firstSnapshot || !wasSeen):
- // 首次快照里已完成的种子:此前一律标记「已见过」并跳过整理,
- // 导致「下载完成时应用恰好不在线/正在重启」的种子永远不会被
- // 自动整理入库。现在对最近完成的种子补一次整理
- // (onTorrentComplete 内部仍受 organize.auto 开关约束,且
- // 整理对已存在的目标文件幂等跳过)。
- d.prevStates[stateKey] = true
- if taskNeedsOrganize {
- shouldQueue = true
- }
- case complete && !wasComplete:
- shouldQueue = true
- case complete && taskNeedsOrganize:
- shouldQueue = true
- case complete:
- d.prevStates[stateKey] = true
- default:
- d.prevStates[stateKey] = false
- }
- d.mu.Unlock()
-
- if shouldQueue && d.enqueueCompletedTorrent(torrent) {
- 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()
-}
-
-// 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 (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 != "" {
- return hash
- }
- parts := []string{torrent.Name, torrent.ContentPath, torrent.SavePath}
- 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 || strings.TrimSpace(torrent.Name) == "" {
- return
- }
- matched, ok := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey)
- if !ok {
- return
- }
- status := torrent.State
- if torrent.Progress >= 1 {
- 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 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))
- for _, row := range rows {
- key := normalizeTorrentName(row.Title)
- if key != "" {
- out[key] = row
- }
- }
- return out
-}
-
-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.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.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"
-}
-
-// 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.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(taskRow),
- MediaCategory: firstNonEmpty(downloadTaskMediaCategory(taskRow), torrent.Category),
- AllowReplace: allowReplace,
- })
- if err != nil {
- 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) 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)
- 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 = "下载任务"
- }
- body := fmt.Sprintf("任务:%s\n保存路径:%s\nHash:%s", name, firstNonEmpty(torrent.ContentPath, torrent.SavePath), torrent.Hash)
- data := map[string]interface{}{}
- // resource_title 供 Telegram 模板从发布名提取季集(SxxEyy)与版本(分辨率/编码/
- // 字幕组等)信息;隐藏不直接展示。优先用 torrent 原始名(信息最全)。
- if rt := strings.TrimSpace(torrent.Name); rt != "" {
- data["resource_title"] = rt
- }
- if task != nil {
- if strings.TrimSpace(task.PosterURL) != "" {
- data["poster_url"] = task.PosterURL
- }
- if strings.TrimSpace(task.BackdropURL) != "" {
- data["backdrop_url"] = task.BackdropURL
- }
- if strings.TrimSpace(task.MediaType) != "" {
- data["media_type"] = task.MediaType
- }
- if strings.TrimSpace(task.MediaCategory) != "" {
- data["media_category"] = task.MediaCategory
- }
- if strings.TrimSpace(task.Title) != "" {
- data["title"] = task.Title
- }
- if strings.TrimSpace(task.Overview) != "" {
- data["overview"] = task.Overview
- }
- if strings.TrimSpace(task.OriginalName) != "" {
- data["original_title"] = task.OriginalName
- }
- if strings.TrimSpace(task.OriginalLanguage) != "" {
- data["original_language"] = task.OriginalLanguage
- }
- if task.Year > 0 {
- data["year"] = task.Year
- }
- if task.Rating > 0 {
- data["rating"] = task.Rating
- }
- if strings.TrimSpace(task.Genres) != "" {
- data["genres"] = task.Genres
- }
- }
- 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 (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 := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey); ok {
- return &task, true
- }
- if strings.TrimSpace(torrent.ContentPath) != "" {
- if task, ok := findMatchingTaskByTorrentIdentity(filepath.Base(torrent.ContentPath), 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 {
- // 常见路径映射:qBittorrent容器路径 -> MediaStationGo容器路径
- 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
- }
- for _, candidate := range []string{
- torrent.ContentPath,
- filepath.Join(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
-}
diff --git a/internal/service/downloads_test.go b/internal/service/downloads_test.go
index 9529121..bcef269 100644
--- a/internal/service/downloads_test.go
+++ b/internal/service/downloads_test.go
@@ -2,19 +2,13 @@ package service
import (
"encoding/json"
- "errors"
- "net/http"
- "net/http/httptest"
"os"
"path/filepath"
"strings"
- "sync/atomic"
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -48,14 +42,84 @@ func TestDownloadViewsDoNotExposePrivateURL(t *testing.T) {
}
}
+func TestLiveTorrentSnapshotUsesPollingSnapshot(t *testing.T) {
+ now := time.Date(2026, 6, 23, 0, 0, 0, 0, time.UTC)
+ svc := NewDownloadService(zap.NewNop(), nil, NewHub(zap.NewNop()), nil)
+ svc.now = func() time.Time { return now }
+
+ live := []QBitTorrent{{
+ Hash: "hash-1",
+ Name: "Release.Name.S01E01",
+ State: "downloading",
+ Progress: 0.5,
+ }}
+ svc.processDownloadSnapshot(t.Context(), live, nil)
+ live[0].Name = "mutated"
+
+ got := svc.LiveTorrentSnapshot(30 * time.Second)
+ if len(got) != 1 || got[0].Name != "Release.Name.S01E01" {
+ t.Fatalf("snapshot = %#v, want cloned live torrent", got)
+ }
+ got[0].Name = "changed"
+ again := svc.LiveTorrentSnapshot(30 * time.Second)
+ if len(again) != 1 || again[0].Name != "Release.Name.S01E01" {
+ t.Fatalf("snapshot was mutable through caller: %#v", again)
+ }
+
+ now = now.Add(31 * time.Second)
+ if stale := svc.LiveTorrentSnapshot(30 * time.Second); len(stale) != 0 {
+ t.Fatalf("stale snapshot = %#v, want empty", stale)
+ }
+}
+
+func TestDownloadCompleteNotificationPayloadUsesTaskMetadata(t *testing.T) {
+ body, data := downloadCompleteNotificationPayload(QBitTorrent{
+ Hash: "done123",
+ Name: "Release.Name.S01E02.1080p",
+ SavePath: "/downloads/show",
+ ContentPath: "/downloads/show/Release.Name.S01E02.1080p.mkv",
+ }, &model.DownloadTask{
+ Title: "正式标题",
+ PosterURL: "https://img.example/poster.jpg",
+ BackdropURL: "https://img.example/backdrop.jpg",
+ MediaType: "tv",
+ MediaCategory: "日番",
+ Overview: "简介",
+ OriginalName: "Original Title",
+ OriginalLanguage: "ja",
+ Year: 2026,
+ Rating: 8.7,
+ Genres: "动画,剧情",
+ })
+
+ if !strings.Contains(body, "任务:正式标题") {
+ t.Fatalf("body should prefer task title, got %q", body)
+ }
+ if !strings.Contains(body, "保存路径:/downloads/show/Release.Name.S01E02.1080p.mkv") {
+ t.Fatalf("body should include content path, got %q", body)
+ }
+ for key, want := range map[string]interface{}{
+ "resource_title": "Release.Name.S01E02.1080p",
+ "title": "正式标题",
+ "poster_url": "https://img.example/poster.jpg",
+ "backdrop_url": "https://img.example/backdrop.jpg",
+ "media_type": "tv",
+ "media_category": "日番",
+ "overview": "简介",
+ "original_title": "Original Title",
+ "original_language": "ja",
+ "year": 2026,
+ "rating": float32(8.7),
+ "genres": "动画,剧情",
+ } {
+ if got := data[key]; got != want {
+ t.Fatalf("data[%s] = %#v, want %#v", key, got, want)
+ }
+ }
+}
+
func TestSyncDownloadTaskProgressSkipsUnchangedCompletedTask(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{})
repos := repository.New(db)
task := &model.DownloadTask{
Source: "qbittorrent",
@@ -89,13 +153,7 @@ func TestSyncDownloadTaskProgressSkipsUnchangedCompletedTask(t *testing.T) {
}
func TestSyncDownloadTaskProgressMatchesSeasonFolderTorrentName(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{})
repos := repository.New(db)
task := &model.DownloadTask{
Source: "qbittorrent",
@@ -126,13 +184,7 @@ func TestSyncDownloadTaskProgressMatchesSeasonFolderTorrentName(t *testing.T) {
}
func TestProcessDownloadSnapshotQueuesCompletedPendingTaskOnFirstSnapshot(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
task := &model.DownloadTask{
Source: "qbittorrent",
@@ -163,13 +215,7 @@ func TestProcessDownloadSnapshotQueuesCompletedPendingTaskOnFirstSnapshot(t *tes
}
func TestProcessDownloadSnapshotSkipsUntrackedCompletedTorrentOnFirstSnapshot(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{}, &model.Setting{})
repos := repository.New(db)
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
@@ -516,6 +562,43 @@ func TestDownloadPollSkipsRecordedCompletedTorrentCatchup(t *testing.T) {
}
}
+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")
@@ -602,779 +685,3 @@ func TestUserPathMappingsParsing(t *testing.T) {
}
}
}
-
-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 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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)
- }
-}
-
-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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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)
- }
-}
-
-func TestAddDownloadWithMetaSkipsExistingLocalMovieBeforeQBAdd(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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)
- }
-}
-
-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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 TestAddDownloadWithMetaFailsClosedWhenNoDownloaderConfigured(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 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 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- 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 TestAddDownloadWithMetaSkipsExistingLocalEpisodeWithReleaseGroup(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}, &model.DownloadTask{}, &model.Setting{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- 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))
- }
-}
diff --git a/internal/service/emby_artwork.go b/internal/service/emby_artwork.go
new file mode 100644
index 0000000..997d380
--- /dev/null
+++ b/internal/service/emby_artwork.go
@@ -0,0 +1,75 @@
+package service
+
+import (
+ "context"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// ImageURL returns artwork for a media/series/season item id.
+func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (string, error) {
+ pick := func(primary, backdrop string) string {
+ switch strings.ToLower(imageType) {
+ case "backdrop", "art":
+ if backdrop != "" {
+ return backdrop
+ }
+ }
+ if primary != "" {
+ return primary
+ }
+ return backdrop
+ }
+ if strings.HasPrefix(id, embyVirtualSeasonPrefix) {
+ if raw, ok := e.cachedArtworkURL(id, imageType); ok {
+ return raw, nil
+ }
+ return "", nil
+ }
+ if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
+ if raw, ok := e.cachedArtworkURL(id, imageType); ok {
+ return raw, nil
+ }
+ return "", nil
+ }
+ m, err := e.repo.Media.FindByID(ctx, id)
+ if err == nil && m != nil {
+ if e.mediaShouldBeEpisode(ctx, m) {
+ switch strings.ToLower(imageType) {
+ case "backdrop", "art":
+ return "", nil
+ }
+ }
+ return pick(e.mediaPrimaryArtwork(ctx, m), e.mediaBackdropArtwork(ctx, m)), nil
+ }
+ if err != nil {
+ return "", err
+ }
+ if series, ok, err := e.findSeriesGroup(ctx, id, ""); err != nil {
+ return "", err
+ } else if ok {
+ return pick(series.PosterURL, series.BackdropURL), nil
+ }
+ return "", nil
+}
+
+func (e *EmbyService) mediaPrimaryArtwork(ctx context.Context, m *model.Media) string {
+ if m == nil {
+ return ""
+ }
+ if e.mediaShouldBeEpisode(ctx, m) && strings.TrimSpace(m.BackdropURL) != "" {
+ return m.BackdropURL
+ }
+ return m.PosterURL
+}
+
+func (e *EmbyService) mediaBackdropArtwork(ctx context.Context, m *model.Media) string {
+ if m == nil {
+ return ""
+ }
+ if e.mediaShouldBeEpisode(ctx, m) {
+ return ""
+ }
+ return m.BackdropURL
+}
diff --git a/internal/service/emby_compat.go b/internal/service/emby_compat.go
index 158c6e7..23dc8ea 100644
--- a/internal/service/emby_compat.go
+++ b/internal/service/emby_compat.go
@@ -13,26 +13,14 @@ package service
import (
"context"
- "crypto/sha256"
- "encoding/hex"
- "errors"
- "fmt"
- "net/url"
- "path/filepath"
"regexp"
- "sort"
- "strconv"
- "strings"
"sync"
"time"
- "go.uber.org/zap"
- "gorm.io/gorm"
-
"github.com/ShukeBta/MediaStationGo/internal/config"
- "github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+ "go.uber.org/zap"
)
// 用一个固定的 ServerId 字符串。Emby 客户端会缓存这个 id,第一次见到
@@ -103,194 +91,6 @@ func (e *EmbyService) SetCloudProbe(storage cloudPlaybackResolver, probe cloudPl
e.probe = probe
}
-// ─── System ──────────────────────────────────────────────────────────────────
-
-// SystemInfo returns the full Emby identity payload.
-func (e *EmbyService) SystemInfo() map[string]any {
- return map[string]any{
- "Id": embyServerID,
- "ServerId": embyServerID,
- "ServerName": "MediaStationGo",
- "Version": embyCompatVersion,
- "ServerVersion": embyCompatVersion,
- "ProductName": "Emby Server",
- "OperatingSystem": "Windows",
- "Architecture": "X64",
- "LocalAddress": "",
- "WanAddress": "",
- "HasPendingRestart": false,
- "IsShuttingDown": false,
- "SupportsLibraryMonitor": true,
- "SupportsHttps": false,
- "SupportsAutoDiscovery": true,
- "HttpServerPortNumber": e.cfg.App.Port,
- "HttpsPortNumber": 0,
- "PublishedServerUrl": "",
- "WebSocketPortNumber": e.cfg.App.Port,
- "CompletedInstallations": []any{},
- "CanSelfRestart": false,
- "CanLaunchWebBrowser": false,
- "CanRestart": false,
- }
-}
-
-// SystemInfoPublic 是不需要认证的精简版(Emby Web 客户端登陆前会拉)。
-func (e *EmbyService) SystemInfoPublic() map[string]any {
- return map[string]any{
- "Id": embyServerID,
- "ServerId": embyServerID,
- "ServerName": "MediaStationGo",
- "Version": embyCompatVersion,
- "ServerVersion": embyCompatVersion,
- "ProductName": "Emby Server",
- "OperatingSystem": "Windows",
- "LocalAddress": "",
- "WanAddress": "",
- "HttpServerPortNumber": e.cfg.App.Port,
- "HttpsPortNumber": 0,
- "SupportsHttps": false,
- "SupportsAutoDiscovery": true,
- "StartupWizardCompleted": true,
- }
-}
-
-// ─── Users ───────────────────────────────────────────────────────────────────
-
-// ListUsers returns Emby-shaped users.
-func (e *EmbyService) ListUsers(ctx context.Context) ([]map[string]any, error) {
- users, err := e.repo.User.List(ctx)
- if err != nil {
- return nil, err
- }
- out := make([]map[string]any, 0, len(users))
- for _, u := range users {
- out = append(out, e.userPayload(&u))
- }
- return out, nil
-}
-
-// FindUser 用 ID 查用户,用于 /Users/Me 与 /Users/{id}。
-func (e *EmbyService) FindUser(ctx context.Context, id string) (map[string]any, error) {
- u, err := e.repo.User.FindByID(ctx, id)
- if err != nil || u == nil {
- return nil, err
- }
- return e.userPayload(u), nil
-}
-
-func (e *EmbyService) userPayload(u *model.User) map[string]any {
- canDownload := u.Role == "admin"
- return map[string]any{
- "Id": u.ID,
- "Name": u.Username,
- "ServerId": embyServerID,
- "ServerName": "MediaStationGo",
- "HasPassword": true,
- "HasConfiguredPassword": true,
- "HasConfiguredEasyPassword": false,
- "EnableAutoLogin": false,
- "LastLoginDate": u.LastLoginAt,
- "LastActivityDate": u.UpdatedAt,
- "Configuration": map[string]any{
- "PlayDefaultAudioTrack": true,
- "DisplayCollectionsView": true,
- "DisplayMissingEpisodes": false,
- "SubtitleMode": "Default",
- "EnableNextEpisodeAutoPlay": true,
- "AudioLanguagePreference": "",
- "SubtitleLanguagePreference": "",
- },
- "Policy": map[string]any{
- "IsAdministrator": u.Role == "admin",
- "IsHidden": false,
- "IsDisabled": !u.IsActive,
- "EnableUserPreferenceAccess": true,
- "EnableRemoteAccess": true,
- "EnableMediaPlayback": true,
- "EnableAudioPlaybackTranscoding": true,
- "EnableVideoPlaybackTranscoding": true,
- "EnablePlaybackRemuxing": true,
- "EnableLiveTvAccess": false,
- "EnableContentDownloading": canDownload,
- "EnableSyncTranscoding": canDownload,
- "EnableMediaConversion": canDownload,
- "EnableAllChannels": true,
- "EnableAllFolders": true,
- "EnableAllDevices": true,
- "AuthenticationProviderId": embyLocalAuthenticationProviderID,
- "PasswordResetProviderId": embyLocalPasswordResetProviderID,
- },
- }
-}
-
-// ─── Views / MediaFolders ────────────────────────────────────────────────────
-
-// Views 返回 Emby 中"虚拟根目录"——每个 library 一个条目。
-func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any, error) {
- libs, err := e.repo.Library.List(ctx)
- if err != nil {
- return nil, err
- }
- libs = FilterDisplayCloudLibraries(ctx, e.repo, libs)
- visibility := e.mediaVisibility(ctx, userID)
- items := make([]map[string]any, 0, len(libs))
- for _, l := range libs {
- if !e.libraryVisibleFromCachedVisibility(l, visibility) {
- continue
- }
- items = append(items, e.libraryAsView(&l))
- }
- return map[string]any{"Items": items, "TotalRecordCount": len(items), "StartIndex": 0}, nil
-}
-
-func (e *EmbyService) libraryAsView(l *model.Library) map[string]any {
- collectionType := "movies"
- switch l.Type {
- case "tv":
- collectionType = "tvshows"
- case "anime":
- collectionType = "tvshows" // Emby 没有专门的 anime CollectionType
- case "variety":
- collectionType = "tvshows"
- case "music":
- collectionType = "music"
- }
- return map[string]any{
- "Id": l.ID,
- "Name": l.Name,
- "CollectionType": collectionType,
- "ServerId": embyServerID,
- "Type": "CollectionFolder",
- "IsFolder": true,
- "Path": l.Path,
- "SortName": strings.ToLower(l.Name),
- "DateCreated": l.CreatedAt.UTC().Format(time.RFC3339),
- "CanDelete": false,
- "CanDownload": false,
- "DisplayPreferencesId": l.ID,
- "PrimaryImageItemId": l.ID,
- "PrimaryImageAspectRatio": 1.7777777777777777,
- "RecursiveItemCount": 0,
- "ChildCount": 0,
- "SpecialFeatureCount": 0,
- "EnableMediaSourceDisplay": true,
- "PlayAccess": "Full",
- "ExternalUrls": []any{},
- "ProviderIds": map[string]string{},
- "Genres": []string{},
- "Tags": []string{},
- "ImageTags": map[string]string{},
- "BackdropImageTags": []string{},
- "UserData": map[string]any{
- "PlaybackPositionTicks": 0,
- "PlayCount": 0,
- "IsFavorite": false,
- "Played": false,
- "UnplayedItemCount": 0,
- },
- }
-}
-
// ─── Items ───────────────────────────────────────────────────────────────────
// ItemsParams 是 /Items 与 /Users/{uid}/Items 共用的查询参数。
@@ -322,47 +122,6 @@ var (
embyEpisodeTitleRE = regexp.MustCompile(`(?i)\s*[-_ ]*s\d{1,2}e\d{1,3}.*$`)
)
-type embySeriesGroup struct {
- ID string
- LibraryID string
- Name string
- PosterURL string
- BackdropURL string
- Overview string
- Rating float32
- Year int
- TMDbID int
- BangumiID int
- CreatedAt time.Time
- Episodes []model.Media
-}
-
-type embySeasonGroup struct {
- ID string
- SeriesID string
- LibraryID string
- Name string
- SeasonNum int
- Series embySeriesGroup
- Episodes []model.Media
-}
-
-type embySeriesCacheEntry struct {
- group embySeriesGroup
- expiresAt time.Time
-}
-
-type embySeasonCacheEntry struct {
- season embySeasonGroup
- expiresAt time.Time
-}
-
-type embyArtworkCacheEntry struct {
- primary string
- backdrop string
- expiresAt time.Time
-}
-
type embyVisibilityCacheEntry struct {
visibility MediaVisibility
expiresAt time.Time
@@ -459,2047 +218,3 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any,
}
return e.mediaItems(ctx, p)
}
-
-// movieLibraryHasEpisodicContent 报告电影类型库里是否混入了「剧集结构」内容
-// (有季集号且路径形如剧集,例如 .../国产剧/某剧/Season 01/某剧 - S01E01.mkv)。
-// 用于决定是否需要走 movieLibraryItems 把这些内容聚成 Series 卡片。普通电影库
-// 没有这类行时返回 false,继续走常规 mediaItems。
-func (e *EmbyService) movieLibraryHasEpisodicContent(ctx context.Context, libraryID string) (bool, error) {
- clause, args := embyLikelyEpisodicPathSQL()
- if clause == "" {
- return false, nil
- }
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("library_id IN ?", e.mergedLibraryIDs(ctx, libraryID)).
- Where("(season_num > 0 OR episode_num > 0) AND ("+clause+")", args...)
- var count int64
- if err := q.Limit(1).Count(&count).Error; err != nil {
- return false, err
- }
- return count > 0, nil
-}
-
-// movieLibraryItems 处理电影类型库的常规浏览,返回「真正的电影(Movie)」与
-// 「库内剧集结构内容聚成的 Series 卡片」的合并列表(按 DateCreated 倒序分页)。
-// 与 mediaItems 的区别: 后者会把剧集结构行当散装 Episode 漏出;这里改为聚合成
-// Series,从根本上消除「电影库里整部剧被拆成单集」的现象。
-func (e *EmbyService) movieLibraryItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
- libIDs := e.mergedLibraryIDs(ctx, p.ParentID)
- apply := func(q *gorm.DB) *gorm.DB {
- q = e.applyUserMediaVisibility(ctx, q, p.UserID)
- q = q.Where("library_id IN ?", libIDs)
- if p.SearchTerm != "" {
- q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
- }
- if containsEmbyFilter(p.Filters, "IsFavorite") {
- if strings.TrimSpace(p.UserID) == "" {
- return nil
- }
- q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
- }
- return q
- }
-
- // 剧集结构内容 → Series 卡片。
- clause, args := embyLikelyEpisodicPathSQL()
- var episodicRows []model.Media
- if clause != "" {
- epQ := apply(e.repo.DB.WithContext(ctx).Model(&model.Media{}))
- if epQ == nil {
- return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
- }
- epQ = epQ.Where("(season_num > 0 OR episode_num > 0) AND ("+clause+")", args...).
- Order("media.created_at desc").Limit(embySeriesGroupingLimit)
- if err := epQ.Find(&episodicRows).Error; err != nil {
- return nil, err
- }
- }
- seriesGroups := e.seriesGroupsFromMedia(episodicRows)
-
- // 真正的电影 → Movie 项(剔除剧集结构行)。
- movieQ := apply(e.repo.DB.WithContext(ctx).Model(&model.Media{}))
- if movieQ == nil {
- return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
- }
- movieQ = filterLikelyEpisodicPathsFromMovieQuery(movieQ).
- Order("media.created_at desc").Limit(embySeriesGroupingLimit)
- var movieRows []model.Media
- if err := movieQ.Find(&movieRows).Error; err != nil {
- return nil, err
- }
- movieItems, err := e.payloadsForMedia(ctx, movieRows, p.UserID)
- if err != nil {
- return nil, err
- }
-
- // 合并: Series 卡片 + Movie 项, 统一按 DateCreated 倒序。
- type entry struct {
- createdAt time.Time
- payload map[string]any
- }
- entries := make([]entry, 0, len(seriesGroups)+len(movieItems))
- for _, g := range seriesGroups {
- entries = append(entries, entry{createdAt: g.CreatedAt, payload: e.seriesPayload(g)})
- }
- for _, item := range movieItems {
- entries = append(entries, entry{createdAt: embyPayloadCreatedAt(item), payload: item})
- }
- sort.SliceStable(entries, func(i, j int) bool {
- return entries[i].createdAt.After(entries[j].createdAt)
- })
- total := len(entries)
- paged := pageSlice(entries, p.StartIndex, p.Limit)
- items := make([]map[string]any, 0, len(paged))
- for _, en := range paged {
- items = append(items, en.payload)
- }
- return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
-}
-
-// embyPayloadCreatedAt 从 item payload 里取 DateCreated(time.Time),用于合并排序。
-func embyPayloadCreatedAt(item map[string]any) time.Time {
- if item == nil {
- return time.Time{}
- }
- if v, ok := item["DateCreated"].(time.Time); ok {
- return v
- }
- return time.Time{}
-}
-
-func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
- cacheKey := e.embyItemsCacheKey("items", p)
- var cached embyItemsCacheValue
- if e.cache != nil && e.cache.GetJSON(ctx, cacheKey, &cached) {
- return map[string]any{"Items": cached.Items, "TotalRecordCount": cached.TotalRecordCount, "StartIndex": cached.StartIndex}, nil
- }
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{})
- q = e.applyUserMediaVisibility(ctx, q, p.UserID)
- if p.ParentID != "" {
- q = q.Where("library_id IN ? OR series_id = ?", e.mergedLibraryIDs(ctx, p.ParentID), p.ParentID)
- }
- if p.SearchTerm != "" {
- q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
- }
- if containsEmbyFilter(p.Filters, "IsFavorite") {
- if strings.TrimSpace(p.UserID) == "" {
- return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
- }
- q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
- }
- resumeFilter := containsEmbyFilter(p.Filters, "IsResumable")
- if resumeFilter {
- if strings.TrimSpace(p.UserID) == "" {
- return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
- }
- q = q.Joins(`JOIN (
- SELECT media_id, MAX(watched_at) AS watched_at
- FROM playback_histories
- WHERE user_id = ? AND completed = ? AND position_ms > 0
- GROUP BY media_id
- ) AS resume ON resume.media_id = media.id`, p.UserID, false)
- }
- filterBySeasonNumbers := true
- parentKnownNonEpisodic := false
- if p.ParentID != "" {
- if episodic, err := e.libraryIsEpisodic(ctx, p.ParentID); err == nil && !episodic {
- filterBySeasonNumbers = false
- parentKnownNonEpisodic = true
- }
- }
- if parentKnownNonEpisodic && containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
- return emptyItemsEnvelope(p.StartIndex), nil
- }
- if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Movie") && !containsItemType(p.IncludeItemTypes, "Episode") {
- q = e.filterMovieItems(ctx, q)
- }
- if parentKnownNonEpisodic && containsItemType(p.IncludeItemTypes, "Movie") && !containsItemType(p.IncludeItemTypes, "Episode") {
- q = filterLikelyEpisodicPathsFromMovieQuery(q)
- }
- if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
- q = e.filterEpisodeItems(ctx, q)
- }
-
- var total int64
- if err := q.Count(&total).Error; err != nil {
- return nil, err
- }
- order := "media.created_at desc"
- switch primarySupportedEmbySort(p.SortBy, resumeFilter) {
- case "sortname", "name":
- order = "media.title"
- case "premieredate", "productionyear":
- order = "media.year"
- case "datecreated":
- order = "media.created_at"
- case "dateplayed":
- order = "resume.watched_at"
- case "communityrating":
- order = "media.rating"
- }
- if strings.EqualFold(firstCSVValue(p.SortOrder), "Descending") {
- if !strings.HasSuffix(order, " desc") {
- order = order + " desc"
- }
- }
-
- fetchLimit := p.Limit
- fetchOffset := p.StartIndex
- if fetchLimit > 0 && e.shouldCollapseMediaVersions(ctx, p) {
- // Duplicates across merged local/cloud libraries collapse into one Emby
- // item with multiple MediaSources. Fetch a wider window so duplicates do
- // not consume the whole requested page.
- fetchOffset = 0
- fetchLimit = p.StartIndex + maxInt(p.Limit*4, p.Limit)
- }
- var rows []model.Media
- if err := q.Order(order).Offset(fetchOffset).Limit(fetchLimit).Find(&rows).Error; err != nil {
- return nil, err
- }
- if e.shouldCollapseMediaVersions(ctx, p) {
- rows = e.collapseMediaVersionRows(ctx, rows)
- rows = pageSlice(rows, p.StartIndex, p.Limit)
- }
- items, err := e.payloadsForMedia(ctx, rows, p.UserID)
- if err != nil {
- return nil, err
- }
- out := map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}
- if e.cache != nil {
- e.cache.SetJSON(ctx, cacheKey, embyItemsCacheValue{Items: items, TotalRecordCount: total, StartIndex: p.StartIndex}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
- }
- return out, nil
-}
-
-type embyItemsCacheValue struct {
- Items []map[string]any `json:"items"`
- TotalRecordCount int64 `json:"total_record_count"`
- StartIndex int `json:"start_index"`
-}
-
-type embyLatestCacheValue struct {
- Items []map[string]any `json:"items"`
-}
-
-func (e *EmbyService) embyItemsCacheKey(kind string, p ItemsParams) string {
- includeTypes := append([]string(nil), p.IncludeItemTypes...)
- filters := append([]string(nil), p.Filters...)
- ids := append([]string(nil), p.IDs...)
- sort.Strings(includeTypes)
- sort.Strings(filters)
- sort.Strings(ids)
- sum := sha256.Sum256([]byte(strings.Join([]string{
- kind,
- p.UserID,
- p.ParentID,
- strings.Join(ids, ","),
- p.SearchTerm,
- strings.Join(includeTypes, ","),
- strings.Join(filters, ","),
- strconv.FormatBool(p.Recursive),
- p.SortBy,
- p.SortOrder,
- strconv.Itoa(p.StartIndex),
- strconv.Itoa(p.Limit),
- }, "|")))
- return "media:emby:" + hex.EncodeToString(sum[:])
-}
-
-func (e *EmbyService) embyLatestCacheKey(userID, parentID string, limit int) string {
- sum := sha256.Sum256([]byte(strings.Join([]string{"latest", userID, parentID, strconv.Itoa(limit)}, "|")))
- return "media:emby:" + hex.EncodeToString(sum[:])
-}
-
-func (e *EmbyService) mediaCacheTTLSeconds() int {
- if e == nil || e.cfg == nil || e.cfg.Cache.MediaTTLSeconds < 1 {
- return 15
- }
- return e.cfg.Cache.MediaTTLSeconds
-}
-
-func (e *EmbyService) episodeItems(ctx context.Context, rows []model.Media, p ItemsParams) (map[string]any, error) {
- rows = e.filterMediaRowsForUser(ctx, rows, p.UserID)
- if p.SearchTerm != "" {
- filtered := rows[:0]
- needle := strings.ToLower(p.SearchTerm)
- for _, row := range rows {
- if strings.Contains(strings.ToLower(row.Title), needle) || strings.Contains(strings.ToLower(row.OriginalName), needle) {
- filtered = append(filtered, row)
- }
- }
- rows = filtered
- }
- sort.SliceStable(rows, func(i, j int) bool {
- if rows[i].SeasonNum != rows[j].SeasonNum {
- return rows[i].SeasonNum < rows[j].SeasonNum
- }
- if rows[i].EpisodeNum != rows[j].EpisodeNum {
- return rows[i].EpisodeNum < rows[j].EpisodeNum
- }
- return rows[i].CreatedAt.Before(rows[j].CreatedAt)
- })
- total := len(rows)
- items, err := e.payloadsForMedia(ctx, pageSlice(rows, p.StartIndex, p.Limit), p.UserID)
- if err != nil {
- return nil, err
- }
- return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
-}
-
-func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media, userID string) ([]map[string]any, error) {
- rows = e.collapseMediaVersionRows(ctx, rows)
- userFavs := map[string]bool{}
- userPos := map[string]int64{}
- if userID != "" && len(rows) > 0 {
- mediaIDs := make([]string, 0, len(rows))
- for _, row := range rows {
- if strings.TrimSpace(row.ID) != "" {
- mediaIDs = append(mediaIDs, row.ID)
- }
- }
- if len(mediaIDs) == 0 {
- mediaIDs = []string{"__none__"}
- }
- var favs []model.Favorite
- favQuery := e.repo.DB.WithContext(ctx).Where("user_id = ?", userID).Where("media_id IN ?", mediaIDs)
- _ = favQuery.Find(&favs).Error
- for _, f := range favs {
- userFavs[f.MediaID] = true
- }
- var hist []model.PlaybackHistory
- histQuery := e.repo.DB.WithContext(ctx).Where("user_id = ?", userID).Where("media_id IN ?", mediaIDs)
- _ = histQuery.Find(&hist).Error
- for _, h := range hist {
- userPos[h.MediaID] = h.PositionMs
- }
- }
-
- items := make([]map[string]any, 0, len(rows))
- for _, m := range rows {
- items = append(items, e.itemPayload(ctx, &m, userFavs[m.ID], userPos[m.ID]))
- }
- return items, nil
-}
-
-func (e *EmbyService) shouldCollapseMediaVersions(ctx context.Context, p ItemsParams) bool {
- if containsItemType(p.IncludeItemTypes, "Series") || containsItemType(p.IncludeItemTypes, "Season") {
- return false
- }
- if containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
- return true
- }
- if p.ParentID == "" {
- return true
- }
- episodic, err := e.libraryIsEpisodic(ctx, p.ParentID)
- return err == nil && !episodic
-}
-
-func (e *EmbyService) collapseMediaVersionRows(ctx context.Context, rows []model.Media) []model.Media {
- if len(rows) < 2 {
- return rows
- }
- out := make([]model.Media, 0, len(rows))
- indexByKey := make(map[string]int, len(rows))
- for _, row := range rows {
- key := e.mediaVersionKey(ctx, &row)
- if key == "" {
- out = append(out, row)
- continue
- }
- if idx, ok := indexByKey[key]; ok {
- if preferMediaVersion(row, out[idx]) {
- out[idx] = row
- }
- continue
- }
- indexByKey[key] = len(out)
- out = append(out, row)
- }
- return out
-}
-
-// Item 单条目详情。
-func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[string]any, error) {
- if lib, err := e.repo.Library.FindByID(ctx, mediaID); err != nil {
- return nil, err
- } else if lib != nil {
- libs := FilterDisplayCloudLibraries(ctx, e.repo, []model.Library{*lib})
- if len(libs) == 0 {
- return nil, nil
- }
- visibility := e.mediaVisibility(ctx, userID)
- if !e.libraryVisibleFromCachedVisibility(libs[0], visibility) {
- return nil, nil
- }
- return e.libraryAsView(&libs[0]), nil
- }
- if strings.HasPrefix(mediaID, embyVirtualSeasonPrefix) {
- if season, ok, err := e.findSeasonGroup(ctx, mediaID, userID); err != nil {
- return nil, err
- } else if ok {
- return e.seasonPayload(season), nil
- }
- }
- if strings.HasPrefix(mediaID, embyVirtualSeriesPrefix) {
- if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
- return nil, err
- } else if ok {
- return e.seriesPayload(series), nil
- }
- }
- m, err := e.repo.Media.FindByID(ctx, mediaID)
- if err != nil {
- return nil, err
- }
- if m == nil {
- if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
- return nil, err
- } else if ok {
- return e.seriesPayload(series), nil
- }
- return nil, nil
- }
- if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
- return nil, nil
- }
- fav := false
- pos := int64(0)
- if userID != "" {
- var f model.Favorite
- ferr := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
- if ferr == nil {
- fav = true
- }
- var h model.PlaybackHistory
- herr := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).
- Order("watched_at desc").First(&h).Error
- if herr == nil {
- pos = h.PositionMs
- }
- }
- return e.itemPayload(ctx, m, fav, pos), nil
-}
-
-// LatestItems 最近添加,全库或指定库。
-func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string, limit int) ([]map[string]any, error) {
- if limit <= 0 || limit > 100 {
- limit = 20
- }
- cacheKey := e.embyLatestCacheKey(userID, parentID, limit)
- var cached embyLatestCacheValue
- if e.cache != nil && e.cache.GetJSON(ctx, cacheKey, &cached) {
- return cached.Items, nil
- }
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("deleted_at IS NULL")
- q = e.applyUserMediaVisibility(ctx, q, userID)
- if parentID != "" {
- if episodic, err := e.libraryIsEpisodic(ctx, parentID); err == nil && episodic {
- out, err := e.latestSeriesItemsForLibrary(ctx, userID, parentID, limit)
- if err == nil && e.cache != nil {
- e.cache.SetJSON(ctx, cacheKey, embyLatestCacheValue{Items: out}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
- }
- return out, err
- }
- q = q.Where("library_id IN ?", e.mergedLibraryIDs(ctx, parentID))
- }
- var rows []model.Media
- if err := q.Order("media.created_at desc").Limit(limit).Find(&rows).Error; err != nil {
- return nil, err
- }
- favs := map[string]bool{}
- if userID != "" && len(rows) > 0 {
- mediaIDs := make([]string, 0, len(rows))
- for _, row := range rows {
- if strings.TrimSpace(row.ID) != "" {
- mediaIDs = append(mediaIDs, row.ID)
- }
- }
- if len(mediaIDs) == 0 {
- mediaIDs = []string{"__none__"}
- }
- var fr []model.Favorite
- _ = e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id IN ?", userID, mediaIDs).Find(&fr).Error
- for _, f := range fr {
- favs[f.MediaID] = true
- }
- }
- out := make([]map[string]any, 0, len(rows))
- for _, m := range rows {
- out = append(out, e.itemPayload(ctx, &m, favs[m.ID], 0))
- }
- if e.cache != nil {
- e.cache.SetJSON(ctx, cacheKey, embyLatestCacheValue{Items: out}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
- }
- return out, nil
-}
-
-func (e *EmbyService) latestSeriesItemsForLibrary(ctx context.Context, userID, libraryID string, limit int) ([]map[string]any, error) {
- if limit <= 0 || limit > 100 {
- limit = 20
- }
- rowLimit := limit * 40
- if rowLimit < 200 {
- rowLimit = 200
- }
- if rowLimit > embySeriesGroupingLimit {
- rowLimit = embySeriesGroupingLimit
- }
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("library_id IN ? AND (season_num > 0 OR episode_num > 0)", e.mergedLibraryIDs(ctx, libraryID))
- q = e.applyUserMediaVisibility(ctx, q, userID)
- var rows []model.Media
- if err := q.Order("media.created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
- return nil, err
- }
- groups := e.seriesGroupsFromMedia(rows)
- sortSeriesGroups(groups, ItemsParams{SortBy: "datecreated", SortOrder: "Descending"})
- if len(groups) > limit {
- groups = groups[:limit]
- }
- items := make([]map[string]any, 0, len(groups))
- for _, group := range groups {
- items = append(items, e.seriesPayload(group))
- }
- return items, nil
-}
-
-// ResumeItems 列出有未完成播放进度的媒体。
-func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int) (map[string]any, error) {
- if limit <= 0 || limit > 100 {
- limit = 20
- }
- type row struct {
- MediaID string
- PositionMs int64
- DurationMs int64
- }
- var hist []model.PlaybackHistory
- if err := e.repo.DB.WithContext(ctx).
- Where("user_id = ? AND completed = ? AND position_ms > 0", userID, false).
- Order("watched_at desc").Limit(limit).Find(&hist).Error; err != nil {
- return nil, err
- }
- if len(hist) == 0 {
- return map[string]any{"Items": []any{}, "TotalRecordCount": 0}, nil
- }
- ids := make([]string, 0, len(hist))
- posByID := map[string]int64{}
- for _, h := range hist {
- ids = append(ids, h.MediaID)
- posByID[h.MediaID] = h.PositionMs
- }
- var medias []model.Media
- q := e.repo.DB.WithContext(ctx).Where("id IN ?", ids)
- q = e.applyUserMediaVisibility(ctx, q, userID)
- if err := q.Find(&medias).Error; err != nil {
- return nil, err
- }
- // 维持时间倒序
- byID := map[string]*model.Media{}
- for i := range medias {
- byID[medias[i].ID] = &medias[i]
- }
- items := make([]map[string]any, 0, len(hist))
- for _, h := range hist {
- if m, ok := byID[h.MediaID]; ok {
- items = append(items, e.itemPayload(ctx, m, false, posByID[h.MediaID]))
- }
- }
- return map[string]any{"Items": items, "TotalRecordCount": len(items)}, nil
-}
-
-func (e *EmbyService) itemPayload(ctx context.Context, m *model.Media, fav bool, posMs int64) map[string]any {
- itemType := "Movie"
- name := m.Title
- parentID := m.LibraryID
- seriesID := m.SeriesID
- seriesName := ""
- seasonID := ""
- if e.mediaShouldBeEpisode(ctx, m) {
- itemType = "Episode"
- seriesID = e.seriesIDForMedia(m)
- seriesName = e.seriesNameForMedia(m)
- seasonID = e.seasonIDForMedia(m)
- parentID = seasonID
- originalName := strings.TrimSpace(m.OriginalName)
- if originalName != "" && !strings.EqualFold(originalName, seriesName) && !strings.EqualFold(originalName, m.Title) {
- name = m.OriginalName
- } else if m.EpisodeNum > 0 {
- name = fmt.Sprintf("第 %d 集", m.EpisodeNum)
- }
- }
- imageTags := map[string]string{}
- backdropTags := []string{}
- primaryArtwork := e.mediaPrimaryArtwork(ctx, m)
- backdropArtwork := e.mediaBackdropArtwork(ctx, m)
- if primaryArtwork != "" {
- imageTags["Primary"] = m.ID
- }
- if backdropArtwork != "" {
- backdropTags = append(backdropTags, m.ID+"-bd")
- }
-
- runTimeTicks := int64(m.DurationSec) * 10_000_000
- durationMs := int64(m.DurationSec) * 1000
- played := posMs > 0 && durationMs > 0 && posMs >= durationMs*9/10
- pct := 0.0
- if durationMs > 0 {
- pct = float64(posMs) / float64(durationMs) * 100
- }
-
- return map[string]any{
- "Id": m.ID,
- "Name": name,
- "OriginalTitle": m.OriginalName,
- "ServerId": embyServerID,
- "Type": itemType,
- "MediaType": "Video",
- "IsFolder": false,
- "ProductionYear": m.Year,
- "ParentIndexNumber": m.SeasonNum,
- "IndexNumber": m.EpisodeNum,
- "Overview": m.Overview,
- "RunTimeTicks": runTimeTicks,
- "CommunityRating": m.Rating,
- "Container": m.Container,
- "Width": m.Width,
- "Height": m.Height,
- "DateCreated": m.CreatedAt,
- "Path": m.Path,
- "ParentId": parentID,
- "SeasonId": seasonID,
- "SeasonName": seasonName(m.SeasonNum),
- "SeriesId": seriesID,
- "SeriesName": seriesName,
- "ImageTags": imageTags,
- "BackdropImageTags": backdropTags,
- "Genres": splitCSV(m.Genres),
- "ProviderIds": map[string]string{
- "Tmdb": intToStr(m.TMDbID),
- "Bangumi": intToStr(m.BangumiID),
- },
- "UserData": map[string]any{
- "PlaybackPositionTicks": posMs * 10_000,
- "PlayCount": 0,
- "IsFavorite": fav,
- "Played": played,
- "PlayedPercentage": pct,
- },
- "MediaSources": e.mediaSourcesForItem(ctx, m, true, false),
- }
-}
-
-func (e *EmbyService) seriesItemsForLibrary(ctx context.Context, libraryID string, p ItemsParams) (map[string]any, error) {
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
- q = e.applyUserMediaVisibility(ctx, q, p.UserID)
- if libraryID != "" {
- q = q.Where("library_id IN ?", e.mergedLibraryIDs(ctx, libraryID))
- }
- if p.SearchTerm != "" {
- q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
- }
- if containsEmbyFilter(p.Filters, "IsFavorite") {
- if strings.TrimSpace(p.UserID) == "" {
- return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
- }
- q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
- }
- rowLimit := p.StartIndex + maxInt(p.Limit*40, 1000)
- if rowLimit < p.Limit {
- rowLimit = p.Limit
- }
- if rowLimit > embySeriesGroupingLimit {
- rowLimit = embySeriesGroupingLimit
- }
- var rows []model.Media
- if err := q.Order("media.created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
- return nil, err
- }
- groups := e.seriesGroupsFromMedia(rows)
- sortSeriesGroups(groups, p)
- total := len(groups)
- items := make([]map[string]any, 0, minInt(p.Limit, len(groups)))
- for _, group := range pageSlice(groups, p.StartIndex, p.Limit) {
- items = append(items, e.seriesPayload(group))
- }
- return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
-}
-
-func (e *EmbyService) libraryIsEpisodic(ctx context.Context, libraryID string) (bool, error) {
- if strings.TrimSpace(libraryID) == "" {
- return false, nil
- }
- if lib, err := e.repo.Library.FindByID(ctx, libraryID); err != nil {
- return false, err
- } else if lib != nil {
- return embyLibraryTypeIsEpisodic(lib.Type), nil
- }
- var count int64
- err := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("library_id IN ? AND (season_num > 0 OR episode_num > 0)", e.mergedLibraryIDs(ctx, libraryID)).
- Count(&count).Error
- return count > 0, err
-}
-
-func (e *EmbyService) mediaBelongsToEpisodicLibrary(ctx context.Context, m *model.Media) bool {
- if e == nil || m == nil || strings.TrimSpace(m.LibraryID) == "" {
- return false
- }
- lib, err := e.repo.Library.FindByID(ctx, m.LibraryID)
- if err != nil || lib == nil {
- return false
- }
- return embyLibraryTypeIsEpisodic(lib.Type)
-}
-
-func (e *EmbyService) mediaShouldBeEpisode(ctx context.Context, m *model.Media) bool {
- if m == nil || (m.SeasonNum <= 0 && m.EpisodeNum <= 0) {
- return false
- }
- if e.mediaBelongsToEpisodicLibrary(ctx, m) {
- return true
- }
- return embyMediaPathLooksEpisodic(m.Path)
-}
-
-func embyLibraryTypeIsEpisodic(typ string) bool {
- switch strings.ToLower(strings.TrimSpace(typ)) {
- case "tv", "anime", "variety":
- return true
- default:
- return false
- }
-}
-
-func (e *EmbyService) filterMovieItems(ctx context.Context, q *gorm.DB) *gorm.DB {
- episodicIDs := e.episodicLibraryIDs(ctx)
- if len(episodicIDs) == 0 {
- return filterLikelyEpisodicPathsFromMovieQuery(q)
- }
- q = q.Where("(media.season_num = 0 AND media.episode_num = 0) OR media.library_id NOT IN ?", episodicIDs)
- return filterLikelyEpisodicPathsFromMovieQuery(q)
-}
-
-func (e *EmbyService) filterEpisodeItems(ctx context.Context, q *gorm.DB) *gorm.DB {
- episodicIDs := e.episodicLibraryIDs(ctx)
- if len(episodicIDs) == 0 {
- return q.Where("1 = 0")
- }
- return q.Where("media.library_id IN ? AND (media.season_num > 0 OR media.episode_num > 0)", episodicIDs)
-}
-
-func (e *EmbyService) episodicLibraryIDs(ctx context.Context) []string {
- if e == nil || e.repo == nil || e.repo.DB == nil {
- return nil
- }
- var ids []string
- if err := e.repo.DB.WithContext(ctx).Model(&model.Library{}).
- Where("LOWER(type) IN ?", []string{"tv", "anime", "variety"}).
- Pluck("id", &ids).Error; err != nil {
- return nil
- }
- return ids
-}
-
-func filterLikelyEpisodicPathsFromMovieQuery(q *gorm.DB) *gorm.DB {
- clause, args := embyLikelyEpisodicPathSQL()
- if clause == "" {
- return q
- }
- return q.Where("NOT ((media.season_num > 0 OR media.episode_num > 0) AND ("+clause+"))", args...)
-}
-
-func embyLikelyEpisodicPathSQL() (string, []any) {
- patterns := []string{
- "%/season %/%", "%/season.%/%", "%/season-%/%", "%/season_%/%",
- "%/s0%/%", "%/s1%/%", "%/s2%/%", "%/s3%/%", "%/s4%/%", "%/s5%/%", "%/s6%/%", "%/s7%/%", "%/s8%/%", "%/s9%/%",
- "%/special/%", "%/specials/%", "%/sp/%", "%/ova/%", "%/oad/%", "%/extra/%", "%/extras/%",
- "%/电视剧/%", "%/剧集/%", "%/国产剧/%", "%/欧美剧/%", "%/日韩剧/%", "%/日剧/%", "%/韩剧/%",
- "%/日番/%", "%/国漫/%", "%/番剧/%", "%/动漫/%", "%/特别篇/%", "%/特別篇/%", "%/番外/%", "%/特典/%",
- }
- clauses := make([]string, 0, len(patterns)*2)
- args := make([]any, 0, len(patterns)*2)
- for _, pattern := range patterns {
- clauses = append(clauses, "LOWER(media.path) LIKE ?")
- args = append(args, pattern)
- if strings.Contains(pattern, "/") {
- clauses = append(clauses, "LOWER(media.path) LIKE ?")
- args = append(args, strings.ReplaceAll(pattern, "/", `\`))
- }
- }
- return strings.Join(clauses, " OR "), args
-}
-
-func embyMediaPathLooksEpisodic(path string) bool {
- normalized := strings.ToLower(strings.ReplaceAll(strings.TrimSpace(path), "\\", "/"))
- if normalized == "" {
- return false
- }
- for _, marker := range []string{
- "/season ", "/season.", "/season-", "/season_", "/special/", "/specials/", "/sp/", "/ova/", "/oad/", "/extra/", "/extras/",
- "/电视剧/", "/剧集/", "/国产剧/", "/欧美剧/", "/日韩剧/", "/日剧/", "/韩剧/",
- "/日番/", "/国漫/", "/番剧/", "/动漫/", "/特别篇/", "/特別篇/", "/番外/", "/特典/",
- } {
- if strings.Contains(normalized, marker) {
- return true
- }
- }
- for _, marker := range []string{"/s0", "/s1", "/s2", "/s3", "/s4", "/s5", "/s6", "/s7", "/s8", "/s9"} {
- if idx := strings.Index(normalized, marker); idx >= 0 {
- after := idx + len(marker)
- if after < len(normalized) && normalized[after] >= '0' && normalized[after] <= '9' {
- slash := after + 1
- if slash < len(normalized) && normalized[slash] == '/' {
- return true
- }
- }
- }
- }
- return false
-}
-
-func (e *EmbyService) rememberSeriesGroup(group embySeriesGroup) {
- if e == nil || strings.TrimSpace(group.ID) == "" {
- return
- }
- expiresAt := time.Now().Add(embyVirtualCacheTTL)
- e.virtualMu.Lock()
- defer e.virtualMu.Unlock()
- if e.virtualSeries == nil {
- e.virtualSeries = make(map[string]embySeriesCacheEntry)
- }
- if e.virtualSeasons == nil {
- e.virtualSeasons = make(map[string]embySeasonCacheEntry)
- }
- if e.virtualArtwork == nil {
- e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
- }
- if len(e.virtualSeries) > 2000 || len(e.virtualSeasons) > 5000 || len(e.virtualArtwork) > 7000 {
- e.virtualSeries = make(map[string]embySeriesCacheEntry)
- e.virtualSeasons = make(map[string]embySeasonCacheEntry)
- e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
- }
- e.virtualSeries[group.ID] = embySeriesCacheEntry{group: group, expiresAt: expiresAt}
- e.virtualArtwork[group.ID] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
- e.virtualArtwork[group.ID+"-bd"] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
- for _, season := range e.seasonsForSeries(group) {
- e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
- e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
- e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
- }
-}
-
-func (e *EmbyService) rememberSeasonGroup(season embySeasonGroup) {
- if e == nil || strings.TrimSpace(season.ID) == "" {
- return
- }
- expiresAt := time.Now().Add(embyVirtualCacheTTL)
- e.virtualMu.Lock()
- defer e.virtualMu.Unlock()
- if e.virtualSeasons == nil {
- e.virtualSeasons = make(map[string]embySeasonCacheEntry)
- }
- if e.virtualArtwork == nil {
- e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
- }
- e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
- e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
- e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
-}
-
-func (e *EmbyService) cachedSeriesGroup(id string) (embySeriesGroup, bool) {
- if e == nil || strings.TrimSpace(id) == "" {
- return embySeriesGroup{}, false
- }
- now := time.Now()
- e.virtualMu.RLock()
- entry, ok := e.virtualSeries[id]
- e.virtualMu.RUnlock()
- if !ok || now.After(entry.expiresAt) {
- if ok {
- e.virtualMu.Lock()
- delete(e.virtualSeries, id)
- e.virtualMu.Unlock()
- }
- return embySeriesGroup{}, false
- }
- return entry.group, true
-}
-
-func (e *EmbyService) cachedSeasonGroup(id string) (embySeasonGroup, bool) {
- if e == nil || strings.TrimSpace(id) == "" {
- return embySeasonGroup{}, false
- }
- now := time.Now()
- e.virtualMu.RLock()
- entry, ok := e.virtualSeasons[id]
- e.virtualMu.RUnlock()
- if !ok || now.After(entry.expiresAt) {
- if ok {
- e.virtualMu.Lock()
- delete(e.virtualSeasons, id)
- e.virtualMu.Unlock()
- }
- return embySeasonGroup{}, false
- }
- return entry.season, true
-}
-
-func (e *EmbyService) cachedArtworkURL(id, imageType string) (string, bool) {
- if e == nil || strings.TrimSpace(id) == "" {
- return "", false
- }
- now := time.Now()
- e.virtualMu.RLock()
- entry, ok := e.virtualArtwork[id]
- e.virtualMu.RUnlock()
- if !ok || now.After(entry.expiresAt) {
- if ok {
- e.virtualMu.Lock()
- delete(e.virtualArtwork, id)
- e.virtualMu.Unlock()
- }
- return "", false
- }
- switch strings.ToLower(imageType) {
- case "backdrop", "art":
- if entry.backdrop != "" {
- return entry.backdrop, true
- }
- }
- if entry.primary != "" {
- return entry.primary, true
- }
- return entry.backdrop, entry.backdrop != ""
-}
-
-func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (embySeriesGroup, bool, error) {
- if strings.TrimSpace(id) == "" {
- return embySeriesGroup{}, false, nil
- }
- if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
- if group, ok := e.cachedSeriesGroup(id); ok {
- return group, true, nil
- }
- }
- var rows []model.Media
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
- q = e.applyUserMediaVisibility(ctx, q, userID)
- if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
- q = q.Where("series_id = ?", id)
- }
- if err := q.Order("media.season_num asc, media.episode_num asc, media.created_at asc").Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
- return embySeriesGroup{}, false, err
- }
- for _, group := range e.seriesGroupsFromMedia(rows) {
- if group.ID == id {
- e.rememberSeriesGroup(group)
- return group, true, nil
- }
- }
- if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
- if series, err := e.repo.Series.FindByID(ctx, id); err != nil {
- return embySeriesGroup{}, false, err
- } else if series != nil {
- return embySeriesGroup{
- ID: series.ID,
- LibraryID: series.LibraryID,
- Name: series.Title,
- PosterURL: series.PosterURL,
- BackdropURL: series.BackdropURL,
- Overview: series.Overview,
- Rating: series.Rating,
- Year: series.Year,
- TMDbID: series.TMDbID,
- BangumiID: series.BangumiID,
- CreatedAt: series.CreatedAt,
- }, true, nil
- }
- }
- return embySeriesGroup{}, false, nil
-}
-
-func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (embySeasonGroup, bool, error) {
- if strings.TrimSpace(id) == "" || !strings.HasPrefix(id, embyVirtualSeasonPrefix) {
- return embySeasonGroup{}, false, nil
- }
- if season, ok := e.cachedSeasonGroup(id); ok {
- return season, true, nil
- }
- var rows []model.Media
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("season_num > 0 OR episode_num > 0")
- q = e.applyUserMediaVisibility(ctx, q, userID)
- if err := q.
- Order("media.season_num asc, media.episode_num asc, media.created_at asc").
- Limit(embySeriesGroupingLimit).
- Find(&rows).Error; err != nil {
- return embySeasonGroup{}, false, err
- }
- for _, series := range e.seriesGroupsFromMedia(rows) {
- for _, season := range e.seasonsForSeries(series) {
- if season.ID == id {
- e.rememberSeriesGroup(series)
- return season, true, nil
- }
- }
- }
- return embySeasonGroup{}, false, nil
-}
-
-func (e *EmbyService) seriesGroupsFromMedia(rows []model.Media) []embySeriesGroup {
- byID := map[string]*embySeriesGroup{}
- order := []string{}
- for _, row := range rows {
- row := row
- seriesID := e.seriesIDForMedia(&row)
- group, ok := byID[seriesID]
- if !ok {
- group = &embySeriesGroup{
- ID: seriesID,
- LibraryID: row.LibraryID,
- Name: e.seriesNameForMedia(&row),
- Year: row.Year,
- TMDbID: row.TMDbID,
- BangumiID: row.BangumiID,
- CreatedAt: row.CreatedAt,
- }
- byID[seriesID] = group
- order = append(order, seriesID)
- }
- if row.CreatedAt.After(group.CreatedAt) {
- group.CreatedAt = row.CreatedAt
- }
- if group.PosterURL == "" && row.PosterURL != "" {
- group.PosterURL = row.PosterURL
- }
- if group.BackdropURL == "" && row.BackdropURL != "" {
- group.BackdropURL = row.BackdropURL
- }
- if group.Overview == "" && row.Overview != "" {
- group.Overview = row.Overview
- }
- if group.Rating == 0 && row.Rating > 0 {
- group.Rating = row.Rating
- }
- if group.Year == 0 && row.Year > 0 {
- group.Year = row.Year
- }
- group.Episodes = append(group.Episodes, row)
- }
- groups := make([]embySeriesGroup, 0, len(order))
- for _, id := range order {
- group := *byID[id]
- sort.SliceStable(group.Episodes, func(i, j int) bool {
- if group.Episodes[i].SeasonNum != group.Episodes[j].SeasonNum {
- return group.Episodes[i].SeasonNum < group.Episodes[j].SeasonNum
- }
- if group.Episodes[i].EpisodeNum != group.Episodes[j].EpisodeNum {
- return group.Episodes[i].EpisodeNum < group.Episodes[j].EpisodeNum
- }
- return group.Episodes[i].CreatedAt.Before(group.Episodes[j].CreatedAt)
- })
- groups = append(groups, group)
- }
- return groups
-}
-
-func (e *EmbyService) seasonsForSeries(series embySeriesGroup) []embySeasonGroup {
- bySeason := map[int]*embySeasonGroup{}
- order := []int{}
- for _, episode := range series.Episodes {
- seasonNum := episode.SeasonNum
- if seasonNum < 0 {
- seasonNum = 1
- }
- season, ok := bySeason[seasonNum]
- if !ok {
- season = &embySeasonGroup{
- ID: seasonID(series.ID, seasonNum),
- SeriesID: series.ID,
- LibraryID: series.LibraryID,
- Name: seasonName(seasonNum),
- SeasonNum: seasonNum,
- Series: series,
- }
- bySeason[seasonNum] = season
- order = append(order, seasonNum)
- }
- season.Episodes = append(season.Episodes, episode)
- }
- sort.Ints(order)
- out := make([]embySeasonGroup, 0, len(order))
- for _, seasonNum := range order {
- out = append(out, *bySeason[seasonNum])
- }
- return out
-}
-
-func (e *EmbyService) seriesPayload(group embySeriesGroup) map[string]any {
- e.rememberSeriesGroup(group)
- imageTags := map[string]string{}
- backdropTags := []string{}
- if group.PosterURL != "" {
- imageTags["Primary"] = group.ID
- }
- if group.BackdropURL != "" {
- backdropTags = append(backdropTags, group.ID+"-bd")
- }
- return map[string]any{
- "Id": group.ID,
- "Name": group.Name,
- "ServerId": embyServerID,
- "Type": "Series",
- "MediaType": "Video",
- "IsFolder": true,
- "ParentId": group.LibraryID,
- "ProductionYear": group.Year,
- "Overview": group.Overview,
- "CommunityRating": group.Rating,
- "RecursiveItemCount": len(group.Episodes),
- "ChildCount": len(e.seasonsForSeries(group)),
- "DateCreated": group.CreatedAt,
- "ImageTags": imageTags,
- "BackdropImageTags": backdropTags,
- "ProviderIds": map[string]string{
- "Tmdb": intToStr(group.TMDbID),
- "Bangumi": intToStr(group.BangumiID),
- },
- "UserData": emptyUserData(),
- }
-}
-
-func (e *EmbyService) seasonPayload(season embySeasonGroup) map[string]any {
- e.rememberSeasonGroup(season)
- imageTags := map[string]string{}
- backdropTags := []string{}
- if season.Series.PosterURL != "" {
- imageTags["Primary"] = season.ID
- }
- if season.Series.BackdropURL != "" {
- backdropTags = append(backdropTags, season.ID+"-bd")
- }
- return map[string]any{
- "Id": season.ID,
- "Name": season.Name,
- "ServerId": embyServerID,
- "Type": "Season",
- "MediaType": "Video",
- "IsFolder": true,
- "ParentId": season.SeriesID,
- "SeriesId": season.SeriesID,
- "SeriesName": season.Series.Name,
- "IndexNumber": season.SeasonNum,
- "ChildCount": len(season.Episodes),
- "ImageTags": imageTags,
- "BackdropImageTags": backdropTags,
- "UserData": emptyUserData(),
- }
-}
-
-// ImageURL returns artwork for a media/series/season item id.
-func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (string, error) {
- pick := func(primary, backdrop string) string {
- switch strings.ToLower(imageType) {
- case "backdrop", "art":
- if backdrop != "" {
- return backdrop
- }
- }
- if primary != "" {
- return primary
- }
- return backdrop
- }
- if strings.HasPrefix(id, embyVirtualSeasonPrefix) {
- if raw, ok := e.cachedArtworkURL(id, imageType); ok {
- return raw, nil
- }
- return "", nil
- }
- if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
- if raw, ok := e.cachedArtworkURL(id, imageType); ok {
- return raw, nil
- }
- return "", nil
- }
- m, err := e.repo.Media.FindByID(ctx, id)
- if err == nil && m != nil {
- if e.mediaShouldBeEpisode(ctx, m) {
- switch strings.ToLower(imageType) {
- case "backdrop", "art":
- return "", nil
- }
- }
- return pick(e.mediaPrimaryArtwork(ctx, m), e.mediaBackdropArtwork(ctx, m)), nil
- }
- if err != nil {
- return "", err
- }
- if series, ok, err := e.findSeriesGroup(ctx, id, ""); err != nil {
- return "", err
- } else if ok {
- return pick(series.PosterURL, series.BackdropURL), nil
- }
- return "", nil
-}
-
-func (e *EmbyService) mediaPrimaryArtwork(ctx context.Context, m *model.Media) string {
- if m == nil {
- return ""
- }
- if e.mediaShouldBeEpisode(ctx, m) && strings.TrimSpace(m.BackdropURL) != "" {
- return m.BackdropURL
- }
- return m.PosterURL
-}
-
-func (e *EmbyService) mediaBackdropArtwork(ctx context.Context, m *model.Media) string {
- if m == nil {
- return ""
- }
- if e.mediaShouldBeEpisode(ctx, m) {
- return ""
- }
- return m.BackdropURL
-}
-
-func (e *EmbyService) seriesIDForMedia(m *model.Media) string {
- if strings.TrimSpace(m.SeriesID) != "" {
- return m.SeriesID
- }
- return stableEmbyID(embyVirtualSeriesPrefix, m.LibraryID, e.seriesNameForMedia(m))
-}
-
-func (e *EmbyService) seasonIDForMedia(m *model.Media) string {
- return seasonID(e.seriesIDForMedia(m), m.SeasonNum)
-}
-
-func (e *EmbyService) seriesNameForMedia(m *model.Media) string {
- if strings.TrimSpace(m.SeriesID) != "" {
- if series, err := e.repo.Series.FindByID(context.Background(), m.SeriesID); err == nil && series != nil && strings.TrimSpace(series.Title) != "" {
- return series.Title
- }
- }
- if name := inferSeriesNameFromPath(m.Path); name != "" {
- return name
- }
- name := strings.TrimSpace(m.Title)
- name = embyEpisodeTitleRE.ReplaceAllString(name, "")
- name = embyYearSuffixRE.ReplaceAllString(name, "")
- if name == "" {
- name = strings.TrimSpace(m.OriginalName)
- }
- return name
-}
-
-func inferSeriesNameFromPath(path string) string {
- path = strings.TrimSpace(path)
- if path == "" {
- return ""
- }
- dir := filepath.Dir(path)
- base := filepath.Base(dir)
- if embySeasonDirRE.MatchString(base) {
- dir = filepath.Dir(dir)
- base = filepath.Base(dir)
- }
- base = strings.TrimSpace(embyYearSuffixRE.ReplaceAllString(base, ""))
- if base == "." || base == string(filepath.Separator) {
- return ""
- }
- return base
-}
-
-func stableEmbyID(prefix string, parts ...string) string {
- h := sha256.New()
- for _, part := range parts {
- _, _ = h.Write([]byte(strings.ToLower(strings.TrimSpace(part))))
- _, _ = h.Write([]byte{0})
- }
- return prefix + hex.EncodeToString(h.Sum(nil))[:32]
-}
-
-func seasonID(seriesID string, seasonNum int) string {
- if seasonNum < 0 {
- seasonNum = 1
- }
- return stableEmbyID(embyVirtualSeasonPrefix, seriesID, strconv.Itoa(seasonNum))
-}
-
-func seasonName(seasonNum int) string {
- if seasonNum == 0 {
- return "特别篇"
- }
- if seasonNum < 0 {
- seasonNum = 1
- }
- return fmt.Sprintf("第 %d 季", seasonNum)
-}
-
-func sortSeriesGroups(groups []embySeriesGroup, p ItemsParams) {
- switch strings.ToLower(p.SortBy) {
- case "sortname", "name":
- sort.SliceStable(groups, func(i, j int) bool {
- if strings.EqualFold(p.SortOrder, "Descending") {
- return groups[i].Name > groups[j].Name
- }
- return groups[i].Name < groups[j].Name
- })
- default:
- sort.SliceStable(groups, func(i, j int) bool {
- if strings.EqualFold(p.SortOrder, "Ascending") {
- return groups[i].CreatedAt.Before(groups[j].CreatedAt)
- }
- return groups[i].CreatedAt.After(groups[j].CreatedAt)
- })
- }
-}
-
-func containsItemType(types []string, want string) bool {
- for _, t := range types {
- if strings.EqualFold(strings.TrimSpace(t), want) {
- return true
- }
- }
- return false
-}
-
-func containsSupportedEmbyItemType(types []string) bool {
- for _, itemType := range types {
- switch strings.ToLower(strings.TrimSpace(itemType)) {
- case "movie", "series", "season", "episode", "video", "folder", "collectionfolder":
- return true
- }
- }
- return false
-}
-
-func containsOnlyFolderItemTypes(types []string) bool {
- if len(types) == 0 {
- return false
- }
- for _, itemType := range types {
- switch strings.ToLower(strings.TrimSpace(itemType)) {
- case "folder", "collectionfolder":
- default:
- return false
- }
- }
- return true
-}
-
-func emptyItemsEnvelope(startIndex int) map[string]any {
- return map[string]any{
- "Items": []map[string]any{},
- "TotalRecordCount": int64(0),
- "StartIndex": startIndex,
- }
-}
-
-func containsEmbyFilter(filters []string, want string) bool {
- for _, filter := range filters {
- if strings.EqualFold(strings.TrimSpace(filter), want) {
- return true
- }
- }
- return false
-}
-
-func firstCSVValue(value string) string {
- if i := strings.Index(value, ","); i >= 0 {
- value = value[:i]
- }
- return strings.TrimSpace(value)
-}
-
-func primarySupportedEmbySort(sortBy string, resumeFilter bool) string {
- for _, part := range strings.Split(sortBy, ",") {
- key := strings.ToLower(strings.TrimSpace(part))
- switch key {
- case "sortname", "name", "premieredate", "productionyear", "datecreated", "communityrating":
- return key
- case "dateplayed":
- if resumeFilter {
- return key
- }
- }
- }
- return strings.ToLower(strings.TrimSpace(firstCSVValue(sortBy)))
-}
-
-func pageSlice[T any](items []T, start, limit int) []T {
- if start < 0 {
- start = 0
- }
- if limit <= 0 {
- limit = len(items)
- }
- if start >= len(items) {
- return []T{}
- }
- end := start + limit
- if end > len(items) {
- end = len(items)
- }
- return items[start:end]
-}
-
-func emptyUserData() map[string]any {
- return map[string]any{
- "PlaybackPositionTicks": 0,
- "PlayCount": 0,
- "IsFavorite": false,
- "Played": false,
- "PlayedPercentage": 0,
- }
-}
-
-func (e *EmbyService) applyUserMediaVisibility(ctx context.Context, q *gorm.DB, userID string) *gorm.DB {
- visibility := e.mediaVisibility(ctx, userID)
- if !visibility.IncludeNSFW {
- q = q.Where("nsfw = ?", false)
- if hidden := visibility.HiddenLibraryIDs; len(hidden) > 0 {
- q = q.Where("library_id NOT IN ?", hidden)
- }
- }
- if len(visibility.AllowedLibraryIDs) > 0 {
- q = q.Where("library_id IN ?", visibility.AllowedLibraryIDs)
- }
- return q
-}
-
-func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.Media, userID string) []model.Media {
- visibility := e.mediaVisibility(ctx, userID)
- if visibility.IncludeNSFW && len(visibility.AllowedLibraryIDs) == 0 {
- return rows
- }
- allowed := map[string]bool{}
- for _, id := range visibility.AllowedLibraryIDs {
- allowed[id] = true
- }
- hiddenLibraries := map[string]bool{}
- for _, id := range visibility.HiddenLibraryIDs {
- hiddenLibraries[id] = true
- }
- out := rows[:0]
- for _, row := range rows {
- if row.NSFW && !visibility.IncludeNSFW {
- continue
- }
- if hiddenLibraries[row.LibraryID] {
- continue
- }
- if len(allowed) > 0 && !allowed[row.LibraryID] {
- continue
- }
- out = append(out, row)
- }
- return out
-}
-
-func (e *EmbyService) mediaVisibility(ctx context.Context, userID string) MediaVisibility {
- if e == nil {
- return MediaVisibility{IncludeNSFW: true}
- }
- key := strings.TrimSpace(userID)
- now := time.Now()
- e.visibilityMu.RLock()
- entry, ok := e.visibilityCache[key]
- e.visibilityMu.RUnlock()
- if ok && now.Before(entry.expiresAt) {
- return cloneMediaVisibility(entry.visibility)
- }
-
- visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
- if !visibility.IncludeNSFW {
- visibility.HiddenLibraryIDs = e.hiddenLibraryIDs(ctx, visibility)
- }
- visibility = ExpandMediaVisibilityForMergedCloudLibraries(ctx, e.repo, visibility)
- visibility = cloneMediaVisibility(visibility)
-
- e.visibilityMu.Lock()
- if e.visibilityCache == nil {
- e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
- }
- if len(e.visibilityCache) > 1000 {
- e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
- }
- e.visibilityCache[key] = embyVisibilityCacheEntry{
- visibility: cloneMediaVisibility(visibility),
- expiresAt: now.Add(embyVisibilityCacheTTL),
- }
- e.visibilityMu.Unlock()
-
- return visibility
-}
-
-func (e *EmbyService) mergedLibraryIDs(ctx context.Context, libraryID string) []string {
- ids, err := MergedLibraryIDsForLibrary(ctx, e.repo, libraryID)
- if err != nil || len(ids) == 0 {
- return []string{libraryID}
- }
- return ids
-}
-
-func cloneMediaVisibility(visibility MediaVisibility) MediaVisibility {
- if visibility.AllowedLibraryIDs != nil {
- visibility.AllowedLibraryIDs = append([]string(nil), visibility.AllowedLibraryIDs...)
- }
- if visibility.HiddenLibraryIDs != nil {
- visibility.HiddenLibraryIDs = append([]string(nil), visibility.HiddenLibraryIDs...)
- }
- return visibility
-}
-
-func (e *EmbyService) libraryVisibleFromCachedVisibility(lib model.Library, visibility MediaVisibility) bool {
- if len(visibility.AllowedLibraryIDs) > 0 {
- allowed := false
- for _, id := range visibility.AllowedLibraryIDs {
- if id == lib.ID {
- allowed = true
- break
- }
- }
- if !allowed {
- return false
- }
- }
- if visibility.IncludeNSFW {
- return true
- }
- for _, id := range visibility.HiddenLibraryIDs {
- if id == lib.ID {
- return false
- }
- }
- return true
-}
-
-func (e *EmbyService) hiddenLibraryIDs(ctx context.Context, visibility MediaVisibility) []string {
- if visibility.IncludeNSFW {
- return nil
- }
- libs, err := e.repo.Library.List(ctx)
- if err != nil {
- return nil
- }
- shadowed := ShadowedCloudLibraryIDSet(libs)
- ids := make([]string, 0)
- for _, lib := range libs {
- if shadowed[lib.ID] || !LibraryVisibleForUser(ctx, e.repo, lib, visibility) {
- ids = append(ids, lib.ID)
- }
- }
- return ids
-}
-
-func minInt(a, b int) int {
- if a < b {
- return a
- }
- return b
-}
-
-func maxInt(a, b int) int {
- if a > b {
- return a
- }
- return b
-}
-
-// ─── Playback ────────────────────────────────────────────────────────────────
-
-// PlaybackInfo returns a PlaybackInfoResponse usable by Emby clients.
-func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string) (map[string]any, error) {
- m, err := e.playableMedia(ctx, mediaID, userID)
- if err != nil || m == nil {
- return nil, err
- }
- e.ensureCloudTrackMetadata(ctx, m)
- return map[string]any{
- "MediaSources": e.mediaSourcesForItem(ctx, m, false, e.directPlayOnly(ctx)),
- "PlaySessionId": fmt.Sprintf("%s-%d", m.ID, time.Now().Unix()),
- }, nil
-}
-
-// ensureCloudTrackMetadata 在后台补齐云盘媒体的轨道元数据。
-//
-// 注意必须是异步的:此前这里在 PlaybackInfo 请求路径上同步执行
-// CloudResolve + ffprobe(HTTP)(最长 8 秒),既把第三方播放器的起播时间
-// 拖长到秒级,又让每一次点开详情/起播都可能触发一次云盘数据下载,是
-// Docker 部署下 CPU/带宽长期居高的来源之一。探测结果落库后,下一次
-// 请求自然能读到完整元数据。
-func (e *EmbyService) ensureCloudTrackMetadata(ctx context.Context, m *model.Media) {
- if e == nil || m == nil || e.storage == nil || e.probe == nil || !mediaTrackMetadataMissing(m) {
- return
- }
- typ, ref, ok := parseCloudMediaPlaybackURL(m.STRMURL)
- if !ok {
- return
- }
- mediaID := m.ID
- e.cloudProbeMu.Lock()
- if e.cloudProbeInFlight == nil {
- e.cloudProbeInFlight = make(map[string]struct{})
- }
- if _, busy := e.cloudProbeInFlight[mediaID]; busy {
- e.cloudProbeMu.Unlock()
- return
- }
- e.cloudProbeInFlight[mediaID] = struct{}{}
- e.cloudProbeMu.Unlock()
-
- go func() {
- defer func() {
- e.cloudProbeMu.Lock()
- delete(e.cloudProbeInFlight, mediaID)
- e.cloudProbeMu.Unlock()
- }()
- probeCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
- defer cancel()
- link, err := e.storage.CloudResolve(probeCtx, typ, ref, "")
- if err != nil {
- if e.log != nil {
- e.log.Debug("resolve cloud media for playback probe failed", zap.String("media_id", mediaID), zap.Error(err))
- }
- return
- }
- probe, err := e.probe.ProbeHTTP(probeCtx, link.URL, link.Headers)
- if err != nil {
- if e.log != nil {
- e.log.Debug("playback cloud ffprobe failed", zap.String("media_id", mediaID), zap.Error(err))
- }
- return
- }
- updates := probeResultUpdates(probe)
- if len(updates) == 0 {
- return
- }
- if err := e.repo.DB.WithContext(probeCtx).Model(&model.Media{}).Where("id = ?", mediaID).Updates(updates).Error; err != nil && e.log != nil {
- e.log.Debug("persist playback cloud probe failed", zap.String("media_id", mediaID), zap.Error(err))
- }
- }()
-}
-
-func mediaTrackMetadataMissing(m *model.Media) bool {
- return m.DurationSec <= 0 ||
- m.Width <= 0 ||
- m.Height <= 0 ||
- strings.TrimSpace(m.VideoCodec) == "" ||
- strings.TrimSpace(m.AudioCodec) == ""
-}
-
-func parseCloudMediaPlaybackURL(raw string) (string, string, bool) {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return "", "", false
- }
- u, err := url.Parse(raw)
- if err != nil {
- return "", "", false
- }
- path := strings.Trim(u.Path, "/")
- const prefix = "api/cloud/play/"
- idx := strings.Index(strings.ToLower(path), prefix)
- if idx < 0 {
- return "", "", false
- }
- typ := strings.TrimSpace(path[idx+len(prefix):])
- ref := strings.TrimSpace(u.Query().Get("ref"))
- return typ, ref, typ != "" && ref != ""
-}
-
-func applyProbeResultToMediaValue(m *model.Media, probe *ProbeResult) {
- if m == nil || probe == nil {
- return
- }
- if probe.DurationSec > 0 {
- m.DurationSec = probe.DurationSec
- }
- if probe.Width > 0 {
- m.Width = probe.Width
- }
- if probe.Height > 0 {
- m.Height = probe.Height
- }
- if strings.TrimSpace(probe.VideoCodec) != "" {
- m.VideoCodec = probe.VideoCodec
- }
- if strings.TrimSpace(probe.AudioCodec) != "" {
- m.AudioCodec = probe.AudioCodec
- }
- if strings.TrimSpace(probe.Container) != "" {
- m.Container = probe.Container
- }
-}
-
-// directPlayOnly reports whether the admin enabled「客户端直连解码」mode.
-// In that mode the host never transcodes; clients must direct-play.
-func (e *EmbyService) directPlayOnly(ctx context.Context) bool {
- if e.repo == nil || e.repo.Setting == nil {
- return false
- }
- v, err := e.repo.Setting.Get(ctx, PlaybackDirectOnlySettingKey)
- if err != nil {
- return false
- }
- return parseBoolSetting(v, false)
-}
-
-func (e *EmbyService) playableMedia(ctx context.Context, id, userID string) (*model.Media, error) {
- if season, ok, err := e.findSeasonGroup(ctx, id, userID); err != nil {
- return nil, err
- } else if ok && len(season.Episodes) > 0 {
- return &season.Episodes[0], nil
- }
- if series, ok, err := e.findSeriesGroup(ctx, id, userID); err != nil {
- return nil, err
- } else if ok && len(series.Episodes) > 0 {
- return &series.Episodes[0], nil
- }
- m, err := e.repo.Media.FindByID(ctx, id)
- if err != nil || m == nil {
- return m, err
- }
- if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
- return nil, nil
- }
- return m, nil
-}
-
-// mediaSource 是 /Items 与 /PlaybackInfo 共享的 MediaSource 结构。
-//
-// asEmbedded=true:嵌在 /Items 列表里,不包含完整 stream URL(避免暴露
-// 直链给搜索接口)。/PlaybackInfo 走 false 路径,URL 指向 Emby 兼容
-// /Videos/{id}/stream(客户端会继续携带 X-Emby-Token 或 append api_key)。
-func (e *EmbyService) mediaSource(ctx context.Context, m *model.Media, asEmbedded, directOnly bool) map[string]any {
- container := strings.Trim(strings.ToLower(m.Container), ". ")
- if container == "" {
- container = strings.TrimPrefix(strings.ToLower(filepath.Ext(m.Path)), ".")
- }
- if container == "" && strings.TrimSpace(m.STRMURL) != "" {
- container = "strm"
- }
- isCloud := strings.TrimSpace(m.STRMURL) != ""
- playURL := embyDirectStreamURL(m.ID, container)
- if isCloud {
- switch CloudPlaybackMode(ctx, e.repo) {
- case CloudPlaybackModeSTRM:
- playURL = embySTRMStreamURL(m.ID)
- case CloudPlaybackModeRedirectProxy:
- playURL = embyDirectStreamURL(m.ID, container)
- default:
- playURL = ""
- }
- }
- if isCloud {
- // Cloud/WebDAV media is already a direct/proxy stream. Advertising HLS
- // transcoding makes some Emby clients pick /master.m3u8, forcing this
- // lightweight server to pull remote bytes through ffmpeg and often
- // surfacing as "network/playback failed". Keep cloud media direct-only.
- directOnly = true
- }
- src := map[string]any{
- "Id": m.ID,
- "Name": m.Title,
- "Path": m.Path,
- "Container": container,
- "Size": m.SizeBytes,
- "Protocol": "Http",
- "Type": "Default",
- "IsRemote": isCloud,
- "RequiresOpening": false,
- "RequiresClosing": false,
- "ReadAtNativeFramerate": false,
- "SupportsTranscoding": !directOnly,
- // 云盘媒体的 Path 在 PlaybackInfo 阶段会被补上 api_key,且最终
- // 302 到云盘直链。Infuse/Emby 官方客户端会优先挑选 DirectPlay
- // 源;如果这里标 false,即使 DirectStreamUrl 可用,也可能被判定
- // 为“没有可播放媒体源”。
- "SupportsDirectStream": !isCloud || playURL != "",
- "SupportsDirectPlay": !isCloud || playURL != "",
- "SupportsProbing": true,
- "RunTimeTicks": int64(m.DurationSec) * 10_000_000,
- "MediaStreams": e.mediaStreams(m),
- }
- if !asEmbedded && playURL != "" {
- src["DirectStreamUrl"] = playURL
- // 直连解码模式下不下发 TranscodingUrl,迫使客户端本地解码直连,
- // 宿主机不参与转码。
- if !directOnly {
- src["TranscodingUrl"] = "/Videos/" + m.ID + "/master.m3u8"
- }
- }
- if strings.TrimSpace(m.STRMURL) != "" && playURL != "" {
- // STRM / cloud:// media must stay behind a token-aware endpoint. When
- // STRM playback is enabled we expose /api/stream so third-party clients
- // follow the same STRM entry as generated .strm files; when disabled we
- // expose /Videos/{id}/stream so playback uses the Emby 302/proxy path.
- src["IsRemote"] = true
- src["Path"] = playURL
- }
- return src
-}
-
-func (e *EmbyService) mediaSourcesForItem(ctx context.Context, m *model.Media, asEmbedded, directOnly bool) []map[string]any {
- siblings := e.mediaVersionSiblings(ctx, m)
- if len(siblings) == 0 {
- return []map[string]any{e.mediaSource(ctx, m, asEmbedded, directOnly)}
- }
- sources := make([]map[string]any, 0, len(siblings))
- for i := range siblings {
- media := siblings[i]
- sources = append(sources, e.mediaSource(ctx, &media, asEmbedded, directOnly))
- }
- return sources
-}
-
-func (e *EmbyService) mediaVersionSiblings(ctx context.Context, m *model.Media) []model.Media {
- if e == nil || e.repo == nil || e.repo.DB == nil || m == nil || strings.TrimSpace(m.ID) == "" {
- return nil
- }
- libraryIDs := e.mergedLibraryIDs(ctx, m.LibraryID)
- if len(libraryIDs) == 0 {
- libraryIDs = []string{m.LibraryID}
- }
- q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("library_id IN ?", libraryIDs).
- Where("season_num = ? AND episode_num = ?", m.SeasonNum, m.EpisodeNum)
- if m.TMDbID > 0 {
- q = q.Where("tm_db_id = ?", m.TMDbID)
- } else if m.BangumiID > 0 {
- q = q.Where("bangumi_id = ?", m.BangumiID)
- } else {
- title := strings.TrimSpace(m.Title)
- if title == "" {
- title = strings.TrimSpace(m.OriginalName)
- }
- if title == "" {
- return []model.Media{*m}
- }
- q = q.Where("LOWER(title) = ?", strings.ToLower(title))
- if m.Year > 0 {
- q = q.Where("year = ?", m.Year)
- }
- }
- var rows []model.Media
- if err := q.Find(&rows).Error; err != nil || len(rows) == 0 {
- return []model.Media{*m}
- }
- rows = e.collapseExactPathRows(rows)
- sort.SliceStable(rows, func(i, j int) bool {
- if rows[i].ID == m.ID {
- return true
- }
- if rows[j].ID == m.ID {
- return false
- }
- return preferMediaVersion(rows[i], rows[j])
- })
- return rows
-}
-
-func (e *EmbyService) collapseExactPathRows(rows []model.Media) []model.Media {
- if len(rows) < 2 {
- return rows
- }
- out := rows[:0]
- seen := map[string]struct{}{}
- for _, row := range rows {
- path := strings.TrimSpace(row.Path)
- if path != "" {
- if _, ok := seen[path]; ok {
- continue
- }
- seen[path] = struct{}{}
- }
- out = append(out, row)
- }
- return out
-}
-
-func (e *EmbyService) mediaVersionKey(ctx context.Context, m *model.Media) string {
- if e == nil || m == nil {
- return ""
- }
- ids := e.mergedLibraryIDs(ctx, m.LibraryID)
- sort.Strings(ids)
- libraryGroup := strings.Join(ids, ",")
- if libraryGroup == "" {
- libraryGroup = strings.TrimSpace(m.LibraryID)
- }
- if m.TMDbID > 0 {
- return fmt.Sprintf("%s|tmdb:%d|s:%d|e:%d", libraryGroup, m.TMDbID, m.SeasonNum, m.EpisodeNum)
- }
- if m.BangumiID > 0 {
- return fmt.Sprintf("%s|bangumi:%d|s:%d|e:%d", libraryGroup, m.BangumiID, m.SeasonNum, m.EpisodeNum)
- }
- title := strings.ToLower(strings.TrimSpace(m.Title))
- if title == "" {
- title = strings.ToLower(strings.TrimSpace(m.OriginalName))
- }
- if title == "" {
- return ""
- }
- return fmt.Sprintf("%s|title:%s|y:%d|s:%d|e:%d", libraryGroup, title, m.Year, m.SeasonNum, m.EpisodeNum)
-}
-
-func preferMediaVersion(candidate, current model.Media) bool {
- candidateCloud := strings.TrimSpace(candidate.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(candidate.Path)), "cloud://")
- currentCloud := strings.TrimSpace(current.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(current.Path)), "cloud://")
- if candidateCloud != currentCloud {
- return !candidateCloud
- }
- if candidate.Width != current.Width {
- return candidate.Width > current.Width
- }
- if candidate.SizeBytes != current.SizeBytes {
- return candidate.SizeBytes > current.SizeBytes
- }
- return candidate.CreatedAt.After(current.CreatedAt)
-}
-
-func embySTRMStreamURL(mediaID string) string {
- return "/api/stream/" + url.PathEscape(strings.TrimSpace(mediaID))
-}
-
-func embyDirectStreamURL(mediaID, container string) string {
- mediaID = strings.TrimSpace(mediaID)
- container = strings.Trim(strings.ToLower(container), ". ")
- if container == "" || container == "strm" {
- return "/Videos/" + mediaID + "/stream"
- }
- return "/Videos/" + mediaID + "/stream." + container
-}
-
-func (e *EmbyService) mediaStreams(m *model.Media) []map[string]any {
- streams := []map[string]any{}
- if m.VideoCodec != "" || m.Width > 0 {
- streams = append(streams, map[string]any{
- "Codec": m.VideoCodec,
- "Type": "Video",
- "Index": 0,
- "Width": m.Width,
- "Height": m.Height,
- "AspectRatio": "",
- "IsDefault": true,
- "IsForced": false,
- "IsExternal": false,
- "DisplayTitle": fmt.Sprintf("%dx%d %s", m.Width, m.Height, m.VideoCodec),
- })
- }
- if m.AudioCodec != "" {
- streams = append(streams, map[string]any{
- "Codec": m.AudioCodec,
- "Type": "Audio",
- "Index": 1,
- "IsDefault": true,
- "IsForced": false,
- "IsExternal": false,
- })
- }
- if len(streams) == 0 {
- streams = append(streams, map[string]any{
- "Codec": "unknown",
- "Type": "Video",
- "Index": 0,
- "IsDefault": true,
- "IsForced": false,
- "IsExternal": false,
- "DisplayTitle": "Video",
- })
- }
- return streams
-}
-
-// ─── 收藏 / 已看(Emby 客户端写路径) ──────────────────────────────────────
-
-// SetFavorite 把 mediaID 标为 userID 的收藏。
-func (e *EmbyService) SetFavorite(ctx context.Context, userID, mediaID string, favorite bool) error {
- if favorite {
- var f model.Favorite
- err := e.repo.DB.WithContext(ctx).
- Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return e.repo.DB.WithContext(ctx).Create(&model.Favorite{
- UserID: userID, MediaID: mediaID,
- }).Error
- }
- return err
- }
- return e.repo.DB.WithContext(ctx).
- Where("user_id = ? AND media_id = ?", userID, mediaID).
- Delete(&model.Favorite{}).Error
-}
-
-// MarkPlayed 把 mediaID 标为已看(写一个 100% 进度的 history 行)。
-func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, played bool) error {
- if !played {
- return e.repo.DB.WithContext(ctx).
- Where("user_id = ? AND media_id = ?", userID, mediaID).
- Delete(&model.PlaybackHistory{}).Error
- }
- m, err := e.repo.Media.FindByID(ctx, mediaID)
- if err != nil || m == nil {
- return errors.New("media not found")
- }
- dur := int64(m.DurationSec) * 1000
- if dur <= 0 {
- dur = 1
- }
- return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
- UserID: userID,
- MediaID: mediaID,
- PositionMs: dur,
- DurationMs: dur,
- WatchedAt: time.Now(),
- Completed: true,
- })
-}
-
-// RecordProgress 记录播放进度(来自 Emby 客户端的 /Sessions/Playing/Progress)。
-func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string, positionTicks, runtimeTicks int64) error {
- pos := positionTicks / 10_000
- dur := runtimeTicks / 10_000
- if dur <= 0 {
- // runtimeTicks 缺失时回退到 media.DurationSec
- if m, _ := e.repo.Media.FindByID(ctx, mediaID); m != nil {
- dur = int64(m.DurationSec) * 1000
- }
- }
- completed := dur > 0 && pos >= dur*9/10
- return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
- UserID: userID,
- MediaID: mediaID,
- PositionMs: pos,
- DurationMs: dur,
- WatchedAt: time.Now(),
- Completed: completed,
- })
-}
-
-// ─── Helpers ─────────────────────────────────────────────────────────────────
-
-func splitCSV(s string) []string {
- if strings.TrimSpace(s) == "" {
- return []string{}
- }
- parts := strings.Split(s, ",")
- out := make([]string, 0, len(parts))
- for _, p := range parts {
- p = strings.TrimSpace(p)
- if p != "" {
- out = append(out, p)
- }
- }
- return out
-}
-
-func intToStr(v int) string {
- if v == 0 {
- return ""
- }
- return strconv.Itoa(v)
-}
diff --git a/internal/service/emby_compat_test.go b/internal/service/emby_compat_test.go
index 3dfc80e..2631990 100644
--- a/internal/service/emby_compat_test.go
+++ b/internal/service/emby_compat_test.go
@@ -5,9 +5,7 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -26,7 +24,8 @@ func TestEmbyItemsExposeSeriesSeasonEpisodeHierarchy(t *testing.T) {
Base: model.Base{ID: "ep-1"},
LibraryID: lib.ID,
Title: "间谍过家家",
- OriginalName: "第 1 集",
+ OriginalName: "SPY×FAMILY",
+ EpisodeTitle: "第 1 集",
Path: `F:\downloads\日番\剧集\间谍过家家\Season 02\间谍过家家 - S02E01.mkv`,
PosterURL: `F:\poster.jpg`,
SeasonNum: 2,
@@ -36,7 +35,8 @@ func TestEmbyItemsExposeSeriesSeasonEpisodeHierarchy(t *testing.T) {
Base: model.Base{ID: "ep-2"},
LibraryID: lib.ID,
Title: "间谍过家家",
- OriginalName: "第 2 集",
+ OriginalName: "SPY×FAMILY",
+ EpisodeTitle: "第 2 集",
Path: `F:\downloads\日番\剧集\间谍过家家\Season 02\间谍过家家 - S02E02.mkv`,
PosterURL: `F:\poster.jpg`,
SeasonNum: 2,
@@ -593,468 +593,14 @@ func TestEmbyMergedLocalCloudMovieVersionsShareMediaSources(t *testing.T) {
}
}
-func TestEmbyRootItemsExposeLibraries(t *testing.T) {
- svc := newTestEmbyService(t)
- for _, lib := range []model.Library{
- {Name: "电影", Path: `F:\downloads\电影`, Type: "movie", Enabled: true},
- {Name: "综艺", Path: `F:\downloads\综艺`, Type: "variety", Enabled: true},
- } {
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- }
-
- root, err := svc.Items(t.Context(), ItemsParams{Limit: 50})
- if err != nil {
- t.Fatalf("root items: %v", err)
- }
- items := root["Items"].([]map[string]any)
- if len(items) != 2 {
- t.Fatalf("expected root items to expose libraries, got %#v", items)
- }
- if items[0]["Type"] != "CollectionFolder" || items[1]["Type"] != "CollectionFolder" {
- t.Fatalf("root should return collection folders: %#v", items)
- }
- if items[1]["CollectionType"] != "tvshows" {
- t.Fatalf("variety libraries should use tvshows collection type: %#v", items[1])
- }
-}
-
-func TestEmbyFolderItemQueryExposesLibrariesForHome(t *testing.T) {
- svc := newTestEmbyService(t)
- lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := svc.repo.DB.Create(&model.Media{Base: model.Base{ID: "movie-1"}, LibraryID: lib.ID, Title: "不应出现在文件夹查询", Path: `/media/movies/a.mkv`}).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- out, err := svc.Items(t.Context(), ItemsParams{
- IncludeItemTypes: []string{"Folder", "CollectionFolder"},
- Limit: 50,
- })
- if err != nil {
- t.Fatalf("folder items: %v", err)
- }
- items := out["Items"].([]map[string]any)
- if len(items) != 1 {
- t.Fatalf("expected one library folder, got %#v", items)
- }
- if items[0]["Type"] != "CollectionFolder" || items[0]["IsFolder"] != true {
- t.Fatalf("folder query should return collection folders, got %#v", items[0])
- }
-}
-
-func TestEmbyUnsupportedItemTypesDoNotLeakAllMedia(t *testing.T) {
- svc := newTestEmbyService(t)
- lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- if err := svc.repo.DB.Create(&model.Media{Base: model.Base{ID: "movie-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/a.mkv`}).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- for _, includeType := range []string{"BoxSet", "Game", "Book", "Audio", "MusicAlbum", "Playlist", "TvChannel"} {
- out, err := svc.Items(t.Context(), ItemsParams{
- IncludeItemTypes: []string{includeType},
- Recursive: true,
- Limit: 50,
- })
- if err != nil {
- t.Fatalf("%s items: %v", includeType, err)
- }
- if out["TotalRecordCount"] != int64(0) {
- t.Fatalf("%s should not return media rows, got %#v", includeType, out)
- }
- items := out["Items"].([]map[string]any)
- if len(items) != 0 {
- t.Fatalf("%s should return an empty list, got %#v", includeType, items)
- }
- }
-}
-
-func TestEmbyItemsFiltersFavorites(t *testing.T) {
- svc := newTestEmbyService(t)
- viewer := &model.User{Base: model.Base{ID: "user-1"}, Username: "viewer", Role: "user", Tier: "free", IsActive: true}
- if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
- t.Fatalf("create viewer: %v", err)
- }
- lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- favorite := model.Media{Base: model.Base{ID: "fav-1"}, LibraryID: lib.ID, Title: "收藏电影", Path: `/media/movies/fav.mkv`}
- normal := model.Media{Base: model.Base{ID: "normal-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/normal.mkv`}
- if err := svc.repo.DB.Create(&favorite).Error; err != nil {
- t.Fatalf("create favorite media: %v", err)
- }
- if err := svc.repo.DB.Create(&normal).Error; err != nil {
- t.Fatalf("create normal media: %v", err)
- }
- if err := svc.repo.DB.Create(&model.Favorite{UserID: viewer.ID, MediaID: favorite.ID}).Error; err != nil {
- t.Fatalf("create favorite: %v", err)
- }
-
- out, err := svc.Items(t.Context(), ItemsParams{
- UserID: viewer.ID,
- Filters: []string{"IsFavorite"},
- Recursive: true,
- Limit: 50,
- })
- if err != nil {
- t.Fatalf("favorite items: %v", err)
- }
- if out["TotalRecordCount"] != int64(1) {
- t.Fatalf("expected one favorite, got %#v", out)
- }
- items := out["Items"].([]map[string]any)
- if len(items) != 1 || items[0]["Id"] != favorite.ID {
- t.Fatalf("favorite filter returned wrong items: %#v", items)
- }
- userData := items[0]["UserData"].(map[string]any)
- if userData["IsFavorite"] != true {
- t.Fatalf("favorite payload should carry IsFavorite=true: %#v", userData)
- }
-}
-
-func TestEmbyItemsFiltersResumableForHome(t *testing.T) {
- svc := newTestEmbyService(t)
- viewer := &model.User{Base: model.Base{ID: "user-1"}, Username: "viewer", Role: "user", Tier: "free", IsActive: true}
- if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
- t.Fatalf("create viewer: %v", err)
- }
- lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- resumable := model.Media{Base: model.Base{ID: "resume-1"}, LibraryID: lib.ID, Title: "继续观看", Path: `/media/movies/resume.mkv`, DurationSec: 120}
- normal := model.Media{Base: model.Base{ID: "normal-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/normal.mkv`, DurationSec: 120}
- if err := svc.repo.DB.Create(&resumable).Error; err != nil {
- t.Fatalf("create resumable media: %v", err)
- }
- if err := svc.repo.DB.Create(&normal).Error; err != nil {
- t.Fatalf("create normal media: %v", err)
- }
- if err := svc.repo.DB.Create(&model.PlaybackHistory{
- UserID: viewer.ID,
- MediaID: resumable.ID,
- PositionMs: 30_000,
- DurationMs: 120_000,
- WatchedAt: time.Now(),
- Completed: false,
- }).Error; err != nil {
- t.Fatalf("create playback history: %v", err)
- }
-
- out, err := svc.Items(t.Context(), ItemsParams{
- UserID: viewer.ID,
- Filters: []string{"IsResumable"},
- Recursive: true,
- SortBy: "DatePlayed",
- SortOrder: "Descending",
- Limit: 50,
- StartIndex: 0,
- })
- if err != nil {
- t.Fatalf("resumable items: %v", err)
- }
- if out["TotalRecordCount"] != int64(1) {
- t.Fatalf("expected one resumable item, got %#v", out)
- }
- items := out["Items"].([]map[string]any)
- if len(items) != 1 || items[0]["Id"] != resumable.ID {
- t.Fatalf("resumable filter returned wrong items: %#v", items)
- }
-}
-
-func TestEmbyUserPolicyDisablesDownloadsForViewers(t *testing.T) {
- svc := newTestEmbyService(t)
- viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true}
- admin := &model.User{Username: "admin", Role: "admin", Tier: "plus", IsActive: true}
- if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
- t.Fatalf("create viewer: %v", err)
- }
- if err := svc.repo.User.Create(t.Context(), admin); err != nil {
- t.Fatalf("create admin: %v", err)
- }
-
- viewerPayload, err := svc.FindUser(t.Context(), viewer.ID)
- if err != nil {
- t.Fatalf("viewer payload: %v", err)
- }
- adminPayload, err := svc.FindUser(t.Context(), admin.ID)
- if err != nil {
- t.Fatalf("admin payload: %v", err)
- }
- viewerPolicy := viewerPayload["Policy"].(map[string]any)
- adminPolicy := adminPayload["Policy"].(map[string]any)
- if viewerPolicy["EnableMediaPlayback"] != true {
- t.Fatalf("viewer must keep playback enabled: %#v", viewerPolicy)
- }
- if viewerPolicy["EnableContentDownloading"] != false ||
- viewerPolicy["EnableSyncTranscoding"] != false ||
- viewerPolicy["EnableMediaConversion"] != false {
- t.Fatalf("viewer must not be allowed to download/sync media: %#v", viewerPolicy)
- }
- if adminPolicy["EnableContentDownloading"] != true {
- t.Fatalf("admin should keep downloading capability: %#v", adminPolicy)
- }
-}
-
-func TestEmbyHidesAdultLibrariesForUserLock(t *testing.T) {
- svc := newTestEmbyService(t)
- viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true, HideAdult: true}
- if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
- t.Fatalf("create viewer: %v", err)
- }
- safe := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- adult := model.Library{Name: "9KG 成人", Path: `/media/9KG`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &safe); err != nil {
- t.Fatalf("create safe library: %v", err)
- }
- if err := svc.repo.Library.Create(t.Context(), &adult); err != nil {
- t.Fatalf("create adult library: %v", err)
- }
- if err := svc.repo.Setting.Set(t.Context(), AdultLibraryIDsSettingKey, `["`+adult.ID+`"]`); err != nil {
- t.Fatalf("set adult libraries: %v", err)
- }
- if err := svc.repo.DB.Create(&model.Media{LibraryID: safe.ID, Title: "安全电影", Path: `/media/movies/a.mkv`}).Error; err != nil {
- t.Fatalf("create safe media: %v", err)
- }
- if err := svc.repo.DB.Create(&model.Media{LibraryID: adult.ID, Title: "成人电影", Path: `/media/9KG/a.mkv`}).Error; err != nil {
- t.Fatalf("create adult media: %v", err)
- }
-
- root, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, Limit: 50})
- if err != nil {
- t.Fatalf("root items: %v", err)
- }
- items := root["Items"].([]map[string]any)
- if len(items) != 1 || items[0]["Name"] != "电影" {
- t.Fatalf("adult library should be hidden: %#v", items)
- }
- adultItems, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, ParentID: adult.ID, Limit: 50})
- if err != nil {
- t.Fatalf("adult items: %v", err)
- }
- if got := adultItems["TotalRecordCount"]; got != int64(0) {
- t.Fatalf("adult media should be hidden, total=%#v payload=%#v", got, adultItems)
- }
-}
-
-func TestEmbyPlaybackInfoRespectsDirectPlayOnly(t *testing.T) {
- svc := newTestEmbyService(t)
- lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- media := model.Media{Base: model.Base{ID: "m-1"}, LibraryID: lib.ID, Title: "Inception", Path: `/media/movies/inception.mkv`}
- if err := svc.repo.DB.Create(&media).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- // 默认(关闭):宿主机可转码,下发 TranscodingUrl。
- pb, err := svc.PlaybackInfo(t.Context(), "m-1", "user-1")
- if err != nil {
- t.Fatalf("playback info: %v", err)
- }
- src := pb["MediaSources"].([]map[string]any)[0]
- if src["SupportsTranscoding"] != true {
- t.Fatalf("expected SupportsTranscoding=true by default, got %#v", src["SupportsTranscoding"])
- }
- if _, ok := src["TranscodingUrl"]; !ok {
- t.Fatalf("expected TranscodingUrl present by default: %#v", src)
- }
- if src["TranscodingUrl"] != "/Videos/m-1/master.m3u8" {
- t.Fatalf("expected HLS TranscodingUrl by default, got %#v", src["TranscodingUrl"])
- }
-
- // 开启「客户端直连解码」:不再下发转码能力 / TranscodingUrl,仍保留 DirectStream。
- if err := svc.repo.Setting.Set(t.Context(), PlaybackDirectOnlySettingKey, "true"); err != nil {
- t.Fatalf("enable direct-only: %v", err)
- }
- pb, err = svc.PlaybackInfo(t.Context(), "m-1", "user-1")
- if err != nil {
- t.Fatalf("playback info (direct-only): %v", err)
- }
- src = pb["MediaSources"].([]map[string]any)[0]
- if src["SupportsTranscoding"] != false {
- t.Fatalf("expected SupportsTranscoding=false in direct-only mode, got %#v", src["SupportsTranscoding"])
- }
- if _, ok := src["TranscodingUrl"]; ok {
- t.Fatalf("expected no TranscodingUrl in direct-only mode: %#v", src)
- }
- if src["SupportsDirectPlay"] != true || src["DirectStreamUrl"] != "/Videos/m-1/stream.mkv" {
- t.Fatalf("direct-only must still allow direct play: %#v", src)
- }
-}
-
-func TestEmbyPlaybackInfoKeepsSTRMBehindStreamEndpoint(t *testing.T) {
- svc := newTestEmbyService(t)
- if err := svc.repo.Setting.Set(t.Context(), CloudPlaybackModeSettingKey, CloudPlaybackModeSTRM); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- lib := model.Library{Name: "夸克网盘", Path: `cloud://quark/0`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- media := model.Media{
- Base: model.Base{ID: "cloud-1"},
- LibraryID: lib.ID,
- Title: "Cloud Movie",
- Path: `cloud://quark/f1`,
- STRMURL: `/api/cloud/play/quark?ref=f1`,
- }
- if err := svc.repo.DB.Create(&media).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- pb, err := svc.PlaybackInfo(t.Context(), "cloud-1", "user-1")
- if err != nil {
- t.Fatalf("playback info: %v", err)
- }
- src := pb["MediaSources"].([]map[string]any)[0]
- if src["IsRemote"] != true {
- t.Fatalf("strm media should be marked remote: %#v", src)
- }
- if src["DirectStreamUrl"] != "/api/stream/cloud-1" {
- t.Fatalf("strm playback should prefer /api/stream when enabled: %#v", src)
- }
- if src["Path"] != "/api/stream/cloud-1" {
- t.Fatalf("path should prefer /api/stream when enabled: %#v", src)
- }
- streams := src["MediaStreams"].([]map[string]any)
- if len(streams) == 0 || streams[0]["Type"] != "Video" {
- t.Fatalf("strm media should expose a fallback video stream for Android clients: %#v", src)
- }
-}
-
-func TestEmbyPlaybackInfoUsesVideoStreamWhenSTRMDisabled(t *testing.T) {
- svc := newTestEmbyService(t)
- if err := svc.repo.Setting.Set(t.Context(), CloudPlaybackModeSettingKey, CloudPlaybackModeRedirectProxy); err != nil {
- t.Fatalf("set cloud playback mode: %v", err)
- }
- lib := model.Library{Name: "OpenList", Path: `cloud://openlist/Movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- media := model.Media{
- Base: model.Base{ID: "cloud-302"},
- LibraryID: lib.ID,
- Title: "Cloud 302 Movie",
- Path: `cloud://openlist/Movies/Movie.mkv`,
- STRMURL: `/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv`,
- Container: "mkv",
- }
- if err := svc.repo.DB.Create(&media).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
-
- pb, err := svc.PlaybackInfo(t.Context(), "cloud-302", "user-1")
- if err != nil {
- t.Fatalf("playback info: %v", err)
- }
- src := pb["MediaSources"].([]map[string]any)[0]
- if src["DirectStreamUrl"] != "/Videos/cloud-302/stream.mkv" {
- t.Fatalf("302/proxy mode should use Emby video stream URL: %#v", src)
- }
- if src["Path"] != "/Videos/cloud-302/stream.mkv" {
- t.Fatalf("302/proxy mode path should use Emby video stream URL: %#v", src)
- }
-}
-
-func TestEmbyPlaybackInfoProbesMissingCloudTrackMetadata(t *testing.T) {
- svc := newTestEmbyService(t)
- lib := model.Library{Name: "OpenList", Path: `cloud://openlist/Movies`, Type: "movie", Enabled: true}
- if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
- t.Fatalf("create library: %v", err)
- }
- media := model.Media{
- Base: model.Base{ID: "cloud-probe-1"},
- LibraryID: lib.ID,
- Title: "云盘电影",
- Path: `cloud://openlist/Movies/Movie.mkv`,
- STRMURL: `http://nas.local/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv`,
- }
- if err := svc.repo.DB.Create(&media).Error; err != nil {
- t.Fatalf("create media: %v", err)
- }
- resolver := &fakeCloudPlaybackResolver{
- link: &cloud.DirectLink{
- URL: "http://cdn.example.test/Movie.mkv",
- Headers: map[string]string{"Authorization": "Bearer probe-token"},
- },
- }
- prober := &fakeCloudPlaybackProber{
- probe: &ProbeResult{
- DurationSec: 3661,
- Width: 3840,
- Height: 2160,
- VideoCodec: "hevc",
- AudioCodec: "eac3",
- Container: "matroska,webm",
- },
- }
- svc.SetCloudProbe(resolver, prober)
-
- if _, err := svc.PlaybackInfo(t.Context(), "cloud-probe-1", "user-1"); err != nil {
- t.Fatalf("playback info: %v", err)
- }
-
- // 探测现在是异步的(同步探测曾把起播拖慢最多 8 秒并放大云盘流量)。
- // 轮询等待后台探测结果落库。
- var persisted model.Media
- deadline := time.Now().Add(3 * time.Second)
- for {
- if err := svc.repo.DB.First(&persisted, "id = ?", "cloud-probe-1").Error; err != nil {
- t.Fatalf("reload media: %v", err)
- }
- if persisted.DurationSec > 0 || time.Now().After(deadline) {
- break
- }
- time.Sleep(10 * time.Millisecond)
- }
- if persisted.DurationSec != 3661 || persisted.Width != 3840 || persisted.Height != 2160 || persisted.VideoCodec != "hevc" || persisted.AudioCodec != "eac3" {
- t.Fatalf("probe metadata not persisted: %#v", persisted)
- }
- if resolver.typ != "openlist" || resolver.ref != "/Movies/Movie.mkv" {
- t.Fatalf("resolver called with typ=%q ref=%q", resolver.typ, resolver.ref)
- }
- if prober.rawURL != "http://cdn.example.test/Movie.mkv" || prober.headers["Authorization"] != "Bearer probe-token" {
- t.Fatalf("probe called with url=%q headers=%#v", prober.rawURL, prober.headers)
- }
-
- // 落库之后,再次请求 PlaybackInfo 应当带上完整轨道元数据。
- pb, err := svc.PlaybackInfo(t.Context(), "cloud-probe-1", "user-1")
- if err != nil {
- t.Fatalf("playback info (second): %v", err)
- }
- src := pb["MediaSources"].([]map[string]any)[0]
- if src["RunTimeTicks"] != int64(3661)*10_000_000 {
- t.Fatalf("runtime ticks not filled after async probe: %#v", src)
- }
- streams := src["MediaStreams"].([]map[string]any)
- if len(streams) != 2 || streams[0]["Codec"] != "hevc" || streams[1]["Codec"] != "eac3" {
- t.Fatalf("media streams not filled after async probe: %#v", streams)
- }
-}
-
func newTestEmbyService(t *testing.T) *EmbyService {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Series{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}, &model.User{}, &model.Setting{})
// 内存库 + 异步探测协程:限制为单连接,避免连接池新建连接时
// 拿到一个空白的 :memory: 实例(no such table)。
if sqlDB, err := db.DB(); err == nil {
sqlDB.SetMaxOpenConns(1)
}
- if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}, &model.User{}, &model.Setting{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
repos := repository.New(db)
return NewEmbyService(&config.Config{}, zap.NewNop(), repos)
}
diff --git a/internal/service/emby_items_cache.go b/internal/service/emby_items_cache.go
new file mode 100644
index 0000000..6c7c3df
--- /dev/null
+++ b/internal/service/emby_items_cache.go
@@ -0,0 +1,55 @@
+package service
+
+import (
+ "crypto/sha256"
+ "encoding/hex"
+ "sort"
+ "strconv"
+ "strings"
+)
+
+type embyItemsCacheValue struct {
+ Items []map[string]any `json:"items"`
+ TotalRecordCount int64 `json:"total_record_count"`
+ StartIndex int `json:"start_index"`
+}
+
+type embyLatestCacheValue struct {
+ Items []map[string]any `json:"items"`
+}
+
+func (e *EmbyService) embyItemsCacheKey(kind string, p ItemsParams) string {
+ includeTypes := append([]string(nil), p.IncludeItemTypes...)
+ filters := append([]string(nil), p.Filters...)
+ ids := append([]string(nil), p.IDs...)
+ sort.Strings(includeTypes)
+ sort.Strings(filters)
+ sort.Strings(ids)
+ sum := sha256.Sum256([]byte(strings.Join([]string{
+ kind,
+ p.UserID,
+ p.ParentID,
+ strings.Join(ids, ","),
+ p.SearchTerm,
+ strings.Join(includeTypes, ","),
+ strings.Join(filters, ","),
+ strconv.FormatBool(p.Recursive),
+ p.SortBy,
+ p.SortOrder,
+ strconv.Itoa(p.StartIndex),
+ strconv.Itoa(p.Limit),
+ }, "|")))
+ return "media:emby:" + hex.EncodeToString(sum[:])
+}
+
+func (e *EmbyService) embyLatestCacheKey(userID, parentID string, limit int) string {
+ sum := sha256.Sum256([]byte(strings.Join([]string{"latest", userID, parentID, strconv.Itoa(limit)}, "|")))
+ return "media:emby:" + hex.EncodeToString(sum[:])
+}
+
+func (e *EmbyService) mediaCacheTTLSeconds() int {
+ if e == nil || e.cfg == nil || e.cfg.Cache.MediaTTLSeconds < 1 {
+ return 15
+ }
+ return e.cfg.Cache.MediaTTLSeconds
+}
diff --git a/internal/service/emby_items_detail.go b/internal/service/emby_items_detail.go
new file mode 100644
index 0000000..b5ddbd1
--- /dev/null
+++ b/internal/service/emby_items_detail.go
@@ -0,0 +1,275 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// Item 单条目详情。
+func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[string]any, error) {
+ if lib, err := e.repo.Library.FindByID(ctx, mediaID); err != nil {
+ return nil, err
+ } else if lib != nil {
+ libs := FilterDisplayCloudLibraries(ctx, e.repo, []model.Library{*lib})
+ if len(libs) == 0 {
+ return nil, nil
+ }
+ visibility := e.mediaVisibility(ctx, userID)
+ if !e.libraryVisibleFromCachedVisibility(libs[0], visibility) {
+ return nil, nil
+ }
+ return e.libraryAsView(&libs[0]), nil
+ }
+ if strings.HasPrefix(mediaID, embyVirtualSeasonPrefix) {
+ if season, ok, err := e.findSeasonGroup(ctx, mediaID, userID); err != nil {
+ return nil, err
+ } else if ok {
+ return e.seasonPayload(season), nil
+ }
+ }
+ if strings.HasPrefix(mediaID, embyVirtualSeriesPrefix) {
+ if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
+ return nil, err
+ } else if ok {
+ return e.seriesPayload(series), nil
+ }
+ }
+ m, err := e.repo.Media.FindByID(ctx, mediaID)
+ if err != nil {
+ return nil, err
+ }
+ if m == nil {
+ if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
+ return nil, err
+ } else if ok {
+ return e.seriesPayload(series), nil
+ }
+ return nil, nil
+ }
+ if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
+ return nil, nil
+ }
+ fav := false
+ pos := int64(0)
+ if userID != "" {
+ var f model.Favorite
+ ferr := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
+ if ferr == nil {
+ fav = true
+ }
+ var h model.PlaybackHistory
+ herr := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).
+ Order("watched_at desc").First(&h).Error
+ if herr == nil {
+ pos = h.PositionMs
+ }
+ }
+ return e.itemPayload(ctx, m, fav, pos), nil
+}
+
+// LatestItems 最近添加,全库或指定库。
+func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string, limit int) ([]map[string]any, error) {
+ if limit <= 0 || limit > 100 {
+ limit = 20
+ }
+ cacheKey := e.embyLatestCacheKey(userID, parentID, limit)
+ var cached embyLatestCacheValue
+ if e.cache != nil && e.cache.GetJSON(ctx, cacheKey, &cached) {
+ return cached.Items, nil
+ }
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("deleted_at IS NULL")
+ q = e.applyUserMediaVisibility(ctx, q, userID)
+ if parentID != "" {
+ if episodic, err := e.libraryIsEpisodic(ctx, parentID); err == nil && episodic {
+ out, err := e.latestSeriesItemsForLibrary(ctx, userID, parentID, limit)
+ if err == nil && e.cache != nil {
+ e.cache.SetJSON(ctx, cacheKey, embyLatestCacheValue{Items: out}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
+ }
+ return out, err
+ }
+ q = q.Where("library_id IN ?", e.mergedLibraryIDs(ctx, parentID))
+ }
+ var rows []model.Media
+ if err := q.Order("media.created_at desc").Limit(limit).Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ favs := map[string]bool{}
+ if userID != "" && len(rows) > 0 {
+ mediaIDs := make([]string, 0, len(rows))
+ for _, row := range rows {
+ if strings.TrimSpace(row.ID) != "" {
+ mediaIDs = append(mediaIDs, row.ID)
+ }
+ }
+ if len(mediaIDs) == 0 {
+ mediaIDs = []string{"__none__"}
+ }
+ var fr []model.Favorite
+ _ = e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id IN ?", userID, mediaIDs).Find(&fr).Error
+ for _, f := range fr {
+ favs[f.MediaID] = true
+ }
+ }
+ out := make([]map[string]any, 0, len(rows))
+ for _, m := range rows {
+ out = append(out, e.itemPayload(ctx, &m, favs[m.ID], 0))
+ }
+ if e.cache != nil {
+ e.cache.SetJSON(ctx, cacheKey, embyLatestCacheValue{Items: out}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
+ }
+ return out, nil
+}
+
+func (e *EmbyService) latestSeriesItemsForLibrary(ctx context.Context, userID, libraryID string, limit int) ([]map[string]any, error) {
+ if limit <= 0 || limit > 100 {
+ limit = 20
+ }
+ rowLimit := limit * 40
+ if rowLimit < 200 {
+ rowLimit = 200
+ }
+ if rowLimit > embySeriesGroupingLimit {
+ rowLimit = embySeriesGroupingLimit
+ }
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("library_id IN ? AND (season_num > 0 OR episode_num > 0)", e.mergedLibraryIDs(ctx, libraryID))
+ q = e.applyUserMediaVisibility(ctx, q, userID)
+ var rows []model.Media
+ if err := q.Order("media.created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ groups := e.seriesGroupsFromMedia(rows)
+ sortSeriesGroups(groups, ItemsParams{SortBy: "datecreated", SortOrder: "Descending"})
+ if len(groups) > limit {
+ groups = groups[:limit]
+ }
+ items := make([]map[string]any, 0, len(groups))
+ for _, group := range groups {
+ items = append(items, e.seriesPayload(group))
+ }
+ return items, nil
+}
+
+// ResumeItems 列出有未完成播放进度的媒体。
+func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int) (map[string]any, error) {
+ if limit <= 0 || limit > 100 {
+ limit = 20
+ }
+ var hist []model.PlaybackHistory
+ if err := e.repo.DB.WithContext(ctx).
+ Where("user_id = ? AND completed = ? AND position_ms > 0", userID, false).
+ Order("watched_at desc").Limit(limit).Find(&hist).Error; err != nil {
+ return nil, err
+ }
+ if len(hist) == 0 {
+ return map[string]any{"Items": []any{}, "TotalRecordCount": 0}, nil
+ }
+ ids := make([]string, 0, len(hist))
+ posByID := map[string]int64{}
+ for _, h := range hist {
+ ids = append(ids, h.MediaID)
+ posByID[h.MediaID] = h.PositionMs
+ }
+ var medias []model.Media
+ q := e.repo.DB.WithContext(ctx).Where("id IN ?", ids)
+ q = e.applyUserMediaVisibility(ctx, q, userID)
+ if err := q.Find(&medias).Error; err != nil {
+ return nil, err
+ }
+ byID := map[string]*model.Media{}
+ for i := range medias {
+ byID[medias[i].ID] = &medias[i]
+ }
+ items := make([]map[string]any, 0, len(hist))
+ for _, h := range hist {
+ if m, ok := byID[h.MediaID]; ok {
+ items = append(items, e.itemPayload(ctx, m, false, posByID[h.MediaID]))
+ }
+ }
+ return map[string]any{"Items": items, "TotalRecordCount": len(items)}, nil
+}
+
+func (e *EmbyService) itemPayload(ctx context.Context, m *model.Media, fav bool, posMs int64) map[string]any {
+ itemType := "Movie"
+ name := m.Title
+ parentID := m.LibraryID
+ seriesID := m.SeriesID
+ seriesName := ""
+ seasonID := ""
+ if e.mediaShouldBeEpisode(ctx, m) {
+ itemType = "Episode"
+ seriesID = e.seriesIDForMedia(m)
+ seriesName = e.seriesNameForMedia(m)
+ seasonID = e.seasonIDForMedia(m)
+ parentID = seasonID
+ episodeTitle := strings.TrimSpace(m.EpisodeTitle)
+ if episodeTitle != "" {
+ name = episodeTitle
+ } else if m.EpisodeNum > 0 {
+ name = fmt.Sprintf("第 %d 集", m.EpisodeNum)
+ }
+ }
+ imageTags := map[string]string{}
+ backdropTags := []string{}
+ primaryArtwork := e.mediaPrimaryArtwork(ctx, m)
+ backdropArtwork := e.mediaBackdropArtwork(ctx, m)
+ if primaryArtwork != "" {
+ imageTags["Primary"] = m.ID
+ }
+ if backdropArtwork != "" {
+ backdropTags = append(backdropTags, m.ID+"-bd")
+ }
+
+ runTimeTicks := int64(m.DurationSec) * 10_000_000
+ durationMs := int64(m.DurationSec) * 1000
+ played := posMs > 0 && durationMs > 0 && posMs >= durationMs*9/10
+ pct := 0.0
+ if durationMs > 0 {
+ pct = float64(posMs) / float64(durationMs) * 100
+ }
+
+ return map[string]any{
+ "Id": m.ID,
+ "Name": name,
+ "OriginalTitle": m.OriginalName,
+ "ServerId": embyServerID,
+ "Type": itemType,
+ "MediaType": "Video",
+ "IsFolder": false,
+ "ProductionYear": m.Year,
+ "ParentIndexNumber": m.SeasonNum,
+ "IndexNumber": m.EpisodeNum,
+ "Overview": m.Overview,
+ "RunTimeTicks": runTimeTicks,
+ "CommunityRating": m.Rating,
+ "Container": m.Container,
+ "Width": m.Width,
+ "Height": m.Height,
+ "DateCreated": m.CreatedAt,
+ "Path": m.Path,
+ "ParentId": parentID,
+ "SeasonId": seasonID,
+ "SeasonName": seasonName(m.SeasonNum),
+ "SeriesId": seriesID,
+ "SeriesName": seriesName,
+ "ImageTags": imageTags,
+ "BackdropImageTags": backdropTags,
+ "Genres": splitCSV(m.Genres),
+ "ProviderIds": map[string]string{
+ "Tmdb": intToStr(m.TMDbID),
+ "Bangumi": intToStr(m.BangumiID),
+ },
+ "UserData": map[string]any{
+ "PlaybackPositionTicks": posMs * 10_000,
+ "PlayCount": 0,
+ "IsFavorite": fav,
+ "Played": played,
+ "PlayedPercentage": pct,
+ },
+ "MediaSources": e.mediaSourcesForItem(ctx, m, true, false),
+ }
+}
diff --git a/internal/service/emby_items_helpers.go b/internal/service/emby_items_helpers.go
new file mode 100644
index 0000000..03ddbad
--- /dev/null
+++ b/internal/service/emby_items_helpers.go
@@ -0,0 +1,116 @@
+package service
+
+import "strings"
+
+func containsItemType(types []string, want string) bool {
+ for _, t := range types {
+ if strings.EqualFold(strings.TrimSpace(t), want) {
+ return true
+ }
+ }
+ return false
+}
+
+func containsSupportedEmbyItemType(types []string) bool {
+ for _, itemType := range types {
+ switch strings.ToLower(strings.TrimSpace(itemType)) {
+ case "movie", "series", "season", "episode", "video", "folder", "collectionfolder":
+ return true
+ }
+ }
+ return false
+}
+
+func containsOnlyFolderItemTypes(types []string) bool {
+ if len(types) == 0 {
+ return false
+ }
+ for _, itemType := range types {
+ switch strings.ToLower(strings.TrimSpace(itemType)) {
+ case "folder", "collectionfolder":
+ default:
+ return false
+ }
+ }
+ return true
+}
+
+func emptyItemsEnvelope(startIndex int) map[string]any {
+ return map[string]any{
+ "Items": []map[string]any{},
+ "TotalRecordCount": int64(0),
+ "StartIndex": startIndex,
+ }
+}
+
+func containsEmbyFilter(filters []string, want string) bool {
+ for _, filter := range filters {
+ if strings.EqualFold(strings.TrimSpace(filter), want) {
+ return true
+ }
+ }
+ return false
+}
+
+func firstCSVValue(value string) string {
+ if i := strings.Index(value, ","); i >= 0 {
+ value = value[:i]
+ }
+ return strings.TrimSpace(value)
+}
+
+func primarySupportedEmbySort(sortBy string, resumeFilter bool) string {
+ for _, part := range strings.Split(sortBy, ",") {
+ key := strings.ToLower(strings.TrimSpace(part))
+ switch key {
+ case "sortname", "name", "premieredate", "productionyear", "datecreated", "communityrating":
+ return key
+ case "dateplayed":
+ if resumeFilter {
+ return key
+ }
+ }
+ }
+ return strings.ToLower(strings.TrimSpace(firstCSVValue(sortBy)))
+}
+
+func pageSlice[T any](items []T, start, limit int) []T {
+ if start < 0 {
+ start = 0
+ }
+ if limit <= 0 {
+ limit = len(items)
+ }
+ if start >= len(items) {
+ return []T{}
+ }
+ end := start + limit
+ if end > len(items) {
+ end = len(items)
+ }
+ return items[start:end]
+}
+
+func emptyUserData() map[string]any {
+ return map[string]any{
+ "PlaybackPositionTicks": 0,
+ "PlayCount": 0,
+ "IsFavorite": false,
+ "Played": false,
+ "PlayedPercentage": 0,
+ }
+}
+
+func minInt(a, b int) int {
+ if a < b {
+ return a
+ }
+ return b
+}
+
+func maxInt(a, b int) int {
+ if a > b {
+ return a
+ }
+ return b
+}
diff --git a/internal/service/emby_items_list.go b/internal/service/emby_items_list.go
new file mode 100644
index 0000000..d819a45
--- /dev/null
+++ b/internal/service/emby_items_list.go
@@ -0,0 +1,252 @@
+package service
+
+import (
+ "context"
+ "sort"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
+ cacheKey := e.embyItemsCacheKey("items", p)
+ var cached embyItemsCacheValue
+ if e.cache != nil && e.cache.GetJSON(ctx, cacheKey, &cached) {
+ return map[string]any{"Items": cached.Items, "TotalRecordCount": cached.TotalRecordCount, "StartIndex": cached.StartIndex}, nil
+ }
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{})
+ q = e.applyUserMediaVisibility(ctx, q, p.UserID)
+ if p.ParentID != "" {
+ q = q.Where("library_id IN ? OR series_id = ?", e.mergedLibraryIDs(ctx, p.ParentID), p.ParentID)
+ }
+ if p.SearchTerm != "" {
+ q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
+ }
+ if containsEmbyFilter(p.Filters, "IsFavorite") {
+ if strings.TrimSpace(p.UserID) == "" {
+ return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
+ }
+ q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
+ }
+ resumeFilter := containsEmbyFilter(p.Filters, "IsResumable")
+ if resumeFilter {
+ if strings.TrimSpace(p.UserID) == "" {
+ return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
+ }
+ q = q.Joins(`JOIN (
+ SELECT media_id, MAX(watched_at) AS watched_at
+ FROM playback_histories
+ WHERE user_id = ? AND completed = ? AND position_ms > 0
+ GROUP BY media_id
+ ) AS resume ON resume.media_id = media.id`, p.UserID, false)
+ }
+ filterBySeasonNumbers := true
+ parentKnownNonEpisodic := false
+ if p.ParentID != "" {
+ if episodic, err := e.libraryIsEpisodic(ctx, p.ParentID); err == nil && !episodic {
+ filterBySeasonNumbers = false
+ parentKnownNonEpisodic = true
+ }
+ }
+ if parentKnownNonEpisodic && containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
+ return emptyItemsEnvelope(p.StartIndex), nil
+ }
+ if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Movie") && !containsItemType(p.IncludeItemTypes, "Episode") {
+ q = e.filterMovieItems(ctx, q)
+ }
+ if parentKnownNonEpisodic && containsItemType(p.IncludeItemTypes, "Movie") && !containsItemType(p.IncludeItemTypes, "Episode") {
+ q = filterLikelyEpisodicPathsFromMovieQuery(q)
+ }
+ if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
+ q = e.filterEpisodeItems(ctx, q)
+ }
+
+ var total int64
+ if err := q.Count(&total).Error; err != nil {
+ return nil, err
+ }
+ order := "media.created_at desc"
+ switch primarySupportedEmbySort(p.SortBy, resumeFilter) {
+ case "sortname", "name":
+ order = "media.title"
+ case "premieredate", "productionyear":
+ order = "media.year"
+ case "datecreated":
+ order = "media.created_at"
+ case "dateplayed":
+ order = "resume.watched_at"
+ case "communityrating":
+ order = "media.rating"
+ }
+ if strings.EqualFold(firstCSVValue(p.SortOrder), "Descending") {
+ if !strings.HasSuffix(order, " desc") {
+ order = order + " desc"
+ }
+ }
+
+ fetchLimit := p.Limit
+ fetchOffset := p.StartIndex
+ if fetchLimit > 0 && e.shouldCollapseMediaVersions(ctx, p) {
+ // Duplicates across merged local/cloud libraries collapse into one Emby
+ // item with multiple MediaSources. Fetch a wider window so duplicates do
+ // not consume the whole requested page.
+ fetchOffset = 0
+ fetchLimit = p.StartIndex + maxInt(p.Limit*4, p.Limit)
+ }
+ var rows []model.Media
+ if err := q.Order(order).Offset(fetchOffset).Limit(fetchLimit).Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ if e.shouldCollapseMediaVersions(ctx, p) {
+ rows = e.collapseMediaVersionRows(ctx, rows)
+ rows = pageSlice(rows, p.StartIndex, p.Limit)
+ }
+ items, err := e.payloadsForMedia(ctx, rows, p.UserID)
+ if err != nil {
+ return nil, err
+ }
+ out := map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}
+ if e.cache != nil {
+ e.cache.SetJSON(ctx, cacheKey, embyItemsCacheValue{Items: items, TotalRecordCount: total, StartIndex: p.StartIndex}, time.Duration(e.mediaCacheTTLSeconds())*time.Second)
+ }
+ return out, nil
+}
+
+func (e *EmbyService) episodeItems(ctx context.Context, rows []model.Media, p ItemsParams) (map[string]any, error) {
+ rows = e.filterMediaRowsForUser(ctx, rows, p.UserID)
+ if p.SearchTerm != "" {
+ filtered := rows[:0]
+ needle := strings.ToLower(p.SearchTerm)
+ for _, row := range rows {
+ if strings.Contains(strings.ToLower(row.Title), needle) || strings.Contains(strings.ToLower(row.OriginalName), needle) {
+ filtered = append(filtered, row)
+ }
+ }
+ rows = filtered
+ }
+ sort.SliceStable(rows, func(i, j int) bool {
+ if rows[i].SeasonNum != rows[j].SeasonNum {
+ return rows[i].SeasonNum < rows[j].SeasonNum
+ }
+ if rows[i].EpisodeNum != rows[j].EpisodeNum {
+ return rows[i].EpisodeNum < rows[j].EpisodeNum
+ }
+ return rows[i].CreatedAt.Before(rows[j].CreatedAt)
+ })
+ total := len(rows)
+ items, err := e.payloadsForMedia(ctx, pageSlice(rows, p.StartIndex, p.Limit), p.UserID)
+ if err != nil {
+ return nil, err
+ }
+ return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
+}
+
+func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media, userID string) ([]map[string]any, error) {
+ rows = e.collapseMediaVersionRows(ctx, rows)
+ userFavs := map[string]bool{}
+ userPos := map[string]int64{}
+ if userID != "" && len(rows) > 0 {
+ mediaIDs := make([]string, 0, len(rows))
+ for _, row := range rows {
+ if strings.TrimSpace(row.ID) != "" {
+ mediaIDs = append(mediaIDs, row.ID)
+ }
+ }
+ if len(mediaIDs) == 0 {
+ mediaIDs = []string{"__none__"}
+ }
+ var favs []model.Favorite
+ favQuery := e.repo.DB.WithContext(ctx).Where("user_id = ?", userID).Where("media_id IN ?", mediaIDs)
+ _ = favQuery.Find(&favs).Error
+ for _, f := range favs {
+ userFavs[f.MediaID] = true
+ }
+ var hist []model.PlaybackHistory
+ histQuery := e.repo.DB.WithContext(ctx).Where("user_id = ?", userID).Where("media_id IN ?", mediaIDs)
+ _ = histQuery.Find(&hist).Error
+ for _, h := range hist {
+ userPos[h.MediaID] = h.PositionMs
+ }
+ }
+
+ items := make([]map[string]any, 0, len(rows))
+ for _, m := range rows {
+ items = append(items, e.itemPayload(ctx, &m, userFavs[m.ID], userPos[m.ID]))
+ }
+ return items, nil
+}
+
+func (e *EmbyService) shouldCollapseMediaVersions(ctx context.Context, p ItemsParams) bool {
+ if containsItemType(p.IncludeItemTypes, "Series") || containsItemType(p.IncludeItemTypes, "Season") {
+ return false
+ }
+ if containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
+ return true
+ }
+ if p.ParentID == "" {
+ return true
+ }
+ episodic, err := e.libraryIsEpisodic(ctx, p.ParentID)
+ return err == nil && !episodic
+}
+
+func (e *EmbyService) collapseMediaVersionRows(ctx context.Context, rows []model.Media) []model.Media {
+ if len(rows) < 2 {
+ return rows
+ }
+ out := make([]model.Media, 0, len(rows))
+ indexByKey := make(map[string]int, len(rows))
+ for _, row := range rows {
+ key := e.mediaVersionKey(ctx, &row)
+ if key == "" {
+ out = append(out, row)
+ continue
+ }
+ if idx, ok := indexByKey[key]; ok {
+ if preferMediaVersion(row, out[idx]) {
+ out[idx] = row
+ }
+ continue
+ }
+ indexByKey[key] = len(out)
+ out = append(out, row)
+ }
+ return out
+}
+
+func (e *EmbyService) seriesItemsForLibrary(ctx context.Context, libraryID string, p ItemsParams) (map[string]any, error) {
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
+ q = e.applyUserMediaVisibility(ctx, q, p.UserID)
+ if libraryID != "" {
+ q = q.Where("library_id IN ?", e.mergedLibraryIDs(ctx, libraryID))
+ }
+ if p.SearchTerm != "" {
+ q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
+ }
+ if containsEmbyFilter(p.Filters, "IsFavorite") {
+ if strings.TrimSpace(p.UserID) == "" {
+ return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
+ }
+ q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
+ }
+ rowLimit := p.StartIndex + maxInt(p.Limit*40, 1000)
+ if rowLimit < p.Limit {
+ rowLimit = p.Limit
+ }
+ if rowLimit > embySeriesGroupingLimit {
+ rowLimit = embySeriesGroupingLimit
+ }
+ var rows []model.Media
+ if err := q.Order("media.created_at desc").Limit(rowLimit).Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ groups := e.seriesGroupsFromMedia(rows)
+ sortSeriesGroups(groups, p)
+ total := len(groups)
+ items := make([]map[string]any, 0, minInt(p.Limit, len(groups)))
+ for _, group := range pageSlice(groups, p.StartIndex, p.Limit) {
+ items = append(items, e.seriesPayload(group))
+ }
+ return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
+}
diff --git a/internal/service/emby_media_sources.go b/internal/service/emby_media_sources.go
new file mode 100644
index 0000000..14f9906
--- /dev/null
+++ b/internal/service/emby_media_sources.go
@@ -0,0 +1,182 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "sort"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (e *EmbyService) mediaSourcesForItem(ctx context.Context, m *model.Media, asEmbedded, directOnly bool) []map[string]any {
+ siblings := e.mediaVersionSiblings(ctx, m)
+ if len(siblings) == 0 {
+ return []map[string]any{e.mediaSource(ctx, m, asEmbedded, directOnly)}
+ }
+ sources := make([]map[string]any, 0, len(siblings))
+ for i := range siblings {
+ media := siblings[i]
+ sources = append(sources, e.mediaSource(ctx, &media, asEmbedded, directOnly))
+ }
+ return sources
+}
+
+func (e *EmbyService) mediaVersionSiblings(ctx context.Context, m *model.Media) []model.Media {
+ if e == nil || e.repo == nil || e.repo.DB == nil || m == nil || strings.TrimSpace(m.ID) == "" {
+ return nil
+ }
+ libraryIDs := e.mergedLibraryIDs(ctx, m.LibraryID)
+ if len(libraryIDs) == 0 {
+ libraryIDs = []string{m.LibraryID}
+ }
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("library_id IN ?", libraryIDs).
+ Where("season_num = ? AND episode_num = ?", m.SeasonNum, m.EpisodeNum)
+ if m.TMDbID > 0 {
+ q = q.Where("tm_db_id = ?", m.TMDbID)
+ } else if m.BangumiID > 0 {
+ q = q.Where("bangumi_id = ?", m.BangumiID)
+ } else {
+ title := strings.TrimSpace(m.Title)
+ if title == "" {
+ title = strings.TrimSpace(m.OriginalName)
+ }
+ if title == "" {
+ return []model.Media{*m}
+ }
+ q = q.Where("LOWER(title) = ?", strings.ToLower(title))
+ if m.Year > 0 {
+ q = q.Where("year = ?", m.Year)
+ }
+ }
+ var rows []model.Media
+ if err := q.Find(&rows).Error; err != nil || len(rows) == 0 {
+ return []model.Media{*m}
+ }
+ rows = e.collapseExactPathRows(rows)
+ sort.SliceStable(rows, func(i, j int) bool {
+ if rows[i].ID == m.ID {
+ return true
+ }
+ if rows[j].ID == m.ID {
+ return false
+ }
+ return preferMediaVersion(rows[i], rows[j])
+ })
+ return rows
+}
+
+func (e *EmbyService) collapseExactPathRows(rows []model.Media) []model.Media {
+ if len(rows) < 2 {
+ return rows
+ }
+ out := rows[:0]
+ seen := map[string]struct{}{}
+ for _, row := range rows {
+ path := strings.TrimSpace(row.Path)
+ if path != "" {
+ if _, ok := seen[path]; ok {
+ continue
+ }
+ seen[path] = struct{}{}
+ }
+ out = append(out, row)
+ }
+ return out
+}
+
+func (e *EmbyService) mediaVersionKey(ctx context.Context, m *model.Media) string {
+ if e == nil || m == nil {
+ return ""
+ }
+ ids := e.mergedLibraryIDs(ctx, m.LibraryID)
+ sort.Strings(ids)
+ libraryGroup := strings.Join(ids, ",")
+ if libraryGroup == "" {
+ libraryGroup = strings.TrimSpace(m.LibraryID)
+ }
+ if m.TMDbID > 0 {
+ return fmt.Sprintf("%s|tmdb:%d|s:%d|e:%d", libraryGroup, m.TMDbID, m.SeasonNum, m.EpisodeNum)
+ }
+ if m.BangumiID > 0 {
+ return fmt.Sprintf("%s|bangumi:%d|s:%d|e:%d", libraryGroup, m.BangumiID, m.SeasonNum, m.EpisodeNum)
+ }
+ title := strings.ToLower(strings.TrimSpace(m.Title))
+ if title == "" {
+ title = strings.ToLower(strings.TrimSpace(m.OriginalName))
+ }
+ if title == "" {
+ return ""
+ }
+ return fmt.Sprintf("%s|title:%s|y:%d|s:%d|e:%d", libraryGroup, title, m.Year, m.SeasonNum, m.EpisodeNum)
+}
+
+func preferMediaVersion(candidate, current model.Media) bool {
+ candidateCloud := strings.TrimSpace(candidate.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(candidate.Path)), "cloud://")
+ currentCloud := strings.TrimSpace(current.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(current.Path)), "cloud://")
+ if candidateCloud != currentCloud {
+ return !candidateCloud
+ }
+ if candidate.Width != current.Width {
+ return candidate.Width > current.Width
+ }
+ if candidate.SizeBytes != current.SizeBytes {
+ return candidate.SizeBytes > current.SizeBytes
+ }
+ return candidate.CreatedAt.After(current.CreatedAt)
+}
+
+func embySTRMStreamURL(mediaID string) string {
+ return "/api/stream/" + url.PathEscape(strings.TrimSpace(mediaID))
+}
+
+func embyDirectStreamURL(mediaID, container string) string {
+ mediaID = strings.TrimSpace(mediaID)
+ container = strings.Trim(strings.ToLower(container), ". ")
+ if container == "" || container == "strm" {
+ return "/Videos/" + mediaID + "/stream"
+ }
+ return "/Videos/" + mediaID + "/stream." + container
+}
+
+func (e *EmbyService) mediaStreams(m *model.Media) []map[string]any {
+ streams := []map[string]any{}
+ if m.VideoCodec != "" || m.Width > 0 {
+ streams = append(streams, map[string]any{
+ "Codec": m.VideoCodec,
+ "Type": "Video",
+ "Index": 0,
+ "Width": m.Width,
+ "Height": m.Height,
+ "AspectRatio": "",
+ "IsDefault": true,
+ "IsForced": false,
+ "IsExternal": false,
+ "DisplayTitle": fmt.Sprintf("%dx%d %s", m.Width, m.Height, m.VideoCodec),
+ })
+ }
+ if m.AudioCodec != "" {
+ streams = append(streams, map[string]any{
+ "Codec": m.AudioCodec,
+ "Type": "Audio",
+ "Index": 1,
+ "IsDefault": true,
+ "IsForced": false,
+ "IsExternal": false,
+ })
+ }
+ if len(streams) == 0 {
+ streams = append(streams, map[string]any{
+ "Codec": "unknown",
+ "Type": "Video",
+ "Index": 0,
+ "IsDefault": true,
+ "IsForced": false,
+ "IsExternal": false,
+ "DisplayTitle": "Video",
+ })
+ }
+ return streams
+}
diff --git a/internal/service/emby_movie_items.go b/internal/service/emby_movie_items.go
new file mode 100644
index 0000000..31cc461
--- /dev/null
+++ b/internal/service/emby_movie_items.go
@@ -0,0 +1,252 @@
+package service
+
+import (
+ "context"
+ "sort"
+ "strings"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// movieLibraryHasEpisodicContent 报告电影类型库里是否混入了「剧集结构」内容
+// (有季集号且路径形如剧集,例如 .../国产剧/某剧/Season 01/某剧 - S01E01.mkv)。
+// 用于决定是否需要走 movieLibraryItems 把这些内容聚成 Series 卡片。普通电影库
+// 没有这类行时返回 false,继续走常规 mediaItems。
+func (e *EmbyService) movieLibraryHasEpisodicContent(ctx context.Context, libraryID string) (bool, error) {
+ clause, args := embyLikelyEpisodicPathSQL()
+ if clause == "" {
+ return false, nil
+ }
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("library_id IN ?", e.mergedLibraryIDs(ctx, libraryID)).
+ Where("(season_num > 0 OR episode_num > 0) AND ("+clause+")", args...)
+ var count int64
+ if err := q.Limit(1).Count(&count).Error; err != nil {
+ return false, err
+ }
+ return count > 0, nil
+}
+
+// movieLibraryItems 处理电影类型库的常规浏览,返回「真正的电影(Movie)」与
+// 「库内剧集结构内容聚成的 Series 卡片」的合并列表(按 DateCreated 倒序分页)。
+// 与 mediaItems 的区别: 后者会把剧集结构行当散装 Episode 漏出;这里改为聚合成
+// Series,从根本上消除「电影库里整部剧被拆成单集」的现象。
+func (e *EmbyService) movieLibraryItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
+ libIDs := e.mergedLibraryIDs(ctx, p.ParentID)
+ apply := func(q *gorm.DB) *gorm.DB {
+ q = e.applyUserMediaVisibility(ctx, q, p.UserID)
+ q = q.Where("library_id IN ?", libIDs)
+ if p.SearchTerm != "" {
+ q = q.Where("title LIKE ? OR original_name LIKE ?", "%"+p.SearchTerm+"%", "%"+p.SearchTerm+"%")
+ }
+ if containsEmbyFilter(p.Filters, "IsFavorite") {
+ if strings.TrimSpace(p.UserID) == "" {
+ return nil
+ }
+ q = q.Joins("JOIN favorites ON favorites.media_id = media.id AND favorites.user_id = ? AND favorites.deleted_at IS NULL", p.UserID)
+ }
+ return q
+ }
+
+ // 剧集结构内容 -> Series 卡片。
+ clause, args := embyLikelyEpisodicPathSQL()
+ var episodicRows []model.Media
+ if clause != "" {
+ epQ := apply(e.repo.DB.WithContext(ctx).Model(&model.Media{}))
+ if epQ == nil {
+ return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
+ }
+ epQ = epQ.Where("(season_num > 0 OR episode_num > 0) AND ("+clause+")", args...).
+ Order("media.created_at desc").Limit(embySeriesGroupingLimit)
+ if err := epQ.Find(&episodicRows).Error; err != nil {
+ return nil, err
+ }
+ }
+ seriesGroups := e.seriesGroupsFromMedia(episodicRows)
+
+ // 真正的电影 -> Movie 项(剔除剧集结构行)。
+ movieQ := apply(e.repo.DB.WithContext(ctx).Model(&model.Media{}))
+ if movieQ == nil {
+ return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": 0, "StartIndex": p.StartIndex}, nil
+ }
+ movieQ = filterLikelyEpisodicPathsFromMovieQuery(movieQ).
+ Order("media.created_at desc").Limit(embySeriesGroupingLimit)
+ var movieRows []model.Media
+ if err := movieQ.Find(&movieRows).Error; err != nil {
+ return nil, err
+ }
+ movieItems, err := e.payloadsForMedia(ctx, movieRows, p.UserID)
+ if err != nil {
+ return nil, err
+ }
+
+ // 合并: Series 卡片 + Movie 项, 统一按 DateCreated 倒序。
+ type entry struct {
+ createdAt time.Time
+ payload map[string]any
+ }
+ entries := make([]entry, 0, len(seriesGroups)+len(movieItems))
+ for _, g := range seriesGroups {
+ entries = append(entries, entry{createdAt: g.CreatedAt, payload: e.seriesPayload(g)})
+ }
+ for _, item := range movieItems {
+ entries = append(entries, entry{createdAt: embyPayloadCreatedAt(item), payload: item})
+ }
+ sort.SliceStable(entries, func(i, j int) bool {
+ return entries[i].createdAt.After(entries[j].createdAt)
+ })
+ total := len(entries)
+ paged := pageSlice(entries, p.StartIndex, p.Limit)
+ items := make([]map[string]any, 0, len(paged))
+ for _, en := range paged {
+ items = append(items, en.payload)
+ }
+ return map[string]any{"Items": items, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
+}
+
+// embyPayloadCreatedAt 从 item payload 里取 DateCreated(time.Time),用于合并排序。
+func embyPayloadCreatedAt(item map[string]any) time.Time {
+ if item == nil {
+ return time.Time{}
+ }
+ if v, ok := item["DateCreated"].(time.Time); ok {
+ return v
+ }
+ return time.Time{}
+}
+
+func (e *EmbyService) libraryIsEpisodic(ctx context.Context, libraryID string) (bool, error) {
+ if strings.TrimSpace(libraryID) == "" {
+ return false, nil
+ }
+ if lib, err := e.repo.Library.FindByID(ctx, libraryID); err != nil {
+ return false, err
+ } else if lib != nil {
+ return embyLibraryTypeIsEpisodic(lib.Type), nil
+ }
+ var count int64
+ err := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("library_id IN ? AND (season_num > 0 OR episode_num > 0)", e.mergedLibraryIDs(ctx, libraryID)).
+ Count(&count).Error
+ return count > 0, err
+}
+
+func (e *EmbyService) mediaBelongsToEpisodicLibrary(ctx context.Context, m *model.Media) bool {
+ if e == nil || m == nil || strings.TrimSpace(m.LibraryID) == "" {
+ return false
+ }
+ lib, err := e.repo.Library.FindByID(ctx, m.LibraryID)
+ if err != nil || lib == nil {
+ return false
+ }
+ return embyLibraryTypeIsEpisodic(lib.Type)
+}
+
+func (e *EmbyService) mediaShouldBeEpisode(ctx context.Context, m *model.Media) bool {
+ if m == nil || (m.SeasonNum <= 0 && m.EpisodeNum <= 0) {
+ return false
+ }
+ if e.mediaBelongsToEpisodicLibrary(ctx, m) {
+ return true
+ }
+ return embyMediaPathLooksEpisodic(m.Path)
+}
+
+func embyLibraryTypeIsEpisodic(typ string) bool {
+ switch strings.ToLower(strings.TrimSpace(typ)) {
+ case "tv", "anime", "variety":
+ return true
+ default:
+ return false
+ }
+}
+
+func (e *EmbyService) filterMovieItems(ctx context.Context, q *gorm.DB) *gorm.DB {
+ episodicIDs := e.episodicLibraryIDs(ctx)
+ if len(episodicIDs) == 0 {
+ return filterLikelyEpisodicPathsFromMovieQuery(q)
+ }
+ q = q.Where("(media.season_num = 0 AND media.episode_num = 0) OR media.library_id NOT IN ?", episodicIDs)
+ return filterLikelyEpisodicPathsFromMovieQuery(q)
+}
+
+func (e *EmbyService) filterEpisodeItems(ctx context.Context, q *gorm.DB) *gorm.DB {
+ episodicIDs := e.episodicLibraryIDs(ctx)
+ if len(episodicIDs) == 0 {
+ return q.Where("1 = 0")
+ }
+ return q.Where("media.library_id IN ? AND (media.season_num > 0 OR media.episode_num > 0)", episodicIDs)
+}
+
+func (e *EmbyService) episodicLibraryIDs(ctx context.Context) []string {
+ if e == nil || e.repo == nil || e.repo.DB == nil {
+ return nil
+ }
+ var ids []string
+ if err := e.repo.DB.WithContext(ctx).Model(&model.Library{}).
+ Where("LOWER(type) IN ?", []string{"tv", "anime", "variety"}).
+ Pluck("id", &ids).Error; err != nil {
+ return nil
+ }
+ return ids
+}
+
+func filterLikelyEpisodicPathsFromMovieQuery(q *gorm.DB) *gorm.DB {
+ clause, args := embyLikelyEpisodicPathSQL()
+ if clause == "" {
+ return q
+ }
+ return q.Where("NOT ((media.season_num > 0 OR media.episode_num > 0) AND ("+clause+"))", args...)
+}
+
+func embyLikelyEpisodicPathSQL() (string, []any) {
+ patterns := []string{
+ "%/season %/%", "%/season.%/%", "%/season-%/%", "%/season_%/%",
+ "%/s0%/%", "%/s1%/%", "%/s2%/%", "%/s3%/%", "%/s4%/%", "%/s5%/%", "%/s6%/%", "%/s7%/%", "%/s8%/%", "%/s9%/%",
+ "%/special/%", "%/specials/%", "%/sp/%", "%/ova/%", "%/oad/%", "%/extra/%", "%/extras/%",
+ "%/电视剧/%", "%/剧集/%", "%/国产剧/%", "%/欧美剧/%", "%/日韩剧/%", "%/日剧/%", "%/韩剧/%",
+ "%/日番/%", "%/国漫/%", "%/番剧/%", "%/动漫/%", "%/特别篇/%", "%/特別篇/%", "%/番外/%", "%/特典/%",
+ }
+ clauses := make([]string, 0, len(patterns)*2)
+ args := make([]any, 0, len(patterns)*2)
+ for _, pattern := range patterns {
+ clauses = append(clauses, "LOWER(media.path) LIKE ?")
+ args = append(args, pattern)
+ if strings.Contains(pattern, "/") {
+ clauses = append(clauses, "LOWER(media.path) LIKE ?")
+ args = append(args, strings.ReplaceAll(pattern, "/", `\`))
+ }
+ }
+ return strings.Join(clauses, " OR "), args
+}
+
+func embyMediaPathLooksEpisodic(path string) bool {
+ normalized := strings.ToLower(strings.ReplaceAll(strings.TrimSpace(path), "\\", "/"))
+ if normalized == "" {
+ return false
+ }
+ for _, marker := range []string{
+ "/season ", "/season.", "/season-", "/season_", "/special/", "/specials/", "/sp/", "/ova/", "/oad/", "/extra/", "/extras/",
+ "/电视剧/", "/剧集/", "/国产剧/", "/欧美剧/", "/日韩剧/", "/日剧/", "/韩剧/",
+ "/日番/", "/国漫/", "/番剧/", "/动漫/", "/特别篇/", "/特別篇/", "/番外/", "/特典/",
+ } {
+ if strings.Contains(normalized, marker) {
+ return true
+ }
+ }
+ for _, marker := range []string{"/s0", "/s1", "/s2", "/s3", "/s4", "/s5", "/s6", "/s7", "/s8", "/s9"} {
+ if idx := strings.Index(normalized, marker); idx >= 0 {
+ after := idx + len(marker)
+ if after < len(normalized) && normalized[after] >= '0' && normalized[after] <= '9' {
+ slash := after + 1
+ if slash < len(normalized) && normalized[slash] == '/' {
+ return true
+ }
+ }
+ }
+ }
+ return false
+}
diff --git a/internal/service/emby_playback.go b/internal/service/emby_playback.go
new file mode 100644
index 0000000..97c5eea
--- /dev/null
+++ b/internal/service/emby_playback.go
@@ -0,0 +1,257 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "path/filepath"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// PlaybackInfo returns a PlaybackInfoResponse usable by Emby clients.
+func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string) (map[string]any, error) {
+ m, err := e.playableMedia(ctx, mediaID, userID)
+ if err != nil || m == nil {
+ return nil, err
+ }
+ e.ensureCloudTrackMetadata(ctx, m)
+ return map[string]any{
+ "MediaSources": e.mediaSourcesForItem(ctx, m, false, e.directPlayOnly(ctx)),
+ "PlaySessionId": fmt.Sprintf("%s-%d", m.ID, time.Now().Unix()),
+ }, nil
+}
+
+// ensureCloudTrackMetadata 在后台补齐云盘媒体的轨道元数据。
+//
+// 注意必须是异步的:此前这里在 PlaybackInfo 请求路径上同步执行
+// CloudResolve + ffprobe(HTTP)(最长 8 秒),既把第三方播放器的起播时间
+// 拖长到秒级,又让每一次点开详情/起播都可能触发一次云盘数据下载,是
+// Docker 部署下 CPU/带宽长期居高的来源之一。探测结果落库后,下一次
+// 请求自然能读到完整元数据。
+func (e *EmbyService) ensureCloudTrackMetadata(ctx context.Context, m *model.Media) {
+ if e == nil || m == nil || e.storage == nil || e.probe == nil || !mediaTrackMetadataMissing(m) {
+ return
+ }
+ typ, ref, ok := parseCloudMediaPlaybackURL(m.STRMURL)
+ if !ok {
+ return
+ }
+ mediaID := m.ID
+ e.cloudProbeMu.Lock()
+ if e.cloudProbeInFlight == nil {
+ e.cloudProbeInFlight = make(map[string]struct{})
+ }
+ if _, busy := e.cloudProbeInFlight[mediaID]; busy {
+ e.cloudProbeMu.Unlock()
+ return
+ }
+ e.cloudProbeInFlight[mediaID] = struct{}{}
+ e.cloudProbeMu.Unlock()
+
+ go e.probeCloudTrackMetadata(mediaID, typ, ref)
+}
+
+func (e *EmbyService) probeCloudTrackMetadata(mediaID, typ, ref string) {
+ defer func() {
+ e.cloudProbeMu.Lock()
+ delete(e.cloudProbeInFlight, mediaID)
+ e.cloudProbeMu.Unlock()
+ }()
+ probeCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ link, err := e.storage.CloudResolve(probeCtx, typ, ref, "")
+ if err != nil {
+ if e.log != nil {
+ e.log.Debug("resolve cloud media for playback probe failed", zap.String("media_id", mediaID), zap.Error(err))
+ }
+ return
+ }
+ probe, err := e.probe.ProbeHTTP(probeCtx, link.URL, link.Headers)
+ if err != nil {
+ if e.log != nil {
+ e.log.Debug("playback cloud ffprobe failed", zap.String("media_id", mediaID), zap.Error(err))
+ }
+ return
+ }
+ updates := probeResultUpdates(probe)
+ if len(updates) == 0 {
+ return
+ }
+ if err := e.repo.DB.WithContext(probeCtx).Model(&model.Media{}).Where("id = ?", mediaID).Updates(updates).Error; err != nil && e.log != nil {
+ e.log.Debug("persist playback cloud probe failed", zap.String("media_id", mediaID), zap.Error(err))
+ }
+}
+
+func mediaTrackMetadataMissing(m *model.Media) bool {
+ return m.DurationSec <= 0 ||
+ m.Width <= 0 ||
+ m.Height <= 0 ||
+ strings.TrimSpace(m.VideoCodec) == "" ||
+ strings.TrimSpace(m.AudioCodec) == ""
+}
+
+func parseCloudMediaPlaybackURL(raw string) (string, string, bool) {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return "", "", false
+ }
+ u, err := url.Parse(raw)
+ if err != nil {
+ return "", "", false
+ }
+ path := strings.Trim(u.Path, "/")
+ const prefix = "api/cloud/play/"
+ idx := strings.Index(strings.ToLower(path), prefix)
+ if idx < 0 {
+ return "", "", false
+ }
+ typ := strings.TrimSpace(path[idx+len(prefix):])
+ ref := strings.TrimSpace(u.Query().Get("ref"))
+ return typ, ref, typ != "" && ref != ""
+}
+
+func applyProbeResultToMediaValue(m *model.Media, probe *ProbeResult) {
+ if m == nil || probe == nil {
+ return
+ }
+ if probe.DurationSec > 0 {
+ m.DurationSec = probe.DurationSec
+ }
+ if probe.Width > 0 {
+ m.Width = probe.Width
+ }
+ if probe.Height > 0 {
+ m.Height = probe.Height
+ }
+ if strings.TrimSpace(probe.VideoCodec) != "" {
+ m.VideoCodec = probe.VideoCodec
+ }
+ if strings.TrimSpace(probe.AudioCodec) != "" {
+ m.AudioCodec = probe.AudioCodec
+ }
+ if strings.TrimSpace(probe.Container) != "" {
+ m.Container = probe.Container
+ }
+}
+
+// directPlayOnly reports whether the admin enabled「客户端直连解码」mode.
+// In that mode the host never transcodes; clients must direct-play.
+func (e *EmbyService) directPlayOnly(ctx context.Context) bool {
+ if e.repo == nil || e.repo.Setting == nil {
+ return false
+ }
+ v, err := e.repo.Setting.Get(ctx, PlaybackDirectOnlySettingKey)
+ if err != nil {
+ return false
+ }
+ return parseBoolSetting(v, false)
+}
+
+func (e *EmbyService) playableMedia(ctx context.Context, id, userID string) (*model.Media, error) {
+ if season, ok, err := e.findSeasonGroup(ctx, id, userID); err != nil {
+ return nil, err
+ } else if ok && len(season.Episodes) > 0 {
+ return &season.Episodes[0], nil
+ }
+ if series, ok, err := e.findSeriesGroup(ctx, id, userID); err != nil {
+ return nil, err
+ } else if ok && len(series.Episodes) > 0 {
+ return &series.Episodes[0], nil
+ }
+ m, err := e.repo.Media.FindByID(ctx, id)
+ if err != nil || m == nil {
+ return m, err
+ }
+ if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
+ return nil, nil
+ }
+ return m, nil
+}
+
+// mediaSource 是 /Items 与 /PlaybackInfo 共享的 MediaSource 结构。
+//
+// asEmbedded=true:嵌在 /Items 列表里,不包含完整 stream URL(避免暴露
+// 直链给搜索接口)。/PlaybackInfo 走 false 路径,URL 指向 Emby 兼容
+// /Videos/{id}/stream(客户端会继续携带 X-Emby-Token 或 append api_key)。
+func (e *EmbyService) mediaSource(ctx context.Context, m *model.Media, asEmbedded, directOnly bool) map[string]any {
+ container := embyMediaContainer(m)
+ isCloud := strings.TrimSpace(m.STRMURL) != ""
+ playURL := e.embyMediaPlayURL(ctx, m, container, isCloud)
+ if isCloud {
+ // Cloud/WebDAV media is already a direct/proxy stream. Advertising HLS
+ // transcoding makes some Emby clients pick /master.m3u8, forcing this
+ // lightweight server to pull remote bytes through ffmpeg and often
+ // surfacing as "network/playback failed". Keep cloud media direct-only.
+ directOnly = true
+ }
+ src := e.baseMediaSource(m, container, isCloud, playURL, directOnly)
+ if !asEmbedded && playURL != "" {
+ src["DirectStreamUrl"] = playURL
+ // 直连解码模式下不下发 TranscodingUrl,迫使客户端本地解码直连,
+ // 宿主机不参与转码。
+ if !directOnly {
+ src["TranscodingUrl"] = "/Videos/" + m.ID + "/master.m3u8"
+ }
+ }
+ if strings.TrimSpace(m.STRMURL) != "" && playURL != "" {
+ // STRM / cloud:// media must stay behind a token-aware endpoint. When
+ // STRM playback is enabled we expose /api/stream so third-party clients
+ // follow the same STRM entry as generated .strm files; when disabled we
+ // expose /Videos/{id}/stream so playback uses the Emby 302/proxy path.
+ src["IsRemote"] = true
+ src["Path"] = playURL
+ }
+ return src
+}
+
+func (e *EmbyService) baseMediaSource(m *model.Media, container string, isCloud bool, playURL string, directOnly bool) map[string]any {
+ return map[string]any{
+ "Id": m.ID,
+ "Name": m.Title,
+ "Path": m.Path,
+ "Container": container,
+ "Size": m.SizeBytes,
+ "Protocol": "Http",
+ "Type": "Default",
+ "IsRemote": isCloud,
+ "RequiresOpening": false,
+ "RequiresClosing": false,
+ "ReadAtNativeFramerate": false,
+ "SupportsTranscoding": !directOnly,
+ "SupportsDirectStream": !isCloud || playURL != "",
+ "SupportsDirectPlay": !isCloud || playURL != "",
+ "SupportsProbing": true,
+ "RunTimeTicks": int64(m.DurationSec) * 10_000_000,
+ "MediaStreams": e.mediaStreams(m),
+ }
+}
+
+func embyMediaContainer(m *model.Media) string {
+ container := strings.Trim(strings.ToLower(m.Container), ". ")
+ if container == "" {
+ container = strings.TrimPrefix(strings.ToLower(filepath.Ext(m.Path)), ".")
+ }
+ if container == "" && strings.TrimSpace(m.STRMURL) != "" {
+ return "strm"
+ }
+ return container
+}
+
+func (e *EmbyService) embyMediaPlayURL(ctx context.Context, m *model.Media, container string, isCloud bool) string {
+ if !isCloud {
+ return embyDirectStreamURL(m.ID, container)
+ }
+ switch CloudPlaybackMode(ctx, e.repo) {
+ case CloudPlaybackModeSTRM:
+ return embySTRMStreamURL(m.ID)
+ case CloudPlaybackModeRedirectProxy:
+ return embyDirectStreamURL(m.ID, container)
+ default:
+ return ""
+ }
+}
diff --git a/internal/service/emby_playback_test.go b/internal/service/emby_playback_test.go
new file mode 100644
index 0000000..87e77c9
--- /dev/null
+++ b/internal/service/emby_playback_test.go
@@ -0,0 +1,452 @@
+package service
+
+import (
+ "testing"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+func TestEmbyRootItemsExposeLibraries(t *testing.T) {
+ svc := newTestEmbyService(t)
+ for _, lib := range []model.Library{
+ {Name: "电影", Path: `F:\downloads\电影`, Type: "movie", Enabled: true},
+ {Name: "综艺", Path: `F:\downloads\综艺`, Type: "variety", Enabled: true},
+ } {
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ }
+
+ root, err := svc.Items(t.Context(), ItemsParams{Limit: 50})
+ if err != nil {
+ t.Fatalf("root items: %v", err)
+ }
+ items := root["Items"].([]map[string]any)
+ if len(items) != 2 {
+ t.Fatalf("expected root items to expose libraries, got %#v", items)
+ }
+ if items[0]["Type"] != "CollectionFolder" || items[1]["Type"] != "CollectionFolder" {
+ t.Fatalf("root should return collection folders: %#v", items)
+ }
+ if items[1]["CollectionType"] != "tvshows" {
+ t.Fatalf("variety libraries should use tvshows collection type: %#v", items[1])
+ }
+}
+
+func TestEmbyFolderItemQueryExposesLibrariesForHome(t *testing.T) {
+ svc := newTestEmbyService(t)
+ lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.Media{Base: model.Base{ID: "movie-1"}, LibraryID: lib.ID, Title: "不应出现在文件夹查询", Path: `/media/movies/a.mkv`}).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ out, err := svc.Items(t.Context(), ItemsParams{
+ IncludeItemTypes: []string{"Folder", "CollectionFolder"},
+ Limit: 50,
+ })
+ if err != nil {
+ t.Fatalf("folder items: %v", err)
+ }
+ items := out["Items"].([]map[string]any)
+ if len(items) != 1 {
+ t.Fatalf("expected one library folder, got %#v", items)
+ }
+ if items[0]["Type"] != "CollectionFolder" || items[0]["IsFolder"] != true {
+ t.Fatalf("folder query should return collection folders, got %#v", items[0])
+ }
+}
+
+func TestEmbyUnsupportedItemTypesDoNotLeakAllMedia(t *testing.T) {
+ svc := newTestEmbyService(t)
+ lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.Media{Base: model.Base{ID: "movie-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/a.mkv`}).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ for _, includeType := range []string{"BoxSet", "Game", "Book", "Audio", "MusicAlbum", "Playlist", "TvChannel"} {
+ out, err := svc.Items(t.Context(), ItemsParams{
+ IncludeItemTypes: []string{includeType},
+ Recursive: true,
+ Limit: 50,
+ })
+ if err != nil {
+ t.Fatalf("%s items: %v", includeType, err)
+ }
+ if out["TotalRecordCount"] != int64(0) {
+ t.Fatalf("%s should not return media rows, got %#v", includeType, out)
+ }
+ items := out["Items"].([]map[string]any)
+ if len(items) != 0 {
+ t.Fatalf("%s should return an empty list, got %#v", includeType, items)
+ }
+ }
+}
+
+func TestEmbyItemsFiltersFavorites(t *testing.T) {
+ svc := newTestEmbyService(t)
+ viewer := &model.User{Base: model.Base{ID: "user-1"}, Username: "viewer", Role: "user", Tier: "free", IsActive: true}
+ if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
+ t.Fatalf("create viewer: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ favorite := model.Media{Base: model.Base{ID: "fav-1"}, LibraryID: lib.ID, Title: "收藏电影", Path: `/media/movies/fav.mkv`}
+ normal := model.Media{Base: model.Base{ID: "normal-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/normal.mkv`}
+ if err := svc.repo.DB.Create(&favorite).Error; err != nil {
+ t.Fatalf("create favorite media: %v", err)
+ }
+ if err := svc.repo.DB.Create(&normal).Error; err != nil {
+ t.Fatalf("create normal media: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.Favorite{UserID: viewer.ID, MediaID: favorite.ID}).Error; err != nil {
+ t.Fatalf("create favorite: %v", err)
+ }
+
+ out, err := svc.Items(t.Context(), ItemsParams{
+ UserID: viewer.ID,
+ Filters: []string{"IsFavorite"},
+ Recursive: true,
+ Limit: 50,
+ })
+ if err != nil {
+ t.Fatalf("favorite items: %v", err)
+ }
+ if out["TotalRecordCount"] != int64(1) {
+ t.Fatalf("expected one favorite, got %#v", out)
+ }
+ items := out["Items"].([]map[string]any)
+ if len(items) != 1 || items[0]["Id"] != favorite.ID {
+ t.Fatalf("favorite filter returned wrong items: %#v", items)
+ }
+ userData := items[0]["UserData"].(map[string]any)
+ if userData["IsFavorite"] != true {
+ t.Fatalf("favorite payload should carry IsFavorite=true: %#v", userData)
+ }
+}
+
+func TestEmbyItemsFiltersResumableForHome(t *testing.T) {
+ svc := newTestEmbyService(t)
+ viewer := &model.User{Base: model.Base{ID: "user-1"}, Username: "viewer", Role: "user", Tier: "free", IsActive: true}
+ if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
+ t.Fatalf("create viewer: %v", err)
+ }
+ lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ resumable := model.Media{Base: model.Base{ID: "resume-1"}, LibraryID: lib.ID, Title: "继续观看", Path: `/media/movies/resume.mkv`, DurationSec: 120}
+ normal := model.Media{Base: model.Base{ID: "normal-1"}, LibraryID: lib.ID, Title: "普通电影", Path: `/media/movies/normal.mkv`, DurationSec: 120}
+ if err := svc.repo.DB.Create(&resumable).Error; err != nil {
+ t.Fatalf("create resumable media: %v", err)
+ }
+ if err := svc.repo.DB.Create(&normal).Error; err != nil {
+ t.Fatalf("create normal media: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.PlaybackHistory{
+ UserID: viewer.ID,
+ MediaID: resumable.ID,
+ PositionMs: 30_000,
+ DurationMs: 120_000,
+ WatchedAt: time.Now(),
+ Completed: false,
+ }).Error; err != nil {
+ t.Fatalf("create playback history: %v", err)
+ }
+
+ out, err := svc.Items(t.Context(), ItemsParams{
+ UserID: viewer.ID,
+ Filters: []string{"IsResumable"},
+ Recursive: true,
+ SortBy: "DatePlayed",
+ SortOrder: "Descending",
+ Limit: 50,
+ StartIndex: 0,
+ })
+ if err != nil {
+ t.Fatalf("resumable items: %v", err)
+ }
+ if out["TotalRecordCount"] != int64(1) {
+ t.Fatalf("expected one resumable item, got %#v", out)
+ }
+ items := out["Items"].([]map[string]any)
+ if len(items) != 1 || items[0]["Id"] != resumable.ID {
+ t.Fatalf("resumable filter returned wrong items: %#v", items)
+ }
+}
+
+func TestEmbyUserPolicyDisablesDownloadsForViewers(t *testing.T) {
+ svc := newTestEmbyService(t)
+ viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true}
+ admin := &model.User{Username: "admin", Role: "admin", Tier: "plus", IsActive: true}
+ if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
+ t.Fatalf("create viewer: %v", err)
+ }
+ if err := svc.repo.User.Create(t.Context(), admin); err != nil {
+ t.Fatalf("create admin: %v", err)
+ }
+
+ viewerPayload, err := svc.FindUser(t.Context(), viewer.ID)
+ if err != nil {
+ t.Fatalf("viewer payload: %v", err)
+ }
+ adminPayload, err := svc.FindUser(t.Context(), admin.ID)
+ if err != nil {
+ t.Fatalf("admin payload: %v", err)
+ }
+ viewerPolicy := viewerPayload["Policy"].(map[string]any)
+ adminPolicy := adminPayload["Policy"].(map[string]any)
+ if viewerPolicy["EnableMediaPlayback"] != true {
+ t.Fatalf("viewer must keep playback enabled: %#v", viewerPolicy)
+ }
+ if viewerPolicy["EnableContentDownloading"] != false ||
+ viewerPolicy["EnableSyncTranscoding"] != false ||
+ viewerPolicy["EnableMediaConversion"] != false {
+ t.Fatalf("viewer must not be allowed to download/sync media: %#v", viewerPolicy)
+ }
+ if adminPolicy["EnableContentDownloading"] != true {
+ t.Fatalf("admin should keep downloading capability: %#v", adminPolicy)
+ }
+}
+
+func TestEmbyHidesAdultLibrariesForUserLock(t *testing.T) {
+ svc := newTestEmbyService(t)
+ viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true, HideAdult: true}
+ if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
+ t.Fatalf("create viewer: %v", err)
+ }
+ safe := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ adult := model.Library{Name: "9KG 成人", Path: `/media/9KG`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &safe); err != nil {
+ t.Fatalf("create safe library: %v", err)
+ }
+ if err := svc.repo.Library.Create(t.Context(), &adult); err != nil {
+ t.Fatalf("create adult library: %v", err)
+ }
+ if err := svc.repo.Setting.Set(t.Context(), AdultLibraryIDsSettingKey, `["`+adult.ID+`"]`); err != nil {
+ t.Fatalf("set adult libraries: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.Media{LibraryID: safe.ID, Title: "安全电影", Path: `/media/movies/a.mkv`}).Error; err != nil {
+ t.Fatalf("create safe media: %v", err)
+ }
+ if err := svc.repo.DB.Create(&model.Media{LibraryID: adult.ID, Title: "成人电影", Path: `/media/9KG/a.mkv`}).Error; err != nil {
+ t.Fatalf("create adult media: %v", err)
+ }
+
+ root, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, Limit: 50})
+ if err != nil {
+ t.Fatalf("root items: %v", err)
+ }
+ items := root["Items"].([]map[string]any)
+ if len(items) != 1 || items[0]["Name"] != "电影" {
+ t.Fatalf("adult library should be hidden: %#v", items)
+ }
+ adultItems, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, ParentID: adult.ID, Limit: 50})
+ if err != nil {
+ t.Fatalf("adult items: %v", err)
+ }
+ if got := adultItems["TotalRecordCount"]; got != int64(0) {
+ t.Fatalf("adult media should be hidden, total=%#v payload=%#v", got, adultItems)
+ }
+}
+
+func TestEmbyPlaybackInfoRespectsDirectPlayOnly(t *testing.T) {
+ svc := newTestEmbyService(t)
+ lib := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ media := model.Media{Base: model.Base{ID: "m-1"}, LibraryID: lib.ID, Title: "Inception", Path: `/media/movies/inception.mkv`}
+ if err := svc.repo.DB.Create(&media).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ pb, err := svc.PlaybackInfo(t.Context(), "m-1", "user-1")
+ if err != nil {
+ t.Fatalf("playback info: %v", err)
+ }
+ src := pb["MediaSources"].([]map[string]any)[0]
+ if src["SupportsTranscoding"] != true {
+ t.Fatalf("expected SupportsTranscoding=true by default, got %#v", src["SupportsTranscoding"])
+ }
+ if _, ok := src["TranscodingUrl"]; !ok {
+ t.Fatalf("expected TranscodingUrl present by default: %#v", src)
+ }
+ if src["TranscodingUrl"] != "/Videos/m-1/master.m3u8" {
+ t.Fatalf("expected HLS TranscodingUrl by default, got %#v", src["TranscodingUrl"])
+ }
+
+ if err := svc.repo.Setting.Set(t.Context(), PlaybackDirectOnlySettingKey, "true"); err != nil {
+ t.Fatalf("enable direct-only: %v", err)
+ }
+ pb, err = svc.PlaybackInfo(t.Context(), "m-1", "user-1")
+ if err != nil {
+ t.Fatalf("playback info (direct-only): %v", err)
+ }
+ src = pb["MediaSources"].([]map[string]any)[0]
+ if src["SupportsTranscoding"] != false {
+ t.Fatalf("expected SupportsTranscoding=false in direct-only mode, got %#v", src["SupportsTranscoding"])
+ }
+ if _, ok := src["TranscodingUrl"]; ok {
+ t.Fatalf("expected no TranscodingUrl in direct-only mode: %#v", src)
+ }
+ if src["SupportsDirectPlay"] != true || src["DirectStreamUrl"] != "/Videos/m-1/stream.mkv" {
+ t.Fatalf("direct-only must still allow direct play: %#v", src)
+ }
+}
+
+func TestEmbyPlaybackInfoKeepsSTRMBehindStreamEndpoint(t *testing.T) {
+ svc := newTestEmbyService(t)
+ if err := svc.repo.Setting.Set(t.Context(), CloudPlaybackModeSettingKey, CloudPlaybackModeSTRM); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: `cloud://openlist/Movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ media := model.Media{
+ Base: model.Base{ID: "cloud-1"},
+ LibraryID: lib.ID,
+ Title: "Cloud Movie",
+ Path: `cloud://openlist/Movies/f1.mkv`,
+ STRMURL: `/api/cloud/play/openlist?ref=%2FMovies%2Ff1.mkv`,
+ }
+ if err := svc.repo.DB.Create(&media).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ pb, err := svc.PlaybackInfo(t.Context(), "cloud-1", "user-1")
+ if err != nil {
+ t.Fatalf("playback info: %v", err)
+ }
+ src := pb["MediaSources"].([]map[string]any)[0]
+ if src["IsRemote"] != true {
+ t.Fatalf("strm media should be marked remote: %#v", src)
+ }
+ if src["DirectStreamUrl"] != "/api/stream/cloud-1" {
+ t.Fatalf("strm playback should prefer /api/stream when enabled: %#v", src)
+ }
+ if src["Path"] != "/api/stream/cloud-1" {
+ t.Fatalf("path should prefer /api/stream when enabled: %#v", src)
+ }
+ streams := src["MediaStreams"].([]map[string]any)
+ if len(streams) == 0 || streams[0]["Type"] != "Video" {
+ t.Fatalf("strm media should expose a fallback video stream for Android clients: %#v", src)
+ }
+}
+
+func TestEmbyPlaybackInfoUsesVideoStreamWhenSTRMDisabled(t *testing.T) {
+ svc := newTestEmbyService(t)
+ if err := svc.repo.Setting.Set(t.Context(), CloudPlaybackModeSettingKey, CloudPlaybackModeRedirectProxy); err != nil {
+ t.Fatalf("set cloud playback mode: %v", err)
+ }
+ lib := model.Library{Name: "OpenList", Path: `cloud://openlist/Movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ media := model.Media{
+ Base: model.Base{ID: "cloud-302"},
+ LibraryID: lib.ID,
+ Title: "Cloud 302 Movie",
+ Path: `cloud://openlist/Movies/Movie.mkv`,
+ STRMURL: `/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv`,
+ Container: "mkv",
+ }
+ if err := svc.repo.DB.Create(&media).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+
+ pb, err := svc.PlaybackInfo(t.Context(), "cloud-302", "user-1")
+ if err != nil {
+ t.Fatalf("playback info: %v", err)
+ }
+ src := pb["MediaSources"].([]map[string]any)[0]
+ if src["DirectStreamUrl"] != "/Videos/cloud-302/stream.mkv" {
+ t.Fatalf("302/proxy mode should use Emby video stream URL: %#v", src)
+ }
+ if src["Path"] != "/Videos/cloud-302/stream.mkv" {
+ t.Fatalf("302/proxy mode path should use Emby video stream URL: %#v", src)
+ }
+}
+
+func TestEmbyPlaybackInfoProbesMissingCloudTrackMetadata(t *testing.T) {
+ svc := newTestEmbyService(t)
+ lib := model.Library{Name: "OpenList", Path: `cloud://openlist/Movies`, Type: "movie", Enabled: true}
+ if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatalf("create library: %v", err)
+ }
+ media := model.Media{
+ Base: model.Base{ID: "cloud-probe-1"},
+ LibraryID: lib.ID,
+ Title: "云盘电影",
+ Path: `cloud://openlist/Movies/Movie.mkv`,
+ STRMURL: `http://nas.local/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv`,
+ }
+ if err := svc.repo.DB.Create(&media).Error; err != nil {
+ t.Fatalf("create media: %v", err)
+ }
+ resolver := &fakeCloudPlaybackResolver{
+ link: &cloud.DirectLink{
+ URL: "http://cdn.example.test/Movie.mkv",
+ Headers: map[string]string{"Authorization": "Bearer probe-token"},
+ },
+ }
+ prober := &fakeCloudPlaybackProber{
+ probe: &ProbeResult{
+ DurationSec: 3661,
+ Width: 3840,
+ Height: 2160,
+ VideoCodec: "hevc",
+ AudioCodec: "eac3",
+ Container: "matroska,webm",
+ },
+ }
+ svc.SetCloudProbe(resolver, prober)
+
+ if _, err := svc.PlaybackInfo(t.Context(), "cloud-probe-1", "user-1"); err != nil {
+ t.Fatalf("playback info: %v", err)
+ }
+
+ var persisted model.Media
+ deadline := time.Now().Add(3 * time.Second)
+ for {
+ if err := svc.repo.DB.First(&persisted, "id = ?", "cloud-probe-1").Error; err != nil {
+ t.Fatalf("reload media: %v", err)
+ }
+ if persisted.DurationSec > 0 || time.Now().After(deadline) {
+ break
+ }
+ time.Sleep(10 * time.Millisecond)
+ }
+ if persisted.DurationSec != 3661 || persisted.Width != 3840 || persisted.Height != 2160 || persisted.VideoCodec != "hevc" || persisted.AudioCodec != "eac3" {
+ t.Fatalf("probe metadata not persisted: %#v", persisted)
+ }
+ if resolver.typ != "openlist" || resolver.ref != "/Movies/Movie.mkv" {
+ t.Fatalf("resolver called with typ=%q ref=%q", resolver.typ, resolver.ref)
+ }
+ if prober.rawURL != "http://cdn.example.test/Movie.mkv" || prober.headers["Authorization"] != "Bearer probe-token" {
+ t.Fatalf("probe called with url=%q headers=%#v", prober.rawURL, prober.headers)
+ }
+
+ pb, err := svc.PlaybackInfo(t.Context(), "cloud-probe-1", "user-1")
+ if err != nil {
+ t.Fatalf("playback info (second): %v", err)
+ }
+ src := pb["MediaSources"].([]map[string]any)[0]
+ if src["RunTimeTicks"] != int64(3661)*10_000_000 {
+ t.Fatalf("runtime ticks not filled after async probe: %#v", src)
+ }
+ streams := src["MediaStreams"].([]map[string]any)
+ if len(streams) != 2 || streams[0]["Codec"] != "hevc" || streams[1]["Codec"] != "eac3" {
+ t.Fatalf("media streams not filled after async probe: %#v", streams)
+ }
+}
diff --git a/internal/service/emby_series.go b/internal/service/emby_series.go
new file mode 100644
index 0000000..fdaba2b
--- /dev/null
+++ b/internal/service/emby_series.go
@@ -0,0 +1,197 @@
+package service
+
+import (
+ "context"
+ "sort"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type embySeriesGroup struct {
+ ID string
+ LibraryID string
+ Name string
+ PosterURL string
+ BackdropURL string
+ Overview string
+ Rating float32
+ Year int
+ TMDbID int
+ BangumiID int
+ CreatedAt time.Time
+ Episodes []model.Media
+}
+
+type embySeasonGroup struct {
+ ID string
+ SeriesID string
+ LibraryID string
+ Name string
+ SeasonNum int
+ Series embySeriesGroup
+ Episodes []model.Media
+}
+
+func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (embySeriesGroup, bool, error) {
+ if strings.TrimSpace(id) == "" {
+ return embySeriesGroup{}, false, nil
+ }
+ if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
+ if group, ok := e.cachedSeriesGroup(id); ok {
+ return group, true, nil
+ }
+ }
+ var rows []model.Media
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
+ q = e.applyUserMediaVisibility(ctx, q, userID)
+ if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
+ q = q.Where("series_id = ?", id)
+ }
+ if err := q.Order("media.season_num asc, media.episode_num asc, media.created_at asc").Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
+ return embySeriesGroup{}, false, err
+ }
+ for _, group := range e.seriesGroupsFromMedia(rows) {
+ if group.ID == id {
+ e.rememberSeriesGroup(group)
+ return group, true, nil
+ }
+ }
+ if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
+ if series, err := e.repo.Series.FindByID(ctx, id); err != nil {
+ return embySeriesGroup{}, false, err
+ } else if series != nil {
+ return embySeriesGroup{
+ ID: series.ID,
+ LibraryID: series.LibraryID,
+ Name: series.Title,
+ PosterURL: series.PosterURL,
+ BackdropURL: series.BackdropURL,
+ Overview: series.Overview,
+ Rating: series.Rating,
+ Year: series.Year,
+ TMDbID: series.TMDbID,
+ BangumiID: series.BangumiID,
+ CreatedAt: series.CreatedAt,
+ }, true, nil
+ }
+ }
+ return embySeriesGroup{}, false, nil
+}
+
+func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (embySeasonGroup, bool, error) {
+ if strings.TrimSpace(id) == "" || !strings.HasPrefix(id, embyVirtualSeasonPrefix) {
+ return embySeasonGroup{}, false, nil
+ }
+ if season, ok := e.cachedSeasonGroup(id); ok {
+ return season, true, nil
+ }
+ var rows []model.Media
+ q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("season_num > 0 OR episode_num > 0")
+ q = e.applyUserMediaVisibility(ctx, q, userID)
+ if err := q.
+ Order("media.season_num asc, media.episode_num asc, media.created_at asc").
+ Limit(embySeriesGroupingLimit).
+ Find(&rows).Error; err != nil {
+ return embySeasonGroup{}, false, err
+ }
+ for _, series := range e.seriesGroupsFromMedia(rows) {
+ for _, season := range e.seasonsForSeries(series) {
+ if season.ID == id {
+ e.rememberSeriesGroup(series)
+ return season, true, nil
+ }
+ }
+ }
+ return embySeasonGroup{}, false, nil
+}
+
+func (e *EmbyService) seriesGroupsFromMedia(rows []model.Media) []embySeriesGroup {
+ byID := map[string]*embySeriesGroup{}
+ order := []string{}
+ for _, row := range rows {
+ row := row
+ seriesID := e.seriesIDForMedia(&row)
+ group, ok := byID[seriesID]
+ if !ok {
+ group = &embySeriesGroup{
+ ID: seriesID,
+ LibraryID: row.LibraryID,
+ Name: e.seriesNameForMedia(&row),
+ Year: row.Year,
+ TMDbID: row.TMDbID,
+ BangumiID: row.BangumiID,
+ CreatedAt: row.CreatedAt,
+ }
+ byID[seriesID] = group
+ order = append(order, seriesID)
+ }
+ if row.CreatedAt.After(group.CreatedAt) {
+ group.CreatedAt = row.CreatedAt
+ }
+ if group.PosterURL == "" && row.PosterURL != "" {
+ group.PosterURL = row.PosterURL
+ }
+ if group.BackdropURL == "" && row.BackdropURL != "" {
+ group.BackdropURL = row.BackdropURL
+ }
+ if group.Overview == "" && row.Overview != "" {
+ group.Overview = row.Overview
+ }
+ if group.Rating == 0 && row.Rating > 0 {
+ group.Rating = row.Rating
+ }
+ if group.Year == 0 && row.Year > 0 {
+ group.Year = row.Year
+ }
+ group.Episodes = append(group.Episodes, row)
+ }
+ groups := make([]embySeriesGroup, 0, len(order))
+ for _, id := range order {
+ group := *byID[id]
+ sort.SliceStable(group.Episodes, func(i, j int) bool {
+ if group.Episodes[i].SeasonNum != group.Episodes[j].SeasonNum {
+ return group.Episodes[i].SeasonNum < group.Episodes[j].SeasonNum
+ }
+ if group.Episodes[i].EpisodeNum != group.Episodes[j].EpisodeNum {
+ return group.Episodes[i].EpisodeNum < group.Episodes[j].EpisodeNum
+ }
+ return group.Episodes[i].CreatedAt.Before(group.Episodes[j].CreatedAt)
+ })
+ groups = append(groups, group)
+ }
+ return groups
+}
+
+func (e *EmbyService) seasonsForSeries(series embySeriesGroup) []embySeasonGroup {
+ bySeason := map[int]*embySeasonGroup{}
+ order := []int{}
+ for _, episode := range series.Episodes {
+ seasonNum := episode.SeasonNum
+ if seasonNum < 0 {
+ seasonNum = 1
+ }
+ season, ok := bySeason[seasonNum]
+ if !ok {
+ season = &embySeasonGroup{
+ ID: seasonID(series.ID, seasonNum),
+ SeriesID: series.ID,
+ LibraryID: series.LibraryID,
+ Name: seasonName(seasonNum),
+ SeasonNum: seasonNum,
+ Series: series,
+ }
+ bySeason[seasonNum] = season
+ order = append(order, seasonNum)
+ }
+ season.Episodes = append(season.Episodes, episode)
+ }
+ sort.Ints(order)
+ out := make([]embySeasonGroup, 0, len(order))
+ for _, seasonNum := range order {
+ out = append(out, *bySeason[seasonNum])
+ }
+ return out
+}
diff --git a/internal/service/emby_series_cache.go b/internal/service/emby_series_cache.go
new file mode 100644
index 0000000..f65fcf8
--- /dev/null
+++ b/internal/service/emby_series_cache.go
@@ -0,0 +1,137 @@
+package service
+
+import (
+ "strings"
+ "time"
+)
+
+type embySeriesCacheEntry struct {
+ group embySeriesGroup
+ expiresAt time.Time
+}
+
+type embySeasonCacheEntry struct {
+ season embySeasonGroup
+ expiresAt time.Time
+}
+
+type embyArtworkCacheEntry struct {
+ primary string
+ backdrop string
+ expiresAt time.Time
+}
+
+func (e *EmbyService) rememberSeriesGroup(group embySeriesGroup) {
+ if e == nil || strings.TrimSpace(group.ID) == "" {
+ return
+ }
+ expiresAt := time.Now().Add(embyVirtualCacheTTL)
+ e.virtualMu.Lock()
+ defer e.virtualMu.Unlock()
+ if e.virtualSeries == nil {
+ e.virtualSeries = make(map[string]embySeriesCacheEntry)
+ }
+ if e.virtualSeasons == nil {
+ e.virtualSeasons = make(map[string]embySeasonCacheEntry)
+ }
+ if e.virtualArtwork == nil {
+ e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
+ }
+ if len(e.virtualSeries) > 2000 || len(e.virtualSeasons) > 5000 || len(e.virtualArtwork) > 7000 {
+ e.virtualSeries = make(map[string]embySeriesCacheEntry)
+ e.virtualSeasons = make(map[string]embySeasonCacheEntry)
+ e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
+ }
+ e.virtualSeries[group.ID] = embySeriesCacheEntry{group: group, expiresAt: expiresAt}
+ e.virtualArtwork[group.ID] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
+ e.virtualArtwork[group.ID+"-bd"] = embyArtworkCacheEntry{primary: group.PosterURL, backdrop: group.BackdropURL, expiresAt: expiresAt}
+ for _, season := range e.seasonsForSeries(group) {
+ e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
+ e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
+ e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
+ }
+}
+
+func (e *EmbyService) rememberSeasonGroup(season embySeasonGroup) {
+ if e == nil || strings.TrimSpace(season.ID) == "" {
+ return
+ }
+ expiresAt := time.Now().Add(embyVirtualCacheTTL)
+ e.virtualMu.Lock()
+ defer e.virtualMu.Unlock()
+ if e.virtualSeasons == nil {
+ e.virtualSeasons = make(map[string]embySeasonCacheEntry)
+ }
+ if e.virtualArtwork == nil {
+ e.virtualArtwork = make(map[string]embyArtworkCacheEntry)
+ }
+ e.virtualSeasons[season.ID] = embySeasonCacheEntry{season: season, expiresAt: expiresAt}
+ e.virtualArtwork[season.ID] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
+ e.virtualArtwork[season.ID+"-bd"] = embyArtworkCacheEntry{primary: season.Series.PosterURL, backdrop: season.Series.BackdropURL, expiresAt: expiresAt}
+}
+
+func (e *EmbyService) cachedSeriesGroup(id string) (embySeriesGroup, bool) {
+ if e == nil || strings.TrimSpace(id) == "" {
+ return embySeriesGroup{}, false
+ }
+ now := time.Now()
+ e.virtualMu.RLock()
+ entry, ok := e.virtualSeries[id]
+ e.virtualMu.RUnlock()
+ if !ok || now.After(entry.expiresAt) {
+ if ok {
+ e.virtualMu.Lock()
+ delete(e.virtualSeries, id)
+ e.virtualMu.Unlock()
+ }
+ return embySeriesGroup{}, false
+ }
+ return entry.group, true
+}
+
+func (e *EmbyService) cachedSeasonGroup(id string) (embySeasonGroup, bool) {
+ if e == nil || strings.TrimSpace(id) == "" {
+ return embySeasonGroup{}, false
+ }
+ now := time.Now()
+ e.virtualMu.RLock()
+ entry, ok := e.virtualSeasons[id]
+ e.virtualMu.RUnlock()
+ if !ok || now.After(entry.expiresAt) {
+ if ok {
+ e.virtualMu.Lock()
+ delete(e.virtualSeasons, id)
+ e.virtualMu.Unlock()
+ }
+ return embySeasonGroup{}, false
+ }
+ return entry.season, true
+}
+
+func (e *EmbyService) cachedArtworkURL(id, imageType string) (string, bool) {
+ if e == nil || strings.TrimSpace(id) == "" {
+ return "", false
+ }
+ now := time.Now()
+ e.virtualMu.RLock()
+ entry, ok := e.virtualArtwork[id]
+ e.virtualMu.RUnlock()
+ if !ok || now.After(entry.expiresAt) {
+ if ok {
+ e.virtualMu.Lock()
+ delete(e.virtualArtwork, id)
+ e.virtualMu.Unlock()
+ }
+ return "", false
+ }
+ switch strings.ToLower(imageType) {
+ case "backdrop", "art":
+ if entry.backdrop != "" {
+ return entry.backdrop, true
+ }
+ }
+ if entry.primary != "" {
+ return entry.primary, true
+ }
+ return entry.backdrop, entry.backdrop != ""
+}
diff --git a/internal/service/emby_series_ids.go b/internal/service/emby_series_ids.go
new file mode 100644
index 0000000..67e42c1
--- /dev/null
+++ b/internal/service/emby_series_ids.go
@@ -0,0 +1,106 @@
+package service
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/hex"
+ "fmt"
+ "path/filepath"
+ "sort"
+ "strconv"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (e *EmbyService) seriesIDForMedia(m *model.Media) string {
+ if strings.TrimSpace(m.SeriesID) != "" {
+ return m.SeriesID
+ }
+ return stableEmbyID(embyVirtualSeriesPrefix, m.LibraryID, e.seriesNameForMedia(m))
+}
+
+func (e *EmbyService) seasonIDForMedia(m *model.Media) string {
+ return seasonID(e.seriesIDForMedia(m), m.SeasonNum)
+}
+
+func (e *EmbyService) seriesNameForMedia(m *model.Media) string {
+ if strings.TrimSpace(m.SeriesID) != "" {
+ if series, err := e.repo.Series.FindByID(context.Background(), m.SeriesID); err == nil && series != nil && strings.TrimSpace(series.Title) != "" {
+ return series.Title
+ }
+ }
+ if name := inferSeriesNameFromPath(m.Path); name != "" {
+ return name
+ }
+ name := strings.TrimSpace(m.Title)
+ name = embyEpisodeTitleRE.ReplaceAllString(name, "")
+ name = embyYearSuffixRE.ReplaceAllString(name, "")
+ if name == "" {
+ name = strings.TrimSpace(m.OriginalName)
+ }
+ return name
+}
+
+func inferSeriesNameFromPath(path string) string {
+ path = strings.TrimSpace(path)
+ if path == "" {
+ return ""
+ }
+ dir := filepath.Dir(path)
+ base := filepath.Base(dir)
+ if embySeasonDirRE.MatchString(base) {
+ dir = filepath.Dir(dir)
+ base = filepath.Base(dir)
+ }
+ base = strings.TrimSpace(embyYearSuffixRE.ReplaceAllString(base, ""))
+ if base == "." || base == string(filepath.Separator) {
+ return ""
+ }
+ return base
+}
+
+func stableEmbyID(prefix string, parts ...string) string {
+ h := sha256.New()
+ for _, part := range parts {
+ _, _ = h.Write([]byte(strings.ToLower(strings.TrimSpace(part))))
+ _, _ = h.Write([]byte{0})
+ }
+ return prefix + hex.EncodeToString(h.Sum(nil))[:32]
+}
+
+func seasonID(seriesID string, seasonNum int) string {
+ if seasonNum < 0 {
+ seasonNum = 1
+ }
+ return stableEmbyID(embyVirtualSeasonPrefix, seriesID, strconv.Itoa(seasonNum))
+}
+
+func seasonName(seasonNum int) string {
+ if seasonNum == 0 {
+ return "特别篇"
+ }
+ if seasonNum < 0 {
+ seasonNum = 1
+ }
+ return fmt.Sprintf("第 %d 季", seasonNum)
+}
+
+func sortSeriesGroups(groups []embySeriesGroup, p ItemsParams) {
+ switch strings.ToLower(p.SortBy) {
+ case "sortname", "name":
+ sort.SliceStable(groups, func(i, j int) bool {
+ if strings.EqualFold(p.SortOrder, "Descending") {
+ return groups[i].Name > groups[j].Name
+ }
+ return groups[i].Name < groups[j].Name
+ })
+ default:
+ sort.SliceStable(groups, func(i, j int) bool {
+ if strings.EqualFold(p.SortOrder, "Ascending") {
+ return groups[i].CreatedAt.Before(groups[j].CreatedAt)
+ }
+ return groups[i].CreatedAt.After(groups[j].CreatedAt)
+ })
+ }
+}
diff --git a/internal/service/emby_series_payload.go b/internal/service/emby_series_payload.go
new file mode 100644
index 0000000..10bab89
--- /dev/null
+++ b/internal/service/emby_series_payload.go
@@ -0,0 +1,63 @@
+package service
+
+func (e *EmbyService) seriesPayload(group embySeriesGroup) map[string]any {
+ e.rememberSeriesGroup(group)
+ imageTags := map[string]string{}
+ backdropTags := []string{}
+ if group.PosterURL != "" {
+ imageTags["Primary"] = group.ID
+ }
+ if group.BackdropURL != "" {
+ backdropTags = append(backdropTags, group.ID+"-bd")
+ }
+ return map[string]any{
+ "Id": group.ID,
+ "Name": group.Name,
+ "ServerId": embyServerID,
+ "Type": "Series",
+ "MediaType": "Video",
+ "IsFolder": true,
+ "ParentId": group.LibraryID,
+ "ProductionYear": group.Year,
+ "Overview": group.Overview,
+ "CommunityRating": group.Rating,
+ "RecursiveItemCount": len(group.Episodes),
+ "ChildCount": len(e.seasonsForSeries(group)),
+ "DateCreated": group.CreatedAt,
+ "ImageTags": imageTags,
+ "BackdropImageTags": backdropTags,
+ "ProviderIds": map[string]string{
+ "Tmdb": intToStr(group.TMDbID),
+ "Bangumi": intToStr(group.BangumiID),
+ },
+ "UserData": emptyUserData(),
+ }
+}
+
+func (e *EmbyService) seasonPayload(season embySeasonGroup) map[string]any {
+ e.rememberSeasonGroup(season)
+ imageTags := map[string]string{}
+ backdropTags := []string{}
+ if season.Series.PosterURL != "" {
+ imageTags["Primary"] = season.ID
+ }
+ if season.Series.BackdropURL != "" {
+ backdropTags = append(backdropTags, season.ID+"-bd")
+ }
+ return map[string]any{
+ "Id": season.ID,
+ "Name": season.Name,
+ "ServerId": embyServerID,
+ "Type": "Season",
+ "MediaType": "Video",
+ "IsFolder": true,
+ "ParentId": season.SeriesID,
+ "SeriesId": season.SeriesID,
+ "SeriesName": season.Series.Name,
+ "IndexNumber": season.SeasonNum,
+ "ChildCount": len(season.Episodes),
+ "ImageTags": imageTags,
+ "BackdropImageTags": backdropTags,
+ "UserData": emptyUserData(),
+ }
+}
diff --git a/internal/service/emby_system.go b/internal/service/emby_system.go
new file mode 100644
index 0000000..e3c6336
--- /dev/null
+++ b/internal/service/emby_system.go
@@ -0,0 +1,191 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// SystemInfo returns the full Emby identity payload.
+func (e *EmbyService) SystemInfo() map[string]any {
+ return map[string]any{
+ "Id": embyServerID,
+ "ServerId": embyServerID,
+ "ServerName": "MediaStationGo",
+ "Version": embyCompatVersion,
+ "ServerVersion": embyCompatVersion,
+ "ProductName": "Emby Server",
+ "OperatingSystem": "Windows",
+ "Architecture": "X64",
+ "LocalAddress": "",
+ "WanAddress": "",
+ "HasPendingRestart": false,
+ "IsShuttingDown": false,
+ "SupportsLibraryMonitor": true,
+ "SupportsHttps": false,
+ "SupportsAutoDiscovery": true,
+ "HttpServerPortNumber": e.cfg.App.Port,
+ "HttpsPortNumber": 0,
+ "PublishedServerUrl": "",
+ "WebSocketPortNumber": e.cfg.App.Port,
+ "CompletedInstallations": []any{},
+ "CanSelfRestart": false,
+ "CanLaunchWebBrowser": false,
+ "CanRestart": false,
+ }
+}
+
+// SystemInfoPublic 是不需要认证的精简版(Emby Web 客户端登陆前会拉)。
+func (e *EmbyService) SystemInfoPublic() map[string]any {
+ return map[string]any{
+ "Id": embyServerID,
+ "ServerId": embyServerID,
+ "ServerName": "MediaStationGo",
+ "Version": embyCompatVersion,
+ "ServerVersion": embyCompatVersion,
+ "ProductName": "Emby Server",
+ "OperatingSystem": "Windows",
+ "LocalAddress": "",
+ "WanAddress": "",
+ "HttpServerPortNumber": e.cfg.App.Port,
+ "HttpsPortNumber": 0,
+ "SupportsHttps": false,
+ "SupportsAutoDiscovery": true,
+ "StartupWizardCompleted": true,
+ }
+}
+
+// ListUsers returns Emby-shaped users.
+func (e *EmbyService) ListUsers(ctx context.Context) ([]map[string]any, error) {
+ users, err := e.repo.User.List(ctx)
+ if err != nil {
+ return nil, err
+ }
+ out := make([]map[string]any, 0, len(users))
+ for _, u := range users {
+ out = append(out, e.userPayload(&u))
+ }
+ return out, nil
+}
+
+// FindUser 用 ID 查用户,用于 /Users/Me 与 /Users/{id}。
+func (e *EmbyService) FindUser(ctx context.Context, id string) (map[string]any, error) {
+ u, err := e.repo.User.FindByID(ctx, id)
+ if err != nil || u == nil {
+ return nil, err
+ }
+ return e.userPayload(u), nil
+}
+
+func (e *EmbyService) userPayload(u *model.User) map[string]any {
+ canDownload := u.Role == "admin"
+ return map[string]any{
+ "Id": u.ID,
+ "Name": u.Username,
+ "ServerId": embyServerID,
+ "ServerName": "MediaStationGo",
+ "HasPassword": true,
+ "HasConfiguredPassword": true,
+ "HasConfiguredEasyPassword": false,
+ "EnableAutoLogin": false,
+ "LastLoginDate": u.LastLoginAt,
+ "LastActivityDate": u.UpdatedAt,
+ "Configuration": map[string]any{
+ "PlayDefaultAudioTrack": true,
+ "DisplayCollectionsView": true,
+ "DisplayMissingEpisodes": false,
+ "SubtitleMode": "Default",
+ "EnableNextEpisodeAutoPlay": true,
+ "AudioLanguagePreference": "",
+ "SubtitleLanguagePreference": "",
+ },
+ "Policy": map[string]any{
+ "IsAdministrator": u.Role == "admin",
+ "IsHidden": false,
+ "IsDisabled": !u.IsActive,
+ "EnableUserPreferenceAccess": true,
+ "EnableRemoteAccess": true,
+ "EnableMediaPlayback": true,
+ "EnableAudioPlaybackTranscoding": true,
+ "EnableVideoPlaybackTranscoding": true,
+ "EnablePlaybackRemuxing": true,
+ "EnableLiveTvAccess": false,
+ "EnableContentDownloading": canDownload,
+ "EnableSyncTranscoding": canDownload,
+ "EnableMediaConversion": canDownload,
+ "EnableAllChannels": true,
+ "EnableAllFolders": true,
+ "EnableAllDevices": true,
+ "AuthenticationProviderId": embyLocalAuthenticationProviderID,
+ "PasswordResetProviderId": embyLocalPasswordResetProviderID,
+ },
+ }
+}
+
+// Views 返回 Emby 中"虚拟根目录"——每个 library 一个条目。
+func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any, error) {
+ libs, err := e.repo.Library.List(ctx)
+ if err != nil {
+ return nil, err
+ }
+ libs = FilterDisplayCloudLibraries(ctx, e.repo, libs)
+ visibility := e.mediaVisibility(ctx, userID)
+ items := make([]map[string]any, 0, len(libs))
+ for _, l := range libs {
+ if !e.libraryVisibleFromCachedVisibility(l, visibility) {
+ continue
+ }
+ items = append(items, e.libraryAsView(&l))
+ }
+ return map[string]any{"Items": items, "TotalRecordCount": len(items), "StartIndex": 0}, nil
+}
+
+func (e *EmbyService) libraryAsView(l *model.Library) map[string]any {
+ collectionType := "movies"
+ switch l.Type {
+ case "tv":
+ collectionType = "tvshows"
+ case "anime":
+ collectionType = "tvshows" // Emby 没有专门的 anime CollectionType
+ case "variety":
+ collectionType = "tvshows"
+ case "music":
+ collectionType = "music"
+ }
+ return map[string]any{
+ "Id": l.ID,
+ "Name": l.Name,
+ "CollectionType": collectionType,
+ "ServerId": embyServerID,
+ "Type": "CollectionFolder",
+ "IsFolder": true,
+ "Path": l.Path,
+ "SortName": strings.ToLower(l.Name),
+ "DateCreated": l.CreatedAt.UTC().Format(time.RFC3339),
+ "CanDelete": false,
+ "CanDownload": false,
+ "DisplayPreferencesId": l.ID,
+ "PrimaryImageItemId": l.ID,
+ "PrimaryImageAspectRatio": 1.7777777777777777,
+ "RecursiveItemCount": 0,
+ "ChildCount": 0,
+ "SpecialFeatureCount": 0,
+ "EnableMediaSourceDisplay": true,
+ "PlayAccess": "Full",
+ "ExternalUrls": []any{},
+ "ProviderIds": map[string]string{},
+ "Genres": []string{},
+ "Tags": []string{},
+ "ImageTags": map[string]string{},
+ "BackdropImageTags": []string{},
+ "UserData": map[string]any{
+ "PlaybackPositionTicks": 0,
+ "PlayCount": 0,
+ "IsFavorite": false,
+ "Played": false,
+ "UnplayedItemCount": 0,
+ },
+ }
+}
diff --git a/internal/service/emby_user_data.go b/internal/service/emby_user_data.go
new file mode 100644
index 0000000..6040b08
--- /dev/null
+++ b/internal/service/emby_user_data.go
@@ -0,0 +1,99 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "strconv"
+ "strings"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// SetFavorite 把 mediaID 标为 userID 的收藏。
+func (e *EmbyService) SetFavorite(ctx context.Context, userID, mediaID string, favorite bool) error {
+ if favorite {
+ var f model.Favorite
+ err := e.repo.DB.WithContext(ctx).
+ Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return e.repo.DB.WithContext(ctx).Create(&model.Favorite{
+ UserID: userID, MediaID: mediaID,
+ }).Error
+ }
+ return err
+ }
+ return e.repo.DB.WithContext(ctx).
+ Where("user_id = ? AND media_id = ?", userID, mediaID).
+ Delete(&model.Favorite{}).Error
+}
+
+// MarkPlayed 把 mediaID 标为已看(写一个 100% 进度的 history 行)。
+func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, played bool) error {
+ if !played {
+ return e.repo.DB.WithContext(ctx).
+ Where("user_id = ? AND media_id = ?", userID, mediaID).
+ Delete(&model.PlaybackHistory{}).Error
+ }
+ m, err := e.repo.Media.FindByID(ctx, mediaID)
+ if err != nil || m == nil {
+ return errors.New("media not found")
+ }
+ dur := int64(m.DurationSec) * 1000
+ if dur <= 0 {
+ dur = 1
+ }
+ return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
+ UserID: userID,
+ MediaID: mediaID,
+ PositionMs: dur,
+ DurationMs: dur,
+ WatchedAt: time.Now(),
+ Completed: true,
+ })
+}
+
+// RecordProgress 记录播放进度(来自 Emby 客户端的 /Sessions/Playing/Progress)。
+func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string, positionTicks, runtimeTicks int64) error {
+ pos := positionTicks / 10_000
+ dur := runtimeTicks / 10_000
+ if dur <= 0 {
+ // runtimeTicks 缺失时回退到 media.DurationSec
+ if m, _ := e.repo.Media.FindByID(ctx, mediaID); m != nil {
+ dur = int64(m.DurationSec) * 1000
+ }
+ }
+ completed := dur > 0 && pos >= dur*9/10
+ return e.repo.History.Upsert(ctx, &model.PlaybackHistory{
+ UserID: userID,
+ MediaID: mediaID,
+ PositionMs: pos,
+ DurationMs: dur,
+ WatchedAt: time.Now(),
+ Completed: completed,
+ })
+}
+
+func splitCSV(s string) []string {
+ if strings.TrimSpace(s) == "" {
+ return []string{}
+ }
+ parts := strings.Split(s, ",")
+ out := make([]string, 0, len(parts))
+ for _, p := range parts {
+ p = strings.TrimSpace(p)
+ if p != "" {
+ out = append(out, p)
+ }
+ }
+ return out
+}
+
+func intToStr(v int) string {
+ if v == 0 {
+ return ""
+ }
+ return strconv.Itoa(v)
+}
diff --git a/internal/service/emby_visibility.go b/internal/service/emby_visibility.go
new file mode 100644
index 0000000..2278a23
--- /dev/null
+++ b/internal/service/emby_visibility.go
@@ -0,0 +1,150 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (e *EmbyService) applyUserMediaVisibility(ctx context.Context, q *gorm.DB, userID string) *gorm.DB {
+ visibility := e.mediaVisibility(ctx, userID)
+ if !visibility.IncludeNSFW {
+ q = q.Where("nsfw = ?", false)
+ if hidden := visibility.HiddenLibraryIDs; len(hidden) > 0 {
+ q = q.Where("library_id NOT IN ?", hidden)
+ }
+ }
+ if len(visibility.AllowedLibraryIDs) > 0 {
+ q = q.Where("library_id IN ?", visibility.AllowedLibraryIDs)
+ }
+ return q
+}
+
+func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.Media, userID string) []model.Media {
+ visibility := e.mediaVisibility(ctx, userID)
+ if visibility.IncludeNSFW && len(visibility.AllowedLibraryIDs) == 0 {
+ return rows
+ }
+ allowed := map[string]bool{}
+ for _, id := range visibility.AllowedLibraryIDs {
+ allowed[id] = true
+ }
+ hiddenLibraries := map[string]bool{}
+ for _, id := range visibility.HiddenLibraryIDs {
+ hiddenLibraries[id] = true
+ }
+ out := rows[:0]
+ for _, row := range rows {
+ if row.NSFW && !visibility.IncludeNSFW {
+ continue
+ }
+ if hiddenLibraries[row.LibraryID] {
+ continue
+ }
+ if len(allowed) > 0 && !allowed[row.LibraryID] {
+ continue
+ }
+ out = append(out, row)
+ }
+ return out
+}
+
+func (e *EmbyService) mediaVisibility(ctx context.Context, userID string) MediaVisibility {
+ if e == nil {
+ return MediaVisibility{IncludeNSFW: true}
+ }
+ key := strings.TrimSpace(userID)
+ now := time.Now()
+ e.visibilityMu.RLock()
+ entry, ok := e.visibilityCache[key]
+ e.visibilityMu.RUnlock()
+ if ok && now.Before(entry.expiresAt) {
+ return cloneMediaVisibility(entry.visibility)
+ }
+
+ visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
+ if !visibility.IncludeNSFW {
+ visibility.HiddenLibraryIDs = e.hiddenLibraryIDs(ctx, visibility)
+ }
+ visibility = ExpandMediaVisibilityForMergedCloudLibraries(ctx, e.repo, visibility)
+ visibility = cloneMediaVisibility(visibility)
+
+ e.visibilityMu.Lock()
+ if e.visibilityCache == nil {
+ e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
+ }
+ if len(e.visibilityCache) > 1000 {
+ e.visibilityCache = make(map[string]embyVisibilityCacheEntry)
+ }
+ e.visibilityCache[key] = embyVisibilityCacheEntry{
+ visibility: cloneMediaVisibility(visibility),
+ expiresAt: now.Add(embyVisibilityCacheTTL),
+ }
+ e.visibilityMu.Unlock()
+
+ return visibility
+}
+
+func (e *EmbyService) mergedLibraryIDs(ctx context.Context, libraryID string) []string {
+ ids, err := MergedLibraryIDsForLibrary(ctx, e.repo, libraryID)
+ if err != nil || len(ids) == 0 {
+ return []string{libraryID}
+ }
+ return ids
+}
+
+func cloneMediaVisibility(visibility MediaVisibility) MediaVisibility {
+ if visibility.AllowedLibraryIDs != nil {
+ visibility.AllowedLibraryIDs = append([]string(nil), visibility.AllowedLibraryIDs...)
+ }
+ if visibility.HiddenLibraryIDs != nil {
+ visibility.HiddenLibraryIDs = append([]string(nil), visibility.HiddenLibraryIDs...)
+ }
+ return visibility
+}
+
+func (e *EmbyService) libraryVisibleFromCachedVisibility(lib model.Library, visibility MediaVisibility) bool {
+ if len(visibility.AllowedLibraryIDs) > 0 {
+ allowed := false
+ for _, id := range visibility.AllowedLibraryIDs {
+ if id == lib.ID {
+ allowed = true
+ break
+ }
+ }
+ if !allowed {
+ return false
+ }
+ }
+ if visibility.IncludeNSFW {
+ return true
+ }
+ for _, id := range visibility.HiddenLibraryIDs {
+ if id == lib.ID {
+ return false
+ }
+ }
+ return true
+}
+
+func (e *EmbyService) hiddenLibraryIDs(ctx context.Context, visibility MediaVisibility) []string {
+ if visibility.IncludeNSFW {
+ return nil
+ }
+ libs, err := e.repo.Library.List(ctx)
+ if err != nil {
+ return nil
+ }
+ shadowed := ShadowedCloudLibraryIDSet(libs)
+ ids := make([]string, 0)
+ for _, lib := range libs {
+ if shadowed[lib.ID] || !LibraryVisibleForUser(ctx, e.repo, lib, visibility) {
+ ids = append(ids, lib.ID)
+ }
+ }
+ return ids
+}
diff --git a/internal/service/episode_metadata_cleanup_test.go b/internal/service/episode_metadata_cleanup_test.go
index cef4395..e57d3f2 100644
--- a/internal/service/episode_metadata_cleanup_test.go
+++ b/internal/service/episode_metadata_cleanup_test.go
@@ -3,7 +3,6 @@ package service
import (
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
@@ -13,16 +12,10 @@ import (
func newCleanupTestContainer(t *testing.T) (*Container, *gorm.DB) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
+ db := newServiceTestDB(t, &model.Media{}, &model.Setting{})
if sqlDB, err := db.DB(); err == nil {
sqlDB.SetMaxOpenConns(1)
}
- if err := db.AutoMigrate(&model.Media{}, &model.Setting{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
repos := repository.New(db)
return &Container{Repo: repos, Log: zap.NewNop()}, db
}
diff --git a/internal/service/external_search.go b/internal/service/external_search.go
index b60fc3d..545efcd 100644
--- a/internal/service/external_search.go
+++ b/internal/service/external_search.go
@@ -25,6 +25,7 @@ type ExternalMediaResult struct {
DoubanID string `json:"douban_id,omitempty"`
TheTVDBID string `json:"thetvdb_id,omitempty"`
SubscribeKeyword string `json:"subscribe_keyword"`
+ SubscribeAliases []string `json:"subscribe_aliases,omitempty"`
TotalEpisodes int `json:"total_episodes,omitempty"`
DownloadedEpisodes int `json:"downloaded_episodes,omitempty"`
LocalMediaCount int `json:"local_media_count,omitempty"`
@@ -67,6 +68,7 @@ func SearchExternalMedia(ctx context.Context, query string, year int, mediaType
TMDbID: m.TMDbID,
BangumiID: m.BangumiID,
SubscribeKeyword: buildSubscribeKeyword(m.Title, m.Year),
+ SubscribeAliases: buildSubscribeAliases(m.Title, m.OriginalName, m.Year),
TotalEpisodes: totalEpisodes,
Languages: m.Languages,
Countries: m.Countries,
@@ -107,6 +109,7 @@ func SearchExternalMedia(ctx context.Context, query string, year int, mediaType
Rating: m.Rating,
DoubanID: m.DoubanID,
SubscribeKeyword: buildSubscribeKeyword(m.Title, yearValue),
+ SubscribeAliases: buildSubscribeAliases(m.Title, "", yearValue),
})
}
}
@@ -116,12 +119,25 @@ func SearchExternalMedia(ctx context.Context, query string, year int, mediaType
func buildSubscribeKeyword(title string, year int) string {
title = strings.TrimSpace(title)
+ if title == "" {
+ return ""
+ }
if year > 0 {
return fmt.Sprintf("%s %d", title, year)
}
return title
}
+func buildSubscribeAliases(title, originalName string, year int) []string {
+ values := []string{
+ title,
+ originalName,
+ buildSubscribeKeyword(title, year),
+ buildSubscribeKeyword(originalName, year),
+ }
+ return compactUniqueStrings(values...)
+}
+
func normalizeDoubanType(doubanType, fallback string) string {
doubanType = strings.ToLower(strings.TrimSpace(doubanType))
switch doubanType {
diff --git a/internal/service/external_search_test.go b/internal/service/external_search_test.go
index 56fced2..f4741a2 100644
--- a/internal/service/external_search_test.go
+++ b/internal/service/external_search_test.go
@@ -9,6 +9,21 @@ func TestBuildSubscribeKeyword(t *testing.T) {
if got := buildSubscribeKeyword("沙丘", 0); got != "沙丘" {
t.Fatalf("keyword without year = %q", got)
}
+ if got := buildSubscribeKeyword("", 2024); got != "" {
+ t.Fatalf("empty keyword = %q, want empty", got)
+ }
+}
+
+func TestBuildSubscribeAliasesIncludesOriginalTitleWithYear(t *testing.T) {
+ got := buildSubscribeAliases("玩具总动员 5", "Toy Story 5", 2026)
+ for _, want := range []string{"玩具总动员 5", "Toy Story 5", "玩具总动员 5 2026", "Toy Story 5 2026"} {
+ if !containsString(got, want) {
+ t.Fatalf("aliases = %#v, missing %q", got, want)
+ }
+ }
+ if containsString(got, "2026") {
+ t.Fatalf("aliases = %#v, must not include bare year", got)
+ }
}
func TestDedupeExternalMedia(t *testing.T) {
diff --git a/internal/service/ffmpeg_auto_install.go b/internal/service/ffmpeg_auto_install.go
index 59ed9b2..4f9ba45 100644
--- a/internal/service/ffmpeg_auto_install.go
+++ b/internal/service/ffmpeg_auto_install.go
@@ -414,7 +414,9 @@ func CheckFFmpegStatus(ffprobePath, ffmpegPath string) map[string]interface{} {
// 提取版本信息(第一行)
lines := bytes.Split(out, []byte("\n"))
if len(lines) > 0 {
- status["ffprobe_version"] = string(bytes.TrimSpace(lines[0]))
+ version := string(bytes.TrimSpace(lines[0]))
+ status["ffprobe_version"] = version
+ status["ffprobe_security"] = EvaluateFFmpegSecurity(version)
}
}
}
@@ -424,6 +426,17 @@ func CheckFFmpegStatus(ffprobePath, ffmpegPath string) map[string]interface{} {
if _, err := os.Stat(ffmpegPath); err == nil {
status["ffmpeg_installed"] = true
status["ffmpeg_path"] = ffmpegPath
+
+ cmd := exec.Command(ffmpegPath, "-version")
+ out, err := cmd.Output()
+ if err == nil {
+ lines := bytes.Split(out, []byte("\n"))
+ if len(lines) > 0 {
+ version := string(bytes.TrimSpace(lines[0]))
+ status["ffmpeg_version"] = version
+ status["ffmpeg_security"] = EvaluateFFmpegSecurity(version)
+ }
+ }
}
}
diff --git a/internal/service/ffmpeg_security.go b/internal/service/ffmpeg_security.go
new file mode 100644
index 0000000..3b5e7a0
--- /dev/null
+++ b/internal/service/ffmpeg_security.go
@@ -0,0 +1,73 @@
+package service
+
+import (
+ "fmt"
+ "regexp"
+ "strconv"
+ "strings"
+)
+
+const (
+ ffmpegPixelSmashCVE = "CVE-2026-8461"
+ ffmpegPixelSmashFixedVersion = "8.1.2"
+)
+
+var ffmpegVersionRE = regexp.MustCompile(`(?i)\bff(?:mpeg|probe)\s+version\s+([0-9]+)\.([0-9]+)(?:\.([0-9]+))?`)
+
+type FFmpegSecurityStatus struct {
+ Status string `json:"status"`
+ CVE string `json:"cve,omitempty"`
+ FixedVersion string `json:"fixed_version,omitempty"`
+ Message string `json:"message,omitempty"`
+ ParsedVersion string `json:"parsed_version,omitempty"`
+}
+
+func EvaluateFFmpegSecurity(versionLine string) FFmpegSecurityStatus {
+ major, minor, patch, parsed := parseFFmpegVersionLine(versionLine)
+ if !parsed {
+ return FFmpegSecurityStatus{
+ Status: "unknown",
+ CVE: ffmpegPixelSmashCVE,
+ FixedVersion: ffmpegPixelSmashFixedVersion,
+ Message: "无法识别 FFmpeg 版本;请确认 ffmpeg/ffprobe 已更新到 8.1.2 或更高版本",
+ }
+ }
+ status := FFmpegSecurityStatus{
+ Status: "ok",
+ CVE: ffmpegPixelSmashCVE,
+ FixedVersion: ffmpegPixelSmashFixedVersion,
+ ParsedVersion: fmt.Sprintf("%d.%d.%d", major, minor, patch),
+ }
+ if major == 8 && minor == 1 && patch < 2 {
+ status.Status = "vulnerable"
+ status.Message = "当前 FFmpeg 8.1.x 版本低于 8.1.2,可能受 PixelSmash 解码漏洞影响;请升级 ffmpeg/ffprobe"
+ return status
+ }
+ if major < 8 {
+ status.Status = "review"
+ status.Message = "当前 FFmpeg 主版本低于 8;请按发行版安全公告确认是否已回补 PixelSmash 修复"
+ }
+ return status
+}
+
+func parseFFmpegVersionLine(versionLine string) (major, minor, patch int, ok bool) {
+ match := ffmpegVersionRE.FindStringSubmatch(strings.TrimSpace(versionLine))
+ if len(match) < 3 {
+ return 0, 0, 0, false
+ }
+ major, err := strconv.Atoi(match[1])
+ if err != nil {
+ return 0, 0, 0, false
+ }
+ minor, err = strconv.Atoi(match[2])
+ if err != nil {
+ return 0, 0, 0, false
+ }
+ if len(match) > 3 && match[3] != "" {
+ patch, err = strconv.Atoi(match[3])
+ if err != nil {
+ return 0, 0, 0, false
+ }
+ }
+ return major, minor, patch, true
+}
diff --git a/internal/service/filemanager_test.go b/internal/service/filemanager_test.go
index a6852f9..d84d76b 100644
--- a/internal/service/filemanager_test.go
+++ b/internal/service/filemanager_test.go
@@ -6,9 +6,7 @@ import (
"path/filepath"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -17,13 +15,7 @@ import (
func newFileManagerTestServiceWithRepo(t *testing.T, root string) (*FileManagerService, *repository.Container) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "downloads", Path: root, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
diff --git a/internal/service/image_proxy.go b/internal/service/image_proxy.go
index 75f8c2a..3e815d8 100644
--- a/internal/service/image_proxy.go
+++ b/internal/service/image_proxy.go
@@ -13,60 +13,16 @@
package service
import (
- "bytes"
- "context"
- "crypto/sha256"
- "encoding/hex"
- "errors"
- "io"
- "net"
"net/http"
- "net/url"
- "os"
"path/filepath"
- "strings"
"sync"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
- "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
-// transparent1x1PNG is a baseline 67-byte PNG used as a fallback when the
-// upstream image cannot be retrieved, so browser layouts never collapse.
-var transparent1x1PNG = []byte{
- 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
- 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
- 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
- 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
- 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
- 0x54, 0x78, 0x9c, 0x63, 0x00, 0x01, 0x00, 0x00,
- 0x05, 0x00, 0x01, 0x0d, 0x0a, 0x2d, 0xb4, 0x00,
- 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae,
- 0x42, 0x60, 0x82,
-}
-
-// knownImageHosts are hosts we explicitly recognize. The list is no longer
-// a hard allow-list — it only short-circuits cases where we can be 100%
-// sure the destination is a public image CDN. Other hosts are accepted as
-// long as the scheme is http/https; this is required so users behind GFW
-// can configure their own TMDb mirror via secrets.tmdb_image_proxy.
-var knownImageHosts = map[string]struct{}{
- "image.tmdb.org": {},
- "www.themoviedb.org": {},
- "lain.bgm.tv": {},
- "img.bgm.tv": {},
- "webdav.bgm.tv": {},
- "img1.doubanio.com": {},
- "img2.doubanio.com": {},
- "img3.doubanio.com": {},
- "img9.doubanio.com": {},
- "assets.fanart.tv": {},
- "artworks.thetvdb.com": {},
-}
-
// ImageProxy fetches and caches remote images on behalf of the browser.
type ImageProxy struct {
cfg *config.Config
@@ -129,515 +85,3 @@ func (p *ImageProxy) libraryRoots() []string {
p.libRootsAt = time.Now()
return p.libRootsCache
}
-
-// validateURL parses raw and ensures the scheme is http/https and the
-// target host is not a private/loopback/link-local address (SSRF guard).
-func (p *ImageProxy) validateURL(raw string) (*url.URL, error) {
- if raw == "" {
- return nil, errors.New("missing url")
- }
- u, err := url.Parse(raw)
- if err != nil || u.Host == "" {
- return nil, errors.New("invalid url")
- }
- scheme := strings.ToLower(u.Scheme)
- if scheme != "http" && scheme != "https" {
- return nil, errors.New("unsupported scheme")
- }
- if isPrivateHost(u.Hostname()) {
- return nil, errors.New("requests to private/internal hosts are not allowed")
- }
- return u, nil
-}
-
-// isPrivateHost returns true only when host is a *literal* loopback, private,
-// link-local or unspecified IP address. This blocks the obvious SSRF vectors
-// (e.g. http://127.0.0.1/… or the cloud metadata IP 169.254.169.254) while
-// NOT blocking hostnames.
-//
-// We deliberately do not resolve hostnames here: under GFW DNS poisoning,
-// public image CDNs such as image.tmdb.org are frequently resolved to
-// loopback/private/bogus IPs. Blocking on resolved addresses would therefore
-// wrongly drop legitimate posters for exactly the users this proxy exists to
-// serve, which is what caused posters to stop displaying.
-func isPrivateHost(host string) bool {
- if host == "" {
- return true
- }
- ip := net.ParseIP(host)
- if ip != nil {
- return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified()
- }
- return false
-}
-
-// isAllowedLocalPath restricts local file reads to known-safe roots to
-// prevent arbitrary file reads via path traversal. Allowed roots are the
-// data dir, cache dir, the configured movies/tv/anime dirs, and — crucially —
-// every configured media library root, because sidecar posters/artwork are
-// stored alongside media under those (arbitrary, user-defined) paths.
-func (p *ImageProxy) isAllowedLocalPath(abs string) bool {
- roots := []string{p.cfg.App.DataDir, p.cfg.Cache.CacheDir, p.cfg.Media.MoviesDir, p.cfg.Media.TVDir, p.cfg.Media.AnimeDir}
- roots = append(roots, p.libraryRoots()...)
- for _, root := range roots {
- if strings.TrimSpace(root) == "" {
- continue
- }
- rootAbs, err := filepath.Abs(root)
- if err != nil {
- continue
- }
- if strings.HasPrefix(abs, rootAbs+string(filepath.Separator)) || abs == rootAbs {
- return true
- }
- }
- return false
-}
-
-func isLocalImagePath(raw string) bool {
- raw = strings.TrimSpace(raw)
- if raw == "" || isHTTPish(raw) {
- return false
- }
- ext := strings.ToLower(filepath.Ext(raw))
- switch ext {
- case ".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp":
- return true
- default:
- return false
- }
-}
-
-func isHTTPish(raw string) bool {
- return strings.HasPrefix(strings.ToLower(raw), "http://") || strings.HasPrefix(strings.ToLower(raw), "https://")
-}
-
-// detectContentType returns the MIME type of data using the first 512 bytes.
-func detectContentType(data []byte) string {
- if len(data) > 512 {
- return http.DetectContentType(data[:512])
- }
- return http.DetectContentType(data)
-}
-
-// servePlaceholder writes a 1×1 transparent PNG to w. Used as a fallback
-// when upstream fetch fails so the browser layout stays intact.
-func servePlaceholder(w http.ResponseWriter) {
- w.Header().Set("Content-Type", "image/png")
- w.Header().Set("Cache-Control", "no-store")
- w.WriteHeader(http.StatusOK)
- _, _ = w.Write(transparent1x1PNG)
-}
-
-func serveCachedPlaceholder(w http.ResponseWriter) {
- w.Header().Set("Content-Type", "image/png")
- w.Header().Set("Cache-Control", imagePlaceholderCacheControl)
- w.WriteHeader(http.StatusOK)
- _, _ = w.Write(transparent1x1PNG)
-}
-
-func (p *ImageProxy) cloudImageCachePaths(stableKey string) (string, string, string) {
- stableKey = strings.TrimSpace(stableKey)
- if stableKey == "" {
- stableKey = "unknown"
- }
- sum := sha256.Sum256([]byte("cloud-image:" + stableKey))
- key := "cloud-" + hex.EncodeToString(sum[:])
- cachePath := filepath.Join(p.cacheDir, key)
- return key, cachePath, cachePath + ".fail"
-}
-
-func serveCachedImageFile(w http.ResponseWriter, r *http.Request, key, cachePath string) bool {
- data, err := os.ReadFile(cachePath) // #nosec G304 -- cachePath is derived from a SHA-256 cache key under the configured cache directory.
- if err != nil || len(data) == 0 {
- return false
- }
- w.Header().Set("Content-Type", detectContentType(data))
- w.Header().Set("Cache-Control", imageBrowserCacheControl)
- stat, _ := os.Stat(cachePath)
- modTime := time.Now()
- if stat != nil {
- modTime = stat.ModTime()
- }
- http.ServeContent(w, r, key, modTime, bytes.NewReader(data))
- return true
-}
-
-func freshNegativeImageCache(failPath string) bool {
- stat, err := os.Stat(failPath)
- if err != nil {
- return false
- }
- if time.Since(stat.ModTime()) < imageNegativeCacheTTL {
- return true
- }
- _ = os.Remove(failPath)
- return false
-}
-
-// CloudImageCached reports whether a stable cloud-image ref already has a
-// usable positive or short-lived negative cache entry. Scanner pre-warm uses it
-// to avoid repeatedly resolving the same cloud sidecar image.
-func (p *ImageProxy) CloudImageCached(stableKey string) bool {
- if p == nil {
- return false
- }
- _, cachePath, failPath := p.cloudImageCachePaths(stableKey)
- if stat, err := os.Stat(cachePath); err == nil && stat.Size() > 0 {
- return true
- }
- return freshNegativeImageCache(failPath)
-}
-
-// ServeCloudCached serves an already-local cloud sidecar image without asking
-// the cloud provider for a fresh direct link. It returns true when it wrote a
-// response, including a fresh negative-cache placeholder.
-func (p *ImageProxy) ServeCloudCached(w http.ResponseWriter, r *http.Request, stableKey string) bool {
- if p == nil {
- return false
- }
- key, cachePath, failPath := p.cloudImageCachePaths(stableKey)
- if serveCachedImageFile(w, r, key, cachePath) {
- return true
- }
- if freshNegativeImageCache(failPath) {
- serveCachedPlaceholder(w)
- return true
- }
- return false
-}
-
-// Serve writes the requested image to w. Caller is expected to validate
-// the JWT before invoking it.
-func (p *ImageProxy) Serve(ctx context.Context, w http.ResponseWriter, r *http.Request, raw string) error {
- if isLocalImagePath(raw) {
- path := filepath.Clean(raw)
- abs, err := filepath.Abs(path)
- if err != nil || !p.isAllowedLocalPath(abs) {
- servePlaceholder(w)
- return nil
- }
- path = abs
- data, err := os.ReadFile(path)
- if err != nil || len(data) == 0 {
- servePlaceholder(w)
- return nil
- }
- stat, _ := os.Stat(path)
- modTime := time.Now()
- if stat != nil {
- modTime = stat.ModTime()
- }
- w.Header().Set("Content-Type", detectContentType(data))
- w.Header().Set("Cache-Control", imageBrowserCacheControl)
- http.ServeContent(w, r, filepath.Base(path), modTime, bytes.NewReader(data))
- return nil
- }
-
- u, err := p.validateURL(raw)
- if err != nil {
- // Bad URL is the only request-side error; everything else falls
- // through to the placeholder so the UI stays clean.
- return err
- }
- host := strings.ToLower(u.Host)
-
- // Cache key = sha256(url)
- sum := sha256.Sum256([]byte(raw))
- key := hex.EncodeToString(sum[:])
- cachePath := filepath.Join(p.cacheDir, key)
- failPath := cachePath + ".fail"
-
- // Cache hit.
- if data, err := os.ReadFile(cachePath); err == nil && len(data) > 0 { // #nosec G304 -- cachePath is derived from a SHA-256 cache key under the configured cache directory.
- w.Header().Set("Content-Type", detectContentType(data))
- w.Header().Set("Cache-Control", imageBrowserCacheControl)
- stat, _ := os.Stat(cachePath)
- modTime := time.Now()
- if stat != nil {
- modTime = stat.ModTime()
- }
- http.ServeContent(w, r, key, modTime, bytes.NewReader(data))
- return nil
- }
- if stat, err := os.Stat(failPath); err == nil && time.Since(stat.ModTime()) < imageNegativeCacheTTL {
- serveCachedPlaceholder(w)
- return nil
- } else if err == nil {
- _ = os.Remove(failPath)
- }
-
- // Cache miss → fetch upstream.
- if err := os.MkdirAll(p.cacheDir, 0o750); err != nil {
- p.log.Warn("imageproxy: mkdir failed", zap.String("dir", p.cacheDir), zap.Error(err))
- servePlaceholder(w)
- return nil
- }
-
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil)
- if err != nil {
- p.log.Warn("imageproxy: build request failed", zap.String("url", raw), zap.Error(err))
- servePlaceholder(w)
- return nil
- }
- req.Header.Set("User-Agent", "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/125.0 Safari/537.36")
- req.Header.Set("Accept", "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8")
- if strings.Contains(host, "doubanio.com") {
- req.Header.Set("Referer", "https://movie.douban.com/")
- }
-
- resp, err := p.client.Do(req)
- if err != nil {
- p.log.Warn("imageproxy: upstream fetch failed",
- zap.String("host", host), zap.Error(err))
- p.markImageFetchFailed(failPath)
- serveCachedPlaceholder(w)
- return nil
- }
- defer resp.Body.Close()
- if resp.StatusCode >= 400 {
- p.log.Warn("imageproxy: upstream returned non-OK",
- zap.String("host", host), zap.String("status", resp.Status))
- p.markImageFetchFailed(failPath)
- serveCachedPlaceholder(w)
- return nil
- }
-
- data, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20)) // 32 MiB cap
- if err != nil || len(data) == 0 {
- p.log.Warn("imageproxy: read upstream body failed",
- zap.String("host", host), zap.Error(err))
- p.markImageFetchFailed(failPath)
- serveCachedPlaceholder(w)
- return nil
- }
-
- // Write to a temp file then rename for atomicity.
- p.mu.Lock()
- tmp, tmpErr := os.CreateTemp(p.cacheDir, "img-*.tmp")
- if tmpErr == nil {
- if _, werr := tmp.Write(data); werr == nil {
- _ = tmp.Close()
- if rerr := os.Rename(tmp.Name(), cachePath); rerr != nil {
- _ = os.Remove(tmp.Name())
- } else {
- _ = os.Remove(failPath)
- }
- } else {
- _ = tmp.Close()
- _ = os.Remove(tmp.Name())
- }
- }
- p.mu.Unlock()
-
- ctype := resp.Header.Get("Content-Type")
- if ctype == "" {
- ctype = detectContentType(data)
- }
- w.Header().Set("Content-Type", ctype)
- if v := resp.Header.Get("Content-Length"); v != "" {
- w.Header().Set("Content-Length", v)
- }
- if v := resp.Header.Get("ETag"); v != "" {
- w.Header().Set("ETag", v)
- }
- if v := resp.Header.Get("Last-Modified"); v != "" {
- w.Header().Set("Last-Modified", v)
- }
- w.Header().Set("Cache-Control", imageBrowserCacheControl)
- http.ServeContent(w, r, key, time.Now(), bytes.NewReader(data))
- return nil
-}
-
-// ServeCloudResolved stores a cloud sidecar image in the same disk cache used
-// by remote posters, then serves it with long browser-cache headers. Cloud
-// direct links are often short-lived, so caching by the stable provider/ref
-// avoids re-resolving and re-downloading artwork every time the web UI or an
-// Emby-compatible client opens a library.
-func (p *ImageProxy) ServeCloudResolved(ctx context.Context, w http.ResponseWriter, r *http.Request, stableKey string, link *cloud.DirectLink) error {
- if p == nil || link == nil || strings.TrimSpace(link.URL) == "" {
- servePlaceholder(w)
- return nil
- }
- stableKey = strings.TrimSpace(stableKey)
- if stableKey == "" {
- stableKey = link.URL
- }
- key, cachePath, failPath := p.cloudImageCachePaths(stableKey)
-
- if serveCachedImageFile(w, r, key, cachePath) {
- return nil
- }
- if freshNegativeImageCache(failPath) {
- serveCachedPlaceholder(w)
- return nil
- }
-
- data, ctype, err := p.fetchAndCacheCloudImage(ctx, stableKey, link, r.UserAgent())
- if err != nil {
- p.log.Warn("imageproxy: cloud image fetch failed", zap.String("url", link.URL), zap.Error(err))
- serveCachedPlaceholder(w)
- return nil
- }
-
- w.Header().Set("Content-Type", ctype)
- w.Header().Set("Cache-Control", imageBrowserCacheControl)
- http.ServeContent(w, r, key, time.Now(), bytes.NewReader(data))
- return nil
-}
-
-// PrefetchCloudResolved downloads a cloud sidecar image into the local cache
-// without writing an HTTP response. It is intentionally best-effort; callers
-// should queue it with low concurrency so large cloud libraries do not overload
-// small NAS devices.
-func (p *ImageProxy) PrefetchCloudResolved(ctx context.Context, stableKey string, link *cloud.DirectLink) error {
- if p == nil || link == nil || strings.TrimSpace(link.URL) == "" {
- return nil
- }
- if p.CloudImageCached(stableKey) {
- return nil
- }
- _, _, err := p.fetchAndCacheCloudImage(ctx, stableKey, link, "MediaStationGo/0.1")
- return err
-}
-
-func (p *ImageProxy) fetchAndCacheCloudImage(ctx context.Context, stableKey string, link *cloud.DirectLink, userAgent string) ([]byte, string, error) {
- if err := os.MkdirAll(p.cacheDir, 0o750); err != nil {
- return nil, "", err
- }
- _, cachePath, failPath := p.cloudImageCachePaths(stableKey)
-
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, link.URL, nil)
- if err != nil {
- return nil, "", err
- }
- for k, v := range link.Headers {
- req.Header.Set(k, v)
- }
- if req.Header.Get("User-Agent") == "" {
- if strings.TrimSpace(userAgent) != "" {
- req.Header.Set("User-Agent", userAgent)
- } else {
- req.Header.Set("User-Agent", "MediaStationGo/0.1")
- }
- }
- req.Header.Set("Accept", "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8")
-
- resp, err := p.client.Do(req)
- if err != nil {
- p.markImageFetchFailed(failPath)
- return nil, "", err
- }
- defer resp.Body.Close()
- if resp.StatusCode >= 400 {
- p.markImageFetchFailed(failPath)
- return nil, "", errors.New("cloud image returned " + resp.Status)
- }
- data, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20))
- if err != nil {
- p.markImageFetchFailed(failPath)
- return nil, "", err
- }
- if len(data) == 0 {
- p.markImageFetchFailed(failPath)
- return nil, "", errors.New("cloud image body is empty")
- }
-
- p.mu.Lock()
- tmp, tmpErr := os.CreateTemp(p.cacheDir, "img-cloud-*.tmp")
- if tmpErr == nil {
- if _, werr := tmp.Write(data); werr == nil {
- _ = tmp.Close()
- if rerr := os.Rename(tmp.Name(), cachePath); rerr != nil {
- _ = os.Remove(tmp.Name())
- } else {
- _ = os.Remove(failPath)
- }
- } else {
- _ = tmp.Close()
- _ = os.Remove(tmp.Name())
- }
- }
- p.mu.Unlock()
-
- ctype := resp.Header.Get("Content-Type")
- if ctype == "" {
- ctype = detectContentType(data)
- }
- return data, ctype, nil
-}
-
-func (p *ImageProxy) markImageFetchFailed(failPath string) {
- if err := os.MkdirAll(filepath.Dir(failPath), 0o750); err != nil {
- return
- }
- p.mu.Lock()
- defer p.mu.Unlock()
- _ = os.WriteFile(failPath, []byte(time.Now().Format(time.RFC3339Nano)), 0o600)
-}
-
-// Fetch 拉取远程图片并返回字节和 Content-Type(带缓存)。
-func (p *ImageProxy) Fetch(ctx context.Context, raw string) ([]byte, string, error) {
- u, err := p.validateURL(raw)
- if err != nil {
- return nil, "", err
- }
-
- // Cache lookup
- sum := sha256.Sum256([]byte(raw))
- key := hex.EncodeToString(sum[:])
- cachePath := filepath.Join(p.cacheDir, key)
-
- if data, err := os.ReadFile(cachePath); err == nil && len(data) > 0 { // #nosec G304 -- cachePath is derived from a SHA-256 cache key under the configured cache directory.
- return data, detectContentType(data), nil
- }
-
- // Fetch upstream
- if err := os.MkdirAll(p.cacheDir, 0o750); err != nil {
- return nil, "", err
- }
-
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil)
- if err != nil {
- return nil, "", err
- }
- req.Header.Set("User-Agent", "MediaStationGo/0.1")
-
- resp, err := p.client.Do(req)
- if err != nil {
- return nil, "", err
- }
- defer resp.Body.Close()
- if resp.StatusCode >= 400 {
- return nil, "", errors.New("upstream returned " + resp.Status)
- }
-
- data, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20))
- if err != nil {
- return nil, "", err
- }
-
- // Write to cache
- p.mu.Lock()
- tmp, terr := os.CreateTemp(p.cacheDir, "img-*.tmp")
- if terr == nil {
- if _, werr := tmp.Write(data); werr == nil {
- _ = tmp.Close()
- if rerr := os.Rename(tmp.Name(), cachePath); rerr != nil {
- _ = os.Remove(tmp.Name())
- }
- } else {
- _ = tmp.Close()
- _ = os.Remove(tmp.Name())
- }
- }
- p.mu.Unlock()
-
- ctype := resp.Header.Get("Content-Type")
- if ctype == "" {
- ctype = detectContentType(data)
- }
- // host is unused here but referenced for log clarity in the future.
- _ = u
- return data, ctype, nil
-}
diff --git a/internal/service/image_proxy_cache.go b/internal/service/image_proxy_cache.go
new file mode 100644
index 0000000..468aa22
--- /dev/null
+++ b/internal/service/image_proxy_cache.go
@@ -0,0 +1,131 @@
+package service
+
+import (
+ "crypto/sha256"
+ "encoding/hex"
+ "io"
+ "net/http"
+ "os"
+ "path/filepath"
+ "strconv"
+ "strings"
+ "time"
+)
+
+// transparent1x1PNG is a baseline 67-byte PNG used as a fallback when the
+// upstream image cannot be retrieved, so browser layouts never collapse.
+var transparent1x1PNG = []byte{
+ 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
+ 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
+ 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
+ 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
+ 0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
+ 0x54, 0x78, 0x9c, 0x63, 0x00, 0x01, 0x00, 0x00,
+ 0x05, 0x00, 0x01, 0x0d, 0x0a, 0x2d, 0xb4, 0x00,
+ 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae,
+ 0x42, 0x60, 0x82,
+}
+
+// detectContentType returns the MIME type of data using the first 512 bytes.
+func detectContentType(data []byte) string {
+ if len(data) > 512 {
+ return http.DetectContentType(data[:512])
+ }
+ return http.DetectContentType(data)
+}
+
+// servePlaceholder writes a 1x1 transparent PNG to w. Used as a fallback
+// when upstream fetch fails so the browser layout stays intact.
+func servePlaceholder(w http.ResponseWriter) {
+ w.Header().Set("Content-Type", "image/png")
+ w.Header().Set("Cache-Control", "no-store")
+ w.WriteHeader(http.StatusOK)
+ _, _ = w.Write(transparent1x1PNG)
+}
+
+func serveCachedPlaceholder(w http.ResponseWriter) {
+ w.Header().Set("Content-Type", "image/png")
+ w.Header().Set("Cache-Control", imagePlaceholderCacheControl)
+ w.WriteHeader(http.StatusOK)
+ _, _ = w.Write(transparent1x1PNG)
+}
+
+func (p *ImageProxy) cloudImageCachePaths(stableKey string) (string, string, string) {
+ stableKey = strings.TrimSpace(stableKey)
+ if stableKey == "" {
+ stableKey = "unknown"
+ }
+ sum := sha256.Sum256([]byte("cloud-image:" + stableKey))
+ key := "cloud-" + hex.EncodeToString(sum[:])
+ cachePath := filepath.Join(p.cacheDir, key)
+ return key, cachePath, cachePath + ".fail"
+}
+
+func (p *ImageProxy) remoteImageCachePaths(raw string) (string, string, string, error) {
+ if _, err := p.validateURL(raw); err != nil {
+ return "", "", "", err
+ }
+ key, cachePath, failPath := p.remoteImageCachePathsForValidated(raw)
+ return key, cachePath, failPath, nil
+}
+
+func (p *ImageProxy) remoteImageCachePathsForValidated(raw string) (string, string, string) {
+ sum := sha256.Sum256([]byte(raw))
+ key := hex.EncodeToString(sum[:])
+ cachePath := filepath.Join(p.cacheDir, key)
+ return key, cachePath, cachePath + ".fail"
+}
+
+func serveCachedImageFile(w http.ResponseWriter, r *http.Request, key, cachePath string) bool {
+ return serveImageFile(w, r, key, cachePath, imageBrowserCacheControl)
+}
+
+func serveImageFile(w http.ResponseWriter, r *http.Request, key, path, cacheControl string) bool {
+ file, err := os.Open(path) // #nosec G304 -- caller only passes validated local paths or SHA-derived cache paths.
+ if err != nil {
+ return false
+ }
+ defer file.Close()
+ stat, err := file.Stat()
+ if err != nil || stat.IsDir() || stat.Size() <= 0 {
+ return false
+ }
+ var sample [512]byte
+ n, _ := file.Read(sample[:])
+ _, _ = file.Seek(0, io.SeekStart)
+ w.Header().Set("Content-Type", detectContentType(sample[:n]))
+ w.Header().Set("Cache-Control", cacheControl)
+ w.Header().Set("ETag", imageFileETag(key, stat))
+ http.ServeContent(w, r, key, stat.ModTime(), file)
+ return true
+}
+
+func imageFileETag(key string, stat os.FileInfo) string {
+ key = strings.TrimSpace(key)
+ if key == "" {
+ key = "image"
+ }
+ sum := sha256.Sum256([]byte(key))
+ return `"img-` + hex.EncodeToString(sum[:8]) + "-" + strconv.FormatInt(stat.Size(), 16) + "-" + strconv.FormatInt(stat.ModTime().Unix(), 16) + `"`
+}
+
+func freshNegativeImageCache(failPath string) bool {
+ stat, err := os.Stat(failPath)
+ if err != nil {
+ return false
+ }
+ if time.Since(stat.ModTime()) < imageNegativeCacheTTL {
+ return true
+ }
+ _ = os.Remove(failPath)
+ return false
+}
+
+func (p *ImageProxy) markImageFetchFailed(failPath string) {
+ if err := os.MkdirAll(filepath.Dir(failPath), 0o750); err != nil {
+ return
+ }
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ _ = os.WriteFile(failPath, []byte(time.Now().Format(time.RFC3339Nano)), 0o600)
+}
diff --git a/internal/service/image_proxy_cloud.go b/internal/service/image_proxy_cloud.go
new file mode 100644
index 0000000..a04f070
--- /dev/null
+++ b/internal/service/image_proxy_cloud.go
@@ -0,0 +1,144 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "io"
+ "net/http"
+ "os"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+// CloudImageCached reports whether a stable cloud-image ref already has a
+// usable positive or short-lived negative cache entry. Scanner pre-warm uses it
+// to avoid repeatedly resolving the same cloud sidecar image.
+func (p *ImageProxy) CloudImageCached(stableKey string) bool {
+ if p == nil {
+ return false
+ }
+ _, cachePath, failPath := p.cloudImageCachePaths(stableKey)
+ if stat, err := os.Stat(cachePath); err == nil && stat.Size() > 0 {
+ return true
+ }
+ return freshNegativeImageCache(failPath)
+}
+
+// ServeCloudCached serves an already-local cloud sidecar image without asking
+// the cloud provider for a fresh direct link. It returns true when it wrote a
+// response, including a fresh negative-cache placeholder.
+func (p *ImageProxy) ServeCloudCached(w http.ResponseWriter, r *http.Request, stableKey string) bool {
+ if p == nil {
+ return false
+ }
+ key, cachePath, failPath := p.cloudImageCachePaths(stableKey)
+ if serveCachedImageFile(w, r, key, cachePath) {
+ return true
+ }
+ if freshNegativeImageCache(failPath) {
+ serveCachedPlaceholder(w)
+ return true
+ }
+ return false
+}
+
+// ServeCloudResolved stores a cloud sidecar image in the same disk cache used
+// by remote posters, then serves it with long browser-cache headers.
+func (p *ImageProxy) ServeCloudResolved(ctx context.Context, w http.ResponseWriter, r *http.Request, stableKey string, link *cloud.DirectLink) error {
+ if p == nil || link == nil || strings.TrimSpace(link.URL) == "" {
+ servePlaceholder(w)
+ return nil
+ }
+ stableKey = strings.TrimSpace(stableKey)
+ if stableKey == "" {
+ stableKey = link.URL
+ }
+ key, cachePath, failPath := p.cloudImageCachePaths(stableKey)
+ if serveCachedImageFile(w, r, key, cachePath) {
+ return nil
+ }
+ if freshNegativeImageCache(failPath) {
+ serveCachedPlaceholder(w)
+ return nil
+ }
+ data, ctype, err := p.fetchAndCacheCloudImage(ctx, stableKey, link, r.UserAgent())
+ if err != nil {
+ p.log.Warn("imageproxy: cloud image fetch failed", zap.String("url", link.URL), zap.Error(err))
+ serveCachedPlaceholder(w)
+ return nil
+ }
+ w.Header().Set("Content-Type", ctype)
+ w.Header().Set("Cache-Control", imageBrowserCacheControl)
+ modTime := time.Now()
+ if stat, err := os.Stat(cachePath); err == nil && stat.Size() > 0 {
+ modTime = stat.ModTime()
+ w.Header().Set("ETag", imageFileETag(key, stat))
+ }
+ http.ServeContent(w, r, key, modTime, bytes.NewReader(data))
+ return nil
+}
+
+// PrefetchCloudResolved downloads a cloud sidecar image into the local cache
+// without writing an HTTP response.
+func (p *ImageProxy) PrefetchCloudResolved(ctx context.Context, stableKey string, link *cloud.DirectLink) error {
+ if p == nil || link == nil || strings.TrimSpace(link.URL) == "" {
+ return nil
+ }
+ if p.CloudImageCached(stableKey) {
+ return nil
+ }
+ _, _, err := p.fetchAndCacheCloudImage(ctx, stableKey, link, "MediaStationGo/0.1")
+ return err
+}
+
+func (p *ImageProxy) fetchAndCacheCloudImage(ctx context.Context, stableKey string, link *cloud.DirectLink, userAgent string) ([]byte, string, error) {
+ if err := os.MkdirAll(p.cacheDir, 0o750); err != nil {
+ return nil, "", err
+ }
+ _, cachePath, failPath := p.cloudImageCachePaths(stableKey)
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, link.URL, nil)
+ if err != nil {
+ return nil, "", err
+ }
+ for k, v := range link.Headers {
+ req.Header.Set(k, v)
+ }
+ if req.Header.Get("User-Agent") == "" {
+ if strings.TrimSpace(userAgent) != "" {
+ req.Header.Set("User-Agent", userAgent)
+ } else {
+ req.Header.Set("User-Agent", "MediaStationGo/0.1")
+ }
+ }
+ req.Header.Set("Accept", "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8")
+ resp, err := p.client.Do(req)
+ if err != nil {
+ p.markImageFetchFailed(failPath)
+ return nil, "", err
+ }
+ defer resp.Body.Close()
+ if resp.StatusCode >= 400 {
+ p.markImageFetchFailed(failPath)
+ return nil, "", errors.New("cloud image returned " + resp.Status)
+ }
+ data, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20))
+ if err != nil {
+ p.markImageFetchFailed(failPath)
+ return nil, "", err
+ }
+ if len(data) == 0 {
+ p.markImageFetchFailed(failPath)
+ return nil, "", errors.New("cloud image body is empty")
+ }
+ p.writeImageCache(cachePath, failPath, "img-cloud-*.tmp", data)
+ ctype := resp.Header.Get("Content-Type")
+ if ctype == "" {
+ ctype = detectContentType(data)
+ }
+ return data, ctype, nil
+}
diff --git a/internal/service/image_proxy_paths.go b/internal/service/image_proxy_paths.go
new file mode 100644
index 0000000..e73280a
--- /dev/null
+++ b/internal/service/image_proxy_paths.go
@@ -0,0 +1,80 @@
+package service
+
+import (
+ "errors"
+ "net"
+ "net/url"
+ "path/filepath"
+ "strings"
+)
+
+// validateURL parses raw and ensures the scheme is http/https and the
+// target host is not a private/loopback/link-local address (SSRF guard).
+func (p *ImageProxy) validateURL(raw string) (*url.URL, error) {
+ if raw == "" {
+ return nil, errors.New("missing url")
+ }
+ u, err := url.Parse(raw)
+ if err != nil || u.Host == "" {
+ return nil, errors.New("invalid url")
+ }
+ scheme := strings.ToLower(u.Scheme)
+ if scheme != "http" && scheme != "https" {
+ return nil, errors.New("unsupported scheme")
+ }
+ if isPrivateHost(u.Hostname()) {
+ return nil, errors.New("requests to private/internal hosts are not allowed")
+ }
+ return u, nil
+}
+
+// isPrivateHost returns true only when host is a literal loopback, private,
+// link-local or unspecified IP address. Hostnames are not resolved here because
+// DNS poisoning can map public image CDNs to bogus private addresses.
+func isPrivateHost(host string) bool {
+ if host == "" {
+ return true
+ }
+ ip := net.ParseIP(host)
+ if ip != nil {
+ return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified()
+ }
+ return false
+}
+
+// isAllowedLocalPath restricts local file reads to known-safe roots.
+func (p *ImageProxy) isAllowedLocalPath(abs string) bool {
+ roots := []string{p.cfg.App.DataDir, p.cfg.Cache.CacheDir, p.cfg.Media.MoviesDir, p.cfg.Media.TVDir, p.cfg.Media.AnimeDir}
+ roots = append(roots, p.libraryRoots()...)
+ for _, root := range roots {
+ if strings.TrimSpace(root) == "" {
+ continue
+ }
+ rootAbs, err := filepath.Abs(root)
+ if err != nil {
+ continue
+ }
+ if strings.HasPrefix(abs, rootAbs+string(filepath.Separator)) || abs == rootAbs {
+ return true
+ }
+ }
+ return false
+}
+
+func isLocalImagePath(raw string) bool {
+ raw = strings.TrimSpace(raw)
+ if raw == "" || isHTTPish(raw) {
+ return false
+ }
+ ext := strings.ToLower(filepath.Ext(raw))
+ switch ext {
+ case ".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn":
+ return true
+ default:
+ return false
+ }
+}
+
+func isHTTPish(raw string) bool {
+ return strings.HasPrefix(strings.ToLower(raw), "http://") || strings.HasPrefix(strings.ToLower(raw), "https://")
+}
diff --git a/internal/service/image_proxy_remote.go b/internal/service/image_proxy_remote.go
new file mode 100644
index 0000000..1add929
--- /dev/null
+++ b/internal/service/image_proxy_remote.go
@@ -0,0 +1,189 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "io"
+ "net/http"
+ "os"
+ "path/filepath"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+var errImageProxyRequestSetup = errors.New("image proxy request setup failed")
+
+func (p *ImageProxy) PrefetchRemote(ctx context.Context, raw string) error {
+ _, _, err := p.Fetch(ctx, raw)
+ return err
+}
+
+func (p *ImageProxy) RemoveCached(raw string) error {
+ if !isHTTPish(raw) {
+ return nil
+ }
+ _, cachePath, failPath, err := p.remoteImageCachePaths(raw)
+ if err != nil {
+ return nil
+ }
+ if err := os.Remove(cachePath); err != nil && !errors.Is(err, os.ErrNotExist) {
+ return err
+ }
+ if err := os.Remove(failPath); err != nil && !errors.Is(err, os.ErrNotExist) {
+ return err
+ }
+ return nil
+}
+
+// Serve writes the requested image to w. Caller is expected to validate
+// the JWT before invoking it.
+func (p *ImageProxy) Serve(ctx context.Context, w http.ResponseWriter, r *http.Request, raw string) error {
+ if isLocalImagePath(raw) {
+ return p.serveLocalImage(w, r, raw)
+ }
+ return p.serveRemoteImage(ctx, w, r, raw)
+}
+
+func (p *ImageProxy) serveLocalImage(w http.ResponseWriter, r *http.Request, raw string) error {
+ path := filepath.Clean(raw)
+ abs, err := filepath.Abs(path)
+ if err != nil || !p.isAllowedLocalPath(abs) {
+ servePlaceholder(w)
+ return nil
+ }
+ if !serveImageFile(w, r, filepath.Base(abs), abs, imageBrowserCacheControl) {
+ servePlaceholder(w)
+ }
+ return nil
+}
+
+func (p *ImageProxy) serveRemoteImage(ctx context.Context, w http.ResponseWriter, r *http.Request, raw string) error {
+ u, err := p.validateURL(raw)
+ if err != nil {
+ return err
+ }
+ host := strings.ToLower(u.Host)
+ key, cachePath, failPath := p.remoteImageCachePathsForValidated(raw)
+ if serveCachedImageFile(w, r, key, cachePath) {
+ return nil
+ }
+ if p.serveFreshRemoteFailure(w, failPath) {
+ return nil
+ }
+ data, ctype, contentLength, err := p.fetchAndCacheRemoteImage(ctx, raw, host, cachePath, failPath)
+ if err != nil {
+ if errors.Is(err, errImageProxyRequestSetup) {
+ servePlaceholder(w)
+ } else {
+ serveCachedPlaceholder(w)
+ }
+ return nil
+ }
+ w.Header().Set("Content-Type", ctype)
+ if contentLength != "" {
+ w.Header().Set("Content-Length", contentLength)
+ }
+ modTime := time.Now()
+ if stat, err := os.Stat(cachePath); err == nil && stat.Size() > 0 {
+ modTime = stat.ModTime()
+ w.Header().Set("ETag", imageFileETag(key, stat))
+ }
+ w.Header().Set("Cache-Control", imageBrowserCacheControl)
+ http.ServeContent(w, r, key, modTime, bytes.NewReader(data))
+ return nil
+}
+
+func (p *ImageProxy) serveFreshRemoteFailure(w http.ResponseWriter, failPath string) bool {
+ if stat, err := os.Stat(failPath); err == nil && time.Since(stat.ModTime()) < imageNegativeCacheTTL {
+ serveCachedPlaceholder(w)
+ return true
+ } else if err == nil {
+ _ = os.Remove(failPath)
+ }
+ return false
+}
+
+func (p *ImageProxy) fetchAndCacheRemoteImage(ctx context.Context, raw, host, cachePath, failPath string) ([]byte, string, string, error) {
+ if err := os.MkdirAll(p.cacheDir, 0o750); err != nil {
+ p.log.Warn("imageproxy: mkdir failed", zap.String("dir", p.cacheDir), zap.Error(err))
+ return nil, "", "", errImageProxyRequestSetup
+ }
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil)
+ if err != nil {
+ p.log.Warn("imageproxy: build request failed", zap.String("url", raw), zap.Error(err))
+ return nil, "", "", errImageProxyRequestSetup
+ }
+ req.Header.Set("User-Agent", "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/125.0 Safari/537.36")
+ req.Header.Set("Accept", "image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8")
+ if strings.Contains(host, "doubanio.com") {
+ req.Header.Set("Referer", "https://movie.douban.com/")
+ }
+ resp, err := p.client.Do(req)
+ if err != nil {
+ p.log.Warn("imageproxy: upstream fetch failed", zap.String("host", host), zap.Error(err))
+ p.markImageFetchFailed(failPath)
+ return nil, "", "", err
+ }
+ defer resp.Body.Close()
+ if resp.StatusCode >= 400 {
+ p.log.Warn("imageproxy: upstream returned non-OK", zap.String("host", host), zap.String("status", resp.Status))
+ p.markImageFetchFailed(failPath)
+ return nil, "", "", errors.New("upstream returned " + resp.Status)
+ }
+ data, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20))
+ if err != nil || len(data) == 0 {
+ p.log.Warn("imageproxy: read upstream body failed", zap.String("host", host), zap.Error(err))
+ p.markImageFetchFailed(failPath)
+ if err == nil {
+ err = errors.New("upstream image body is empty")
+ }
+ return nil, "", "", err
+ }
+ p.writeImageCache(cachePath, failPath, "img-*.tmp", data)
+ ctype := resp.Header.Get("Content-Type")
+ if ctype == "" {
+ ctype = detectContentType(data)
+ }
+ return data, ctype, resp.Header.Get("Content-Length"), nil
+}
+
+// Fetch pulls a remote image and returns bytes plus Content-Type using cache.
+func (p *ImageProxy) Fetch(ctx context.Context, raw string) ([]byte, string, error) {
+ if _, err := p.validateURL(raw); err != nil {
+ return nil, "", err
+ }
+ _, cachePath, failPath := p.remoteImageCachePathsForValidated(raw)
+ if data, err := os.ReadFile(cachePath); err == nil && len(data) > 0 { // #nosec G304 -- cachePath is SHA-derived under cacheDir.
+ return data, detectContentType(data), nil
+ }
+ if stat, err := os.Stat(failPath); err == nil && time.Since(stat.ModTime()) < imageNegativeCacheTTL {
+ return nil, "", errors.New("recent image fetch failure")
+ } else if err == nil {
+ _ = os.Remove(failPath)
+ }
+ data, ctype, _, err := p.fetchAndCacheRemoteImage(ctx, raw, "", cachePath, failPath)
+ return data, ctype, err
+}
+
+func (p *ImageProxy) writeImageCache(cachePath, failPath, pattern string, data []byte) {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ tmp, tmpErr := os.CreateTemp(p.cacheDir, pattern)
+ if tmpErr != nil {
+ return
+ }
+ if _, err := tmp.Write(data); err != nil {
+ _ = tmp.Close()
+ _ = os.Remove(tmp.Name())
+ return
+ }
+ _ = tmp.Close()
+ if err := os.Rename(tmp.Name(), cachePath); err != nil {
+ _ = os.Remove(tmp.Name())
+ return
+ }
+ _ = os.Remove(failPath)
+}
diff --git a/internal/service/image_proxy_test.go b/internal/service/image_proxy_test.go
index e7cc9fb..0234594 100644
--- a/internal/service/image_proxy_test.go
+++ b/internal/service/image_proxy_test.go
@@ -80,6 +80,22 @@ func TestImageProxyServesPosterUnderLibraryRoot(t *testing.T) {
if got := rec.Body.Bytes(); string(got) != string(realPoster) {
t.Fatalf("served %q, want real poster bytes", string(got))
}
+ etag := rec.Header().Get("ETag")
+ if etag == "" {
+ t.Fatal("expected static image ETag")
+ }
+ req := httptest.NewRequest(http.MethodGet, "/api/img", nil)
+ req.Header.Set("If-None-Match", etag)
+ rec = httptest.NewRecorder()
+ if err := proxy.Serve(t.Context(), rec, req, posterPath); err != nil {
+ t.Fatal(err)
+ }
+ if rec.Code != http.StatusNotModified {
+ t.Fatalf("conditional status = %d, want 304", rec.Code)
+ }
+ if rec.Body.Len() != 0 {
+ t.Fatalf("conditional body length = %d, want 0", rec.Body.Len())
+ }
}
func TestImageProxyCachesFailedRemoteImageFetch(t *testing.T) {
diff --git a/internal/service/local_metadata.go b/internal/service/local_metadata.go
index fdc7da4..ae205ff 100644
--- a/internal/service/local_metadata.go
+++ b/internal/service/local_metadata.go
@@ -3,11 +3,6 @@ package service
import (
"encoding/xml"
"errors"
- "image"
- _ "image/gif"
- _ "image/jpeg"
- _ "image/png"
- "net/url"
"os"
"path/filepath"
"strconv"
@@ -18,6 +13,7 @@ import (
type LocalMetadata struct {
Title string
OriginalName string
+ EpisodeTitle string
AdultCode string
Year int
Overview string
@@ -311,6 +307,9 @@ func metadataFromDoc(doc *nfoDocument, baseDir string, seriesLike bool) *LocalMe
Languages: joinNFOValues(doc.Languages),
HasNFO: true,
}
+ if nfoIsEpisodeDetails(doc) {
+ meta.EpisodeTitle = cleanXMLText(doc.Title)
+ }
if meta.AdultCode == "" {
meta.AdultCode = normalizeAdultCode(firstText(doc.OriginalTitle, doc.SortTitle, doc.Title))
}
@@ -429,6 +428,9 @@ func mergeEpisodeMetadata(dst, episode *LocalMetadata, doc *nfoDocument) {
}
}
// 注意: 不要把单集名 / 单集 originaltitle 写进 OriginalName(整剧原名,分组键)。
+ if episodeTitle := firstText(episode.EpisodeTitle, doc.Title); episodeTitle != "" && !strings.EqualFold(episodeTitle, showTitle) {
+ dst.EpisodeTitle = episodeTitle
+ }
// 单集级展示字段: 每个媒体行本就对应一集,这些可安全按集回填。
if episode.Year > 0 {
@@ -465,6 +467,13 @@ func mergeEpisodeMetadata(dst, episode *LocalMetadata, doc *nfoDocument) {
}
}
+func nfoIsEpisodeDetails(doc *nfoDocument) bool {
+ if doc == nil {
+ return false
+ }
+ return strings.EqualFold(strings.TrimSpace(doc.XMLName.Local), "episodedetails")
+}
+
func tmdbIDFromUniqueIDs(ids []nfoUniqueID) int {
value := externalIDFromUniqueIDs(ids, "tmdb")
if value == "" {
@@ -508,243 +517,6 @@ func firstRemoteURL(baseDir string, values ...string) string {
return ""
}
-func localPosterCandidates(mediaPath string) []string {
- base := strings.TrimSuffix(filepath.Base(mediaPath), filepath.Ext(mediaPath))
- names := []string{
- base + "-poster",
- base + ".poster",
- "poster",
- "folder",
- "cover",
- "movie",
- "show",
- base + "-cover",
- base + ".cover",
- base,
- base + "-thumb",
- base + ".thumb",
- "thumb",
- }
- return append(adultArtworkNameCandidates(mediaPath, "poster"), names...)
-}
-
-func localBackdropCandidates(mediaPath string) []string {
- base := strings.TrimSuffix(filepath.Base(mediaPath), filepath.Ext(mediaPath))
- names := []string{
- base + "-fanart",
- base + ".fanart",
- base + "-backdrop",
- base + ".backdrop",
- base + "-background",
- "fanart",
- "backdrop",
- "background",
- "landscape",
- "banner",
- "clearart",
- }
- return append(adultArtworkNameCandidates(mediaPath, "backdrop"), names...)
-}
-
-func adultArtworkNameCandidates(mediaPath, kind string) []string {
- code := AdultCodeFromMediaPath(mediaPath)
- if code == "" {
- return nil
- }
- compact := strings.ReplaceAll(code, "-", "")
- bases := []string{code, compact}
- bases = append(bases, adultDMMNameCandidates(code)...)
- out := make([]string, 0, len(bases)*6)
- for _, base := range bases {
- if base == "" {
- continue
- }
- if kind == "poster" {
- out = append(out, base, base+"-poster", base+".poster", base+"-cover", base+".cover", base+"-thumb", base+".thumb", base+"pl", base+"-pl")
- } else {
- out = append(out, base+"-fanart", base+".fanart", base+"-backdrop", base+".backdrop", base+"-background", base+"-landscape", base+"jp", base+"jp-1")
- }
- }
- return out
-}
-
-func adultDMMNameCandidates(code string) []string {
- parts := adultStandardPattern.FindStringSubmatch(code)
- if len(parts) < 3 {
- return nil
- }
- prefix := strings.ToLower(parts[1])
- num := strings.TrimLeft(parts[2], "0")
- if num == "" {
- num = "0"
- }
- padded := num
- for len(padded) < 5 {
- padded = "0" + padded
- }
- return []string{prefix + padded}
-}
-
-func firstExistingImage(dir string, names ...string) string {
- if dir == "" {
- return ""
- }
- for _, name := range names {
- for _, ext := range []string{".jpg", ".jpeg", ".png", ".webp"} {
- path := filepath.Join(dir, name+ext)
- if fileExists(path) {
- return filepath.Clean(path)
- }
- }
- }
- return ""
-}
-
-func nfoPosterValues(doc *nfoDocument) []string {
- if doc == nil {
- return nil
- }
- values := []string{doc.Poster, doc.Art.Poster}
- for _, thumb := range doc.Thumbs {
- aspect := strings.ToLower(strings.TrimSpace(thumb.Aspect))
- if aspect == "" || aspect == "poster" || aspect == "cover" || aspect == "default" {
- values = append(values, thumb.Value)
- }
- }
- values = append(values, doc.Art.Thumb)
- return values
-}
-
-func nfoBackdropValues(doc *nfoDocument) []string {
- if doc == nil {
- return nil
- }
- values := []string{doc.Fanart.Value, doc.Art.Fanart, doc.Art.Backdrop, doc.Art.Background, doc.Art.Landscape, doc.Art.Banner}
- for _, thumb := range doc.Thumbs {
- aspect := strings.ToLower(strings.TrimSpace(thumb.Aspect))
- if aspect == "fanart" || aspect == "backdrop" || aspect == "background" || aspect == "landscape" {
- values = append(values, thumb.Value)
- }
- }
- values = append(values, doc.Fanart.Thumbs...)
- return values
-}
-
-func firstLocalPoster(mediaPath, showBaseDir string) string {
- mediaDir := filepath.Dir(mediaPath)
- dirs := []string{}
- if showBaseDir != "" && !samePath(showBaseDir, mediaDir) {
- dirs = append(dirs, showBaseDir)
- }
- dirs = append(dirs, mediaDir)
- for _, dir := range dirs {
- if localPoster := firstExistingPosterImage(dir, localPosterCandidates(mediaPath)...); localPoster != "" {
- return localPoster
- }
- }
- return ""
-}
-
-func firstExistingPosterImage(dir string, names ...string) string {
- if dir == "" {
- return ""
- }
- for _, name := range names {
- if isRejectedPosterName(name) {
- continue
- }
- for _, ext := range []string{".jpg", ".jpeg", ".png", ".webp"} {
- path := filepath.Join(dir, name+ext)
- if fileExists(path) && likelyPosterImage(path) {
- return filepath.Clean(path)
- }
- }
- }
- return ""
-}
-
-func firstAdultLooseImage(dir, kind string) string {
- if dir == "" {
- return ""
- }
- matches, _ := filepath.Glob(filepath.Join(dir, "*"))
- preferred := []string{}
- fallback := []string{}
- for _, path := range matches {
- ext := strings.ToLower(filepath.Ext(path))
- if ext != ".jpg" && ext != ".jpeg" && ext != ".png" && ext != ".webp" {
- continue
- }
- name := strings.ToLower(strings.TrimSuffix(filepath.Base(path), ext))
- if kind == "poster" {
- if isRejectedPosterName(name) {
- continue
- }
- if strings.Contains(name, "poster") || strings.Contains(name, "cover") || strings.Contains(name, "folder") || strings.Contains(name, "movie") || strings.HasSuffix(name, "pl") {
- preferred = append(preferred, path)
- }
- } else if strings.Contains(name, "fanart") || strings.Contains(name, "backdrop") || strings.Contains(name, "background") || strings.Contains(name, "landscape") || strings.Contains(name, "jp") {
- preferred = append(preferred, path)
- }
- if kind != "poster" || likelyPosterImage(path) {
- fallback = append(fallback, path)
- }
- }
- if len(preferred) > 0 {
- return filepath.Clean(preferred[0])
- }
- if kind == "poster" && len(fallback) == 1 {
- return filepath.Clean(fallback[0])
- }
- return ""
-}
-
-func isRejectedPosterName(name string) bool {
- name = strings.ToLower(name)
- rejected := []string{
- "actor", "actors", "actress", "cast", "avatar", "portrait", "person",
- "sample", "screenshot", "screen", "still", "scene", "extrafanart", "extrathumb",
- "fanart", "backdrop", "background", "landscape", "banner", "clearlogo", "clearart", "logo", "disc",
- }
- for _, token := range rejected {
- if strings.Contains(name, token) {
- return true
- }
- }
- return false
-}
-
-func likelyPosterImage(path string) bool {
- file, err := os.Open(path) // #nosec G304 -- path is a discovered artwork sidecar under the configured library root.
- if err != nil {
- return false
- }
- defer file.Close()
- cfg, _, err := image.DecodeConfig(file)
- if err != nil || cfg.Width <= 0 || cfg.Height <= 0 {
- return true
- }
- return cfg.Height >= cfg.Width
-}
-
-func fileExists(path string) bool {
- info, err := os.Stat(path)
- return err == nil && !info.IsDir()
-}
-
-func isHTTPURL(raw string) bool {
- u, err := url.Parse(raw)
- if err != nil {
- return false
- }
- return (u.Scheme == "http" || u.Scheme == "https") && u.Host != ""
-}
-
-func isLocalPath(raw string) bool {
- raw = strings.TrimSpace(raw)
- return raw != "" && !isHTTPURL(raw)
-}
-
func firstText(values ...string) string {
for _, value := range values {
if text := cleanXMLText(value); text != "" {
diff --git a/internal/service/local_metadata_apply.go b/internal/service/local_metadata_apply.go
new file mode 100644
index 0000000..cc142b9
--- /dev/null
+++ b/internal/service/local_metadata_apply.go
@@ -0,0 +1,110 @@
+package service
+
+import "github.com/ShukeBta/MediaStationGo/internal/model"
+
+func applyLocalMetadata(m *model.Media, local *LocalMetadata) {
+ applyLocalIdentityMetadata(m, local)
+ applyLocalArtworkMetadata(m, local)
+ applyLocalExternalIDMetadata(m, local)
+ applyLocalEpisodeMetadata(m, local)
+ applyLocalTaxonomyMetadata(m, local)
+ if local.NSFW {
+ m.NSFW = true
+ }
+ if localMetadataMarksMatched(local) {
+ m.ScrapeStatus = "matched"
+ }
+}
+
+func applyLocalIdentityMetadata(m *model.Media, local *LocalMetadata) {
+ if local.Title != "" {
+ m.Title = local.Title
+ }
+ if local.OriginalName != "" {
+ m.OriginalName = local.OriginalName
+ }
+ if local.AdultCode != "" {
+ m.OriginalName = local.AdultCode
+ }
+ if local.Year > 0 {
+ m.Year = local.Year
+ }
+ if local.Overview != "" {
+ m.Overview = local.Overview
+ }
+ if local.Rating > 0 {
+ m.Rating = local.Rating
+ }
+}
+
+func applyLocalArtworkMetadata(m *model.Media, local *LocalMetadata) {
+ if local.PosterURL != "" {
+ m.PosterURL = local.PosterURL
+ }
+ if local.BackdropURL != "" {
+ m.BackdropURL = local.BackdropURL
+ }
+}
+
+func applyLocalExternalIDMetadata(m *model.Media, local *LocalMetadata) {
+ if local.TMDbID > 0 {
+ m.TMDbID = local.TMDbID
+ }
+ if local.BangumiID > 0 {
+ m.BangumiID = local.BangumiID
+ }
+ if local.DoubanID != "" {
+ m.DoubanID = local.DoubanID
+ }
+ if local.TheTVDBID != "" {
+ m.TheTVDBID = local.TheTVDBID
+ }
+}
+
+func applyLocalEpisodeMetadata(m *model.Media, local *LocalMetadata) {
+ if local.EpisodeTitle != "" {
+ m.EpisodeTitle = local.EpisodeTitle
+ }
+ if local.SeasonNum > 0 || local.EpisodeNum > 0 {
+ m.SeasonNum = local.SeasonNum
+ }
+ if local.EpisodeNum > 0 {
+ m.EpisodeNum = local.EpisodeNum
+ }
+}
+
+func applyLocalTaxonomyMetadata(m *model.Media, local *LocalMetadata) {
+ if local.Genres != "" {
+ m.Genres = local.Genres
+ }
+ if local.Countries != "" {
+ m.Countries = local.Countries
+ }
+ if local.Languages != "" {
+ m.Languages = local.Languages
+ }
+}
+
+func localMetadataMarksMatched(local *LocalMetadata) bool {
+ return local != nil && (local.HasNFO || (!local.PathHint && localHasDescriptiveMetadata(local)))
+}
+
+func localHasDescriptiveMetadata(local *LocalMetadata) bool {
+ if local == nil {
+ return false
+ }
+ return local.Title != "" ||
+ local.OriginalName != "" ||
+ local.EpisodeTitle != "" ||
+ local.AdultCode != "" ||
+ local.Year > 0 ||
+ local.Overview != "" ||
+ local.Rating > 0 ||
+ local.TMDbID > 0 ||
+ local.BangumiID > 0 ||
+ local.DoubanID != "" ||
+ local.TheTVDBID != "" ||
+ local.Genres != "" ||
+ local.Countries != "" ||
+ local.Languages != ""
+}
diff --git a/internal/service/local_metadata_apply_test.go b/internal/service/local_metadata_apply_test.go
new file mode 100644
index 0000000..3d75bf6
--- /dev/null
+++ b/internal/service/local_metadata_apply_test.go
@@ -0,0 +1,38 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestApplyLocalMetadataPreservesPathHintAsPending(t *testing.T) {
+ media := &model.Media{Title: "原始标题", ScrapeStatus: "pending"}
+ applyLocalMetadata(media, &LocalMetadata{
+ Title: "路径标题",
+ Year: 2026,
+ TMDbID: 12345,
+ PathHint: true,
+ })
+
+ if media.Title != "路径标题" || media.Year != 2026 || media.TMDbID != 12345 {
+ t.Fatalf("path hint metadata was not applied: %+v", media)
+ }
+ if media.ScrapeStatus != "pending" {
+ t.Fatalf("path hints alone must stay enrichable, got scrape_status=%q", media.ScrapeStatus)
+ }
+}
+
+func TestApplyLocalMetadataMarksNFOAndDescriptiveMetadataMatched(t *testing.T) {
+ nfoMedia := &model.Media{ScrapeStatus: "pending"}
+ applyLocalMetadata(nfoMedia, &LocalMetadata{HasNFO: true})
+ if nfoMedia.ScrapeStatus != "matched" {
+ t.Fatalf("NFO metadata should mark matched, got %q", nfoMedia.ScrapeStatus)
+ }
+
+ descriptiveMedia := &model.Media{ScrapeStatus: "pending"}
+ applyLocalMetadata(descriptiveMedia, &LocalMetadata{Overview: "剧情简介"})
+ if descriptiveMedia.ScrapeStatus != "matched" {
+ t.Fatalf("descriptive metadata should mark matched, got %q", descriptiveMedia.ScrapeStatus)
+ }
+}
diff --git a/internal/service/local_metadata_artwork.go b/internal/service/local_metadata_artwork.go
new file mode 100644
index 0000000..eafb7bc
--- /dev/null
+++ b/internal/service/local_metadata_artwork.go
@@ -0,0 +1,249 @@
+package service
+
+import (
+ "image"
+ _ "image/gif"
+ _ "image/jpeg"
+ _ "image/png"
+ "net/url"
+ "os"
+ "path/filepath"
+ "strings"
+)
+
+func localPosterCandidates(mediaPath string) []string {
+ base := strings.TrimSuffix(filepath.Base(mediaPath), filepath.Ext(mediaPath))
+ names := []string{
+ base + "-poster",
+ base + ".poster",
+ "poster",
+ "folder",
+ "cover",
+ "movie",
+ "show",
+ base + "-cover",
+ base + ".cover",
+ base,
+ base + "-thumb",
+ base + ".thumb",
+ "thumb",
+ }
+ return append(adultArtworkNameCandidates(mediaPath, "poster"), names...)
+}
+
+func localBackdropCandidates(mediaPath string) []string {
+ base := strings.TrimSuffix(filepath.Base(mediaPath), filepath.Ext(mediaPath))
+ names := []string{
+ base + "-fanart",
+ base + ".fanart",
+ base + "-backdrop",
+ base + ".backdrop",
+ base + "-background",
+ "fanart",
+ "backdrop",
+ "background",
+ "landscape",
+ "banner",
+ "clearart",
+ }
+ return append(adultArtworkNameCandidates(mediaPath, "backdrop"), names...)
+}
+
+func adultArtworkNameCandidates(mediaPath, kind string) []string {
+ code := AdultCodeFromMediaPath(mediaPath)
+ if code == "" {
+ return nil
+ }
+ compact := strings.ReplaceAll(code, "-", "")
+ bases := []string{code, compact}
+ bases = append(bases, adultDMMNameCandidates(code)...)
+ out := make([]string, 0, len(bases)*6)
+ for _, base := range bases {
+ if base == "" {
+ continue
+ }
+ if kind == "poster" {
+ out = append(out, base, base+"-poster", base+".poster", base+"-cover", base+".cover", base+"-thumb", base+".thumb", base+"pl", base+"-pl")
+ } else {
+ out = append(out, base+"-fanart", base+".fanart", base+"-backdrop", base+".backdrop", base+"-background", base+"-landscape", base+"jp", base+"jp-1")
+ }
+ }
+ return out
+}
+
+func adultDMMNameCandidates(code string) []string {
+ parts := adultStandardPattern.FindStringSubmatch(code)
+ if len(parts) < 3 {
+ return nil
+ }
+ prefix := strings.ToLower(parts[1])
+ num := strings.TrimLeft(parts[2], "0")
+ if num == "" {
+ num = "0"
+ }
+ padded := num
+ for len(padded) < 5 {
+ padded = "0" + padded
+ }
+ return []string{prefix + padded}
+}
+
+func firstExistingImage(dir string, names ...string) string {
+ if dir == "" {
+ return ""
+ }
+ for _, name := range names {
+ for _, ext := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
+ path := filepath.Join(dir, name+ext)
+ if fileExists(path) {
+ return filepath.Clean(path)
+ }
+ }
+ }
+ return ""
+}
+
+func nfoPosterValues(doc *nfoDocument) []string {
+ if doc == nil {
+ return nil
+ }
+ values := []string{doc.Poster, doc.Art.Poster}
+ for _, thumb := range doc.Thumbs {
+ aspect := strings.ToLower(strings.TrimSpace(thumb.Aspect))
+ if aspect == "" || aspect == "poster" || aspect == "cover" || aspect == "default" {
+ values = append(values, thumb.Value)
+ }
+ }
+ values = append(values, doc.Art.Thumb)
+ return values
+}
+
+func nfoBackdropValues(doc *nfoDocument) []string {
+ if doc == nil {
+ return nil
+ }
+ values := []string{doc.Fanart.Value, doc.Art.Fanart, doc.Art.Backdrop, doc.Art.Background, doc.Art.Landscape, doc.Art.Banner}
+ for _, thumb := range doc.Thumbs {
+ aspect := strings.ToLower(strings.TrimSpace(thumb.Aspect))
+ if aspect == "fanart" || aspect == "backdrop" || aspect == "background" || aspect == "landscape" {
+ values = append(values, thumb.Value)
+ }
+ }
+ values = append(values, doc.Fanart.Thumbs...)
+ return values
+}
+
+func firstLocalPoster(mediaPath, showBaseDir string) string {
+ mediaDir := filepath.Dir(mediaPath)
+ dirs := []string{}
+ if showBaseDir != "" && !samePath(showBaseDir, mediaDir) {
+ dirs = append(dirs, showBaseDir)
+ }
+ dirs = append(dirs, mediaDir)
+ for _, dir := range dirs {
+ if localPoster := firstExistingPosterImage(dir, localPosterCandidates(mediaPath)...); localPoster != "" {
+ return localPoster
+ }
+ }
+ return ""
+}
+
+func firstExistingPosterImage(dir string, names ...string) string {
+ if dir == "" {
+ return ""
+ }
+ for _, name := range names {
+ if isRejectedPosterName(name) {
+ continue
+ }
+ for _, ext := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
+ path := filepath.Join(dir, name+ext)
+ if fileExists(path) && likelyPosterImage(path) {
+ return filepath.Clean(path)
+ }
+ }
+ }
+ return ""
+}
+
+func firstAdultLooseImage(dir, kind string) string {
+ if dir == "" {
+ return ""
+ }
+ matches, _ := filepath.Glob(filepath.Join(dir, "*"))
+ preferred := []string{}
+ fallback := []string{}
+ for _, path := range matches {
+ ext := strings.ToLower(filepath.Ext(path))
+ if ext != ".jpg" && ext != ".jpeg" && ext != ".png" && ext != ".webp" && ext != ".gif" && ext != ".bmp" && ext != ".tbn" {
+ continue
+ }
+ name := strings.ToLower(strings.TrimSuffix(filepath.Base(path), ext))
+ if kind == "poster" {
+ if isRejectedPosterName(name) {
+ continue
+ }
+ if strings.Contains(name, "poster") || strings.Contains(name, "cover") || strings.Contains(name, "folder") || strings.Contains(name, "movie") || strings.HasSuffix(name, "pl") {
+ preferred = append(preferred, path)
+ }
+ } else if strings.Contains(name, "fanart") || strings.Contains(name, "backdrop") || strings.Contains(name, "background") || strings.Contains(name, "landscape") || strings.Contains(name, "jp") {
+ preferred = append(preferred, path)
+ }
+ if kind != "poster" || likelyPosterImage(path) {
+ fallback = append(fallback, path)
+ }
+ }
+ if len(preferred) > 0 {
+ return filepath.Clean(preferred[0])
+ }
+ if kind == "poster" && len(fallback) == 1 {
+ return filepath.Clean(fallback[0])
+ }
+ return ""
+}
+
+func isRejectedPosterName(name string) bool {
+ name = strings.ToLower(name)
+ rejected := []string{
+ "actor", "actors", "actress", "cast", "avatar", "portrait", "person",
+ "sample", "screenshot", "screen", "still", "scene", "extrafanart", "extrathumb",
+ "fanart", "backdrop", "background", "landscape", "banner", "clearlogo", "clearart", "logo", "disc",
+ }
+ for _, token := range rejected {
+ if strings.Contains(name, token) {
+ return true
+ }
+ }
+ return false
+}
+
+func likelyPosterImage(path string) bool {
+ file, err := os.Open(path) // #nosec G304 -- path is a discovered artwork sidecar under the configured library root.
+ if err != nil {
+ return false
+ }
+ defer file.Close()
+ cfg, _, err := image.DecodeConfig(file)
+ if err != nil || cfg.Width <= 0 || cfg.Height <= 0 {
+ return true
+ }
+ return cfg.Height >= cfg.Width
+}
+
+func fileExists(path string) bool {
+ info, err := os.Stat(path)
+ return err == nil && !info.IsDir()
+}
+
+func isHTTPURL(raw string) bool {
+ u, err := url.Parse(raw)
+ if err != nil {
+ return false
+ }
+ return (u.Scheme == "http" || u.Scheme == "https") && u.Host != ""
+}
+
+func isLocalPath(raw string) bool {
+ raw = strings.TrimSpace(raw)
+ return raw != "" && !isHTTPURL(raw)
+}
diff --git a/internal/service/local_metadata_test.go b/internal/service/local_metadata_test.go
index a707b2a..11abc26 100644
--- a/internal/service/local_metadata_test.go
+++ b/internal/service/local_metadata_test.go
@@ -5,9 +5,7 @@ import (
"path/filepath"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -78,6 +76,9 @@ func TestReadLocalEpisodeMetadataMergesShowAndEpisode(t *testing.T) {
if got.OriginalName != "" {
t.Fatalf("episode title must not pollute OriginalName, got %q", got.OriginalName)
}
+ if got.EpisodeTitle != "第三集" {
+ t.Fatalf("episode title metadata = %q, want 第三集", got.EpisodeTitle)
+ }
// 单集级简介按集回填;整剧 tmdb 仍取 tvshow.nfo 的 123,单集 id 不得覆盖。
if got.Overview != "本集简介" || got.TMDbID != 123 {
t.Fatalf("episode/show merge failed: %+v", got)
@@ -443,13 +444,7 @@ func TestScanLibraryUsesLocalMetadata(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "TV", Path: root, Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -476,13 +471,16 @@ func TestScanLibraryUsesLocalMetadata(t *testing.T) {
if media.Title != "本地剧名" || media.OriginalName != "" || media.SeasonNum != 2 || media.EpisodeNum != 3 || media.ScrapeStatus != "matched" {
t.Fatalf("unexpected scanned media: %+v", media)
}
+ if media.EpisodeTitle != "本地第三集" {
+ t.Fatalf("episode_title = %q, want 本地第三集", media.EpisodeTitle)
+ }
res, err = scanner.ScanLibrary(t.Context(), lib.ID)
if err != nil {
t.Fatal(err)
}
- if res.Added != 0 || res.Updated != 1 {
- t.Fatalf("repeat scan counts added=%d updated=%d, want 0/1", res.Added, res.Updated)
+ if res.Added != 0 || res.Updated != 0 || res.Skipped != 1 {
+ t.Fatalf("repeat scan counts added=%d updated=%d skipped=%d, want 0/0/1", res.Added, res.Updated, res.Skipped)
}
}
@@ -497,13 +495,7 @@ func TestScanLibraryDoesNotMarkArtworkOnlyAsMatched(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "Adult", Path: root, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -541,13 +533,7 @@ func TestScanLibraryRefreshesArtworkOnlyMetadata(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "Movies", Path: root, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -590,13 +576,7 @@ func TestScanLibraryParsesEpisodesForMovieTypedLibrary(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "综艺", Path: root, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -623,13 +603,7 @@ func TestScanLibraryPrunesMissingMedia(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "TV", Path: root, Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
diff --git a/internal/service/manual_scrape.go b/internal/service/manual_scrape.go
index ba5710f..7f7351b 100644
--- a/internal/service/manual_scrape.go
+++ b/internal/service/manual_scrape.go
@@ -11,23 +11,32 @@ import (
)
type ManualScrapeRequest struct {
- Source string `json:"source"`
- MediaType string `json:"media_type"`
- Title string `json:"title"`
- OriginalName string `json:"original_name"`
- Overview string `json:"overview"`
- PosterURL string `json:"poster_url"`
- BackdropURL string `json:"backdrop_url"`
- Year int `json:"year"`
- Rating float32 `json:"rating"`
- TMDbID int `json:"tmdb_id"`
- BangumiID int `json:"bangumi_id"`
- DoubanID string `json:"douban_id"`
- TheTVDBID string `json:"thetvdb_id"`
- Languages []string `json:"languages"`
- Countries []string `json:"countries"`
- Genres []string `json:"genres"`
- NSFW bool `json:"nsfw"`
+ Source string `json:"source"`
+ MediaType string `json:"media_type"`
+ Title string `json:"title"`
+ OriginalName string `json:"original_name"`
+ Overview string `json:"overview"`
+ PosterURL string `json:"poster_url"`
+ BackdropURL string `json:"backdrop_url"`
+ Year int `json:"year"`
+ Rating float32 `json:"rating"`
+ TMDbID int `json:"tmdb_id"`
+ BangumiID int `json:"bangumi_id"`
+ DoubanID string `json:"douban_id"`
+ TheTVDBID string `json:"thetvdb_id"`
+ Languages []string `json:"languages"`
+ Countries []string `json:"countries"`
+ Genres []string `json:"genres"`
+ NSFW bool `json:"nsfw"`
+ EpisodeArtwork *bool `json:"episode_artwork,omitempty"`
+ EpisodeImages *bool `json:"episode_images,omitempty"`
+}
+
+func (r ManualScrapeRequest) EpisodeArtworkOption() *bool {
+ if r.EpisodeImages != nil {
+ return r.EpisodeImages
+ }
+ return r.EpisodeArtwork
}
func (s *ScraperService) ManualSearch(ctx context.Context, media *model.Media, query, provider, mediaType string) ([]ExternalMediaResult, error) {
@@ -35,27 +44,22 @@ func (s *ScraperService) ManualSearch(ctx context.Context, media *model.Media, q
return nil, errors.New("media required")
}
lib, _ := s.repo.Library.FindByID(ctx, media.LibraryID)
- query = strings.TrimSpace(query)
- if query == "" {
- query = firstText(media.Title, media.OriginalName)
- }
- if query == "" {
- query, _ = CleanQuery(media.Path)
- }
- if query == "" {
+ queries := manualSearchQueries(media, lib, query)
+ if len(queries) == 0 {
return nil, errors.New("search query required")
}
- if mediaType == "" && lib != nil {
- mediaType = lib.Type
- }
- mediaType = normalizeMediaType(mediaType, query, "")
- provider = strings.ToLower(strings.TrimSpace(provider))
- if provider == "" || provider == "all" {
- provider = "all"
+ if mediaType == "" {
+ if mediaIsEpisodic(media, lib) {
+ mediaType = "tv"
+ } else if lib != nil {
+ mediaType = lib.Type
+ }
}
+ mediaType = normalizeMediaType(mediaType, queries[0], "")
+ providers := manualSearchProviderSet(provider)
year := mediaYearHint(media)
if year <= 0 {
- _, year = CleanQuery(query)
+ _, year = CleanQuery(queries[0])
}
out := make([]ExternalMediaResult, 0, 6)
@@ -78,6 +82,7 @@ func (s *ScraperService) ManualSearch(ctx context.Context, media *model.Media, q
DoubanID: match.DoubanID,
TheTVDBID: match.TheTVDBID,
SubscribeKeyword: buildSubscribeKeyword(match.Title, match.Year),
+ SubscribeAliases: buildSubscribeAliases(match.Title, match.OriginalName, match.Year),
Languages: match.Languages,
Countries: match.Countries,
Genres: match.Genres,
@@ -85,42 +90,108 @@ func (s *ScraperService) ManualSearch(ctx context.Context, media *model.Media, q
})
}
- if provider == "all" || provider == "adult" {
- for _, match := range s.manualAdultMatches(ctx, media, query) {
- add("adult", "adult", match)
- }
- }
- if provider == "all" || provider == "tmdb" {
- for _, match := range s.manualTMDbMatches(ctx, query, year, mediaType) {
- typ := mediaType
- if typ == "" {
- typ = "movie"
+ if providers.want("adult") {
+ for _, candidateQuery := range queries {
+ for _, match := range s.manualAdultMatches(ctx, media, candidateQuery) {
+ add("adult", "adult", match)
}
- if match.TMDbID > 0 && isTVLikeTMDbMatch(match, mediaType) {
- typ = "tv"
+ }
+ }
+ if providers.want("tmdb") {
+ for _, candidateQuery := range queries {
+ for _, candidate := range s.manualTMDbCandidates(ctx, candidateQuery, year, mediaType) {
+ add("tmdb", candidate.MediaType, candidate.Match)
}
- add("tmdb", typ, match)
}
}
- if provider == "all" || provider == "douban" {
- if match := s.manualDoubanMatch(ctx, query); match != nil {
- add("douban", normalizeMediaType(mediaType, query, ""), match)
+ if providers.want("douban") {
+ for _, candidateQuery := range queries {
+ if match := s.manualDoubanMatch(ctx, candidateQuery); match != nil {
+ add("douban", normalizeMediaType(mediaType, candidateQuery, ""), match)
+ }
}
}
- if provider == "all" || provider == "bangumi" {
- if match := s.manualBangumiMatch(ctx, query); match != nil {
- add("bangumi", "anime", match)
+ if providers.want("bangumi") {
+ for _, candidateQuery := range queries {
+ if match := s.manualBangumiMatch(ctx, candidateQuery); match != nil {
+ add("bangumi", "anime", match)
+ }
}
}
- if provider == "all" || provider == "thetvdb" {
- if match := s.manualTheTVDBMatch(ctx, query); match != nil {
- add("thetvdb", "tv", match)
+ if providers.want("thetvdb") {
+ for _, candidateQuery := range queries {
+ if match := s.manualTheTVDBMatch(ctx, candidateQuery); match != nil {
+ add("thetvdb", "tv", match)
+ }
}
}
return dedupeExternalMedia(out), nil
}
+type manualSearchProviders map[string]struct{}
+
+func manualSearchProviderSet(provider string) manualSearchProviders {
+ out := manualSearchProviders{}
+ for _, field := range strings.FieldsFunc(provider, func(r rune) bool {
+ return r == ',' || r == ';' || r == '|' || r == ' '
+ }) {
+ field = strings.ToLower(strings.TrimSpace(field))
+ if field == "" || field == "all" {
+ return nil
+ }
+ out[field] = struct{}{}
+ }
+ if len(out) == 0 {
+ return nil
+ }
+ return out
+}
+
+func (p manualSearchProviders) want(provider string) bool {
+ if len(p) == 0 {
+ return true
+ }
+ _, ok := p[provider]
+ return ok
+}
+
+func manualSearchQueries(media *model.Media, lib *model.Library, query string) []string {
+ seen := map[string]struct{}{}
+ out := make([]string, 0, 4)
+ add := func(value string) {
+ value = strings.Join(strings.Fields(strings.TrimSpace(value)), " ")
+ if value == "" {
+ return
+ }
+ key := strings.ToLower(value)
+ if _, ok := seen[key]; ok {
+ return
+ }
+ seen[key] = struct{}{}
+ out = append(out, value)
+ }
+
+ add(query)
+ if strings.TrimSpace(query) == "" && media != nil {
+ add(firstText(media.Title, media.OriginalName))
+ }
+ if media != nil {
+ for _, candidate := range scrapeQueryCandidates(media, lib) {
+ add(candidate)
+ }
+ if len(out) == 0 {
+ title, _ := CleanQuery(media.Path)
+ add(title)
+ }
+ }
+ return out
+}
+
func (s *ScraperService) ApplyManualMatch(ctx context.Context, mediaID string, req ManualScrapeRequest) (*model.Media, error) {
+ return s.ApplyManualMatchWithOptions(ctx, mediaID, req, ScrapeOptions{EpisodeArtwork: req.EpisodeArtworkOption()})
+}
+
+func (s *ScraperService) ApplyManualMatchWithOptions(ctx context.Context, mediaID string, req ManualScrapeRequest, options ScrapeOptions) (*model.Media, error) {
media, err := s.repo.Media.FindByID(ctx, mediaID)
if err != nil || media == nil {
return nil, errors.New("media not found")
@@ -133,7 +204,7 @@ func (s *ScraperService) ApplyManualMatch(ctx context.Context, mediaID string, r
if strings.TrimSpace(match.Title) == "" {
return nil, errors.New("manual match title required")
}
- if err := s.applyProviderMatch(ctx, media, lib, match); err != nil {
+ if err := s.applyProviderMatchWithOptions(ctx, media, lib, match, options); err != nil {
return nil, err
}
return s.repo.Media.FindByID(ctx, mediaID)
@@ -183,29 +254,78 @@ func (s *ScraperService) manualRequestMatch(ctx context.Context, req ManualScrap
return fallback()
}
-func (s *ScraperService) manualTMDbMatches(ctx context.Context, query string, year int, mediaType string) []*Match {
+type manualTMDbCandidate struct {
+ MediaType string
+ Match *Match
+}
+
+func (s *ScraperService) manualTMDbCandidates(ctx context.Context, query string, year int, mediaType string) []manualTMDbCandidate {
if s.tmdb == nil || !s.tmdb.Enabled() {
return nil
}
if id, ok := parsePositiveInt(query); ok {
- if match := s.manualTMDbMatchByID(ctx, id, mediaType); match != nil {
- return []*Match{match}
+ out := make([]manualTMDbCandidate, 0, 2)
+ for _, typ := range manualTMDbIDSearchTypes(mediaType) {
+ if match := s.manualTMDbMatchByIDForType(ctx, id, typ); match != nil {
+ out = append(out, manualTMDbCandidate{MediaType: typ, Match: match})
+ }
}
+ return out
}
- out := make([]*Match, 0, 2)
- if mediaType == "" || mediaType == "movie" {
- if matches, err := s.tmdb.SearchMovieCandidates(ctx, query, year); err == nil {
- out = append(out, matches...)
- }
- }
- if mediaType == "" || mediaType == "tv" || mediaType == "anime" || mediaType == "variety" {
- if matches, err := s.tmdb.SearchTVCandidates(ctx, query, year); err == nil {
- out = append(out, matches...)
+ out := make([]manualTMDbCandidate, 0, 4)
+ for _, typ := range manualTMDbSearchTypes(mediaType) {
+ switch typ {
+ case "movie":
+ if matches, err := s.tmdb.SearchMovieCandidates(ctx, query, year); err == nil {
+ for _, match := range matches {
+ out = append(out, manualTMDbCandidate{MediaType: "movie", Match: match})
+ }
+ }
+ case "tv":
+ if matches, err := s.tmdb.SearchTVCandidates(ctx, query, year); err == nil {
+ for _, match := range matches {
+ out = append(out, manualTMDbCandidate{MediaType: "tv", Match: match})
+ }
+ }
}
}
return out
}
+func (s *ScraperService) manualTMDbMatches(ctx context.Context, query string, year int, mediaType string) []*Match {
+ candidates := s.manualTMDbCandidates(ctx, query, year, mediaType)
+ out := make([]*Match, 0, len(candidates))
+ for _, candidate := range candidates {
+ out = append(out, candidate.Match)
+ }
+ return out
+}
+
+func manualTMDbIDSearchTypes(mediaType string) []string {
+ switch normalizeMediaType(mediaType, "", "") {
+ case "tv", "anime", "variety":
+ return []string{"tv", "movie"}
+ case "movie", "adult":
+ return []string{"movie", "tv"}
+ default:
+ return []string{"movie", "tv"}
+ }
+}
+
+func manualTMDbSearchTypes(mediaType string) []string {
+ if strings.TrimSpace(mediaType) == "" {
+ return []string{"movie", "tv"}
+ }
+ switch normalizeMediaType(mediaType, "", "") {
+ case "tv", "anime", "variety":
+ return []string{"tv", "movie"}
+ case "movie", "adult":
+ return []string{"movie"}
+ default:
+ return []string{"movie", "tv"}
+ }
+}
+
func (s *ScraperService) manualAdultMatches(ctx context.Context, media *model.Media, query string) []*Match {
candidates := []string{query}
if media != nil {
@@ -258,6 +378,23 @@ func (s *ScraperService) manualTMDbMatchByID(ctx context.Context, id int, mediaT
return nil
}
+func (s *ScraperService) manualTMDbMatchByIDForType(ctx context.Context, id int, mediaType string) *Match {
+ if s.tmdb == nil || !s.tmdb.Enabled() || id <= 0 {
+ return nil
+ }
+ switch normalizeMediaType(mediaType, "", "") {
+ case "tv", "anime", "variety":
+ if match, err := s.tmdb.GetTVMatch(ctx, id); err == nil && match != nil {
+ return match
+ }
+ case "movie", "adult":
+ if match, err := s.tmdb.GetMovieMatch(ctx, id); err == nil && match != nil {
+ return match
+ }
+ }
+ return nil
+}
+
func (s *ScraperService) manualDoubanMatch(ctx context.Context, query string) *Match {
if s.douban == nil || !s.douban.Enabled() {
return nil
diff --git a/internal/service/manual_scrape_test.go b/internal/service/manual_scrape_test.go
new file mode 100644
index 0000000..5998b97
--- /dev/null
+++ b/internal/service/manual_scrape_test.go
@@ -0,0 +1,401 @@
+package service
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestManualRequestMatchFallsBackToCandidatePayload(t *testing.T) {
+ scraper := &ScraperService{}
+ match, err := scraper.manualRequestMatch(t.Context(), ManualScrapeRequest{
+ Source: "douban",
+ Title: "手动选择的电影",
+ DoubanID: "1234567",
+ Year: 2026,
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if match.Title != "手动选择的电影" || match.DoubanID != "1234567" || match.Year != 2026 {
+ t.Fatalf("fallback match = %#v", match)
+ }
+}
+
+func TestManualSearchReturnsTMDbCandidatePage(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ if r.URL.Path != "/search/movie" {
+ http.NotFound(w, r)
+ return
+ }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{
+ {
+ "id": 101,
+ "title": "错误的同名电影",
+ "poster_path": "/wrong.jpg",
+ "release_date": "2021-01-01",
+ "vote_average": 5.1,
+ "genre_ids": []int{18},
+ "backdrop_path": "/wrong-backdrop.jpg",
+ },
+ {
+ "id": 202,
+ "title": "正确的同名电影",
+ "poster_path": "/right.jpg",
+ "release_date": "2021-08-01",
+ "vote_average": 8.2,
+ "genre_ids": []int{28},
+ "backdrop_path": "/right-backdrop.jpg",
+ },
+ },
+ })
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{LibraryID: lib.ID, Title: "同名电影", Path: "/media/movie/同名电影.mkv"}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ results, err := scraper.ManualSearch(t.Context(), &media, "同名电影", "tmdb", "movie")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(results) != 2 || results[0].TMDbID != 101 || results[1].TMDbID != 202 {
+ t.Fatalf("manual TMDb candidates = %#v", results)
+ }
+}
+
+func TestManualSearchFallsBackToMovieFolderForGenericQuery(t *testing.T) {
+ var queries []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ queries = append(queries, r.URL.Query().Get("query"))
+ w.Header().Set("Content-Type", "application/json")
+ if r.URL.Path != "/search/movie" {
+ http.NotFound(w, r)
+ return
+ }
+ if r.URL.Query().Get("query") != "inception" {
+ _ = json.NewEncoder(w).Encode(map[string]any{"results": []any{}})
+ return
+ }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 27205,
+ "title": "Inception",
+ "overview": "A thief enters dreams.",
+ "poster_path": "/inception.jpg",
+ "release_date": "2010-07-16",
+ "vote_average": 8.4,
+ "original_title": "Inception",
+ }},
+ })
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "Movies", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "00000",
+ Path: `/media/movies/Inception (2010)/BDMV/STREAM/00000.m2ts`,
+ }
+
+ results, err := scraper.ManualSearch(t.Context(), &media, "00000", "tmdb", "movie")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(results) != 1 || results[0].TMDbID != 27205 {
+ t.Fatalf("manual search results=%#v, want folder fallback candidate; queries=%v", results, queries)
+ }
+ if len(queries) < 2 || queries[0] != "00000" || queries[1] != "inception" {
+ t.Fatalf("manual search queries=%v, want explicit query then folder fallback", queries)
+ }
+}
+
+func TestManualSearchReturnsMovieFallbackForTVTypedTMDbSearch(t *testing.T) {
+ var paths []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ paths = append(paths, r.URL.Path)
+ w.Header().Set("Content-Type", "application/json")
+ switch r.URL.Path {
+ case "/search/tv":
+ _ = json.NewEncoder(w).Encode(map[string]any{"results": []any{}})
+ case "/search/movie":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 808,
+ "title": "正义女神",
+ "poster_path": "/movie.jpg",
+ "release_date": "2024-01-01",
+ "vote_average": 7.1,
+ "backdrop_path": "/movie-backdrop.jpg",
+ }},
+ })
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "电视剧", Path: `/media/tv`, Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{LibraryID: lib.ID, Title: "正义女神", Path: `/media/tv/正义女神.mkv`}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ results, err := scraper.ManualSearch(t.Context(), &media, "正义女神", "tmdb", "tv")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(results) != 1 || results[0].TMDbID != 808 || results[0].MediaType != "movie" {
+ t.Fatalf("manual TMDb fallback results=%#v, paths=%v", results, paths)
+ }
+ if len(paths) < 2 || paths[0] != "/search/tv" || paths[1] != "/search/movie" {
+ t.Fatalf("tmdb search paths=%v, want tv first then movie fallback", paths)
+ }
+}
+
+func TestManualSearchTMDbNumericIDTriesMovieAndTVNamespaces(t *testing.T) {
+ var paths []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ paths = append(paths, r.URL.Path)
+ w.Header().Set("Content-Type", "application/json")
+ switch r.URL.Path {
+ case "/movie/12345":
+ http.NotFound(w, r)
+ case "/tv/12345":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": 12345,
+ "name": "数字 ID 剧集",
+ "original_name": "Numeric ID Show",
+ "overview": "Matched from TV namespace.",
+ "poster_path": "/tv.jpg",
+ "first_air_date": "2025-01-01",
+ "vote_average": 7.8,
+ })
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "电影", Path: `/media/movie`, Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{LibraryID: lib.ID, Title: "待匹配", Path: `/media/movie/raw.mkv`}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ results, err := scraper.ManualSearch(t.Context(), &media, "12345", "tmdb", "movie")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(results) != 1 || results[0].TMDbID != 12345 || results[0].MediaType != "tv" {
+ t.Fatalf("manual TMDb numeric results=%#v, paths=%v", results, paths)
+ }
+ if len(paths) < 2 || paths[0] != "/movie/12345" || paths[1] != "/tv/12345" {
+ t.Fatalf("tmdb numeric paths=%v, want movie then tv", paths)
+ }
+}
+
+func TestManualSearchIncludesAdultProvider(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.URL.Path {
+ case "/search":
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ _, _ = w.Write([]byte(`SSIS-001 手动候选`))
+ case "/v/ssis001":
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ _, _ = w.Write([]byte(`
SSIS-001 手动成人标题
`))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}, &model.APIConfig{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ apiConfig := NewAPIConfigService(zap.NewNop(), repos, NewCryptoService("", zap.NewNop()))
+ baseURL := upstream.URL
+ if _, err := apiConfig.Update(t.Context(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
+ t.Fatal(err)
+ }
+ log := zap.NewNop()
+ scraper := NewScraperService(&config.Config{}, log, repos, nil, nil, nil, nil, NewHub(log), NewAdultProvider(log, apiConfig))
+
+ lib := model.Library{Name: "成人", Path: "/media/adult", Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{LibraryID: lib.ID, Title: "SSIS-001", OriginalName: "SSIS-001", Path: "/media/adult/SSIS-001.mkv"}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ results, err := scraper.ManualSearch(t.Context(), &media, "SSIS-001", "adult", "adult")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(results) != 1 || results[0].Source != "adult" || results[0].MediaType != "adult" || !results[0].NSFW || results[0].OriginalName != "SSIS-001" {
+ t.Fatalf("manual adult candidates = %#v", results)
+ }
+}
+
+func TestApplyManualMatchSavesSelectedCloudMatchWhenDetailsSlow(t *testing.T) {
+ oldTimeout := tmdbDetailsTimeout
+ tmdbDetailsTimeout = 20 * time.Millisecond
+ defer func() { tmdbDetailsTimeout = oldTimeout }()
+
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/movie/77" {
+ http.NotFound(w, r)
+ return
+ }
+ select {
+ case <-r.Context().Done():
+ return
+ case <-time.After(time.Second):
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": 77,
+ "title": "Slow Details",
+ })
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "bad cloud title",
+ Path: "cloud://openlist/Movies/Bad.Title.2026.mkv",
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ start := time.Now()
+ if _, err := scraper.ApplyManualMatch(t.Context(), media.ID, ManualScrapeRequest{
+ Source: "manual",
+ MediaType: "movie",
+ Title: "Correct Cloud Movie",
+ TMDbID: 77,
+ Year: 2026,
+ }); err != nil {
+ t.Fatal(err)
+ }
+ if elapsed := time.Since(start); elapsed > 500*time.Millisecond {
+ t.Fatalf("manual apply waited for optional details: %s", elapsed)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.Title != "Correct Cloud Movie" || got.ScrapeStatus != "matched" || got.TMDbID != 77 {
+ t.Fatalf("manual cloud match was not saved: title=%q status=%q tmdb=%d", got.Title, got.ScrapeStatus, got.TMDbID)
+ }
+}
diff --git a/internal/service/media.go b/internal/service/media.go
index 4b30291..b70d17b 100644
--- a/internal/service/media.go
+++ b/internal/service/media.go
@@ -7,8 +7,6 @@ import (
"encoding/hex"
"errors"
"fmt"
- "os"
- "path/filepath"
"sort"
"strings"
"time"
@@ -96,229 +94,6 @@ func (s *MediaService) CreateLibrary(ctx context.Context, name, path, kind strin
return lib, nil
}
-func inferLibraryKind(name, path, requested string) string {
- requested = normalizeOrganizeMediaType(requested)
- text := strings.ToLower(name + " " + filepath.ToSlash(path))
- switch {
- case containsAnyText(text, "成人", "番号", "jav", "9kg", "adult", "nsfw"):
- return "adult"
- case containsAnyText(text, "综艺", "真人秀", "variety"):
- return "variety"
- case containsAnyText(text, "国漫", "日漫", "日番", "动漫", "动画", "anime", "bangumi") && !containsAnyText(text, "动画电影"):
- return "anime"
- case containsAnyText(text, "电视剧", "国产剧", "欧美剧", "日韩剧", "日剧", "韩剧", "剧集", "tv", "series"):
- return "tv"
- case containsAnyText(text, "电影", "movie", "film"):
- return "movie"
- }
- if requested != "" {
- return requested
- }
- return "movie"
-}
-
-func resolveAccessibleLibraryPath(path string) (string, error) {
- input := strings.TrimSpace(path)
- if input == "" {
- return "", errors.New("path required")
- }
- for _, candidate := range mappedPathCandidates(input) {
- if isAccessibleDir(candidate) {
- return filepath.Clean(candidate), nil
- }
- }
- abs, err := filepath.Abs(input)
- if err != nil {
- return "", fmt.Errorf("invalid path: %w", err)
- }
- return "", fmt.Errorf("path is not an accessible directory: %s", abs)
-}
-
-func resolveAccessibleMappedPath(path string) (string, os.FileInfo, error) {
- input := strings.TrimSpace(path)
- if input == "" {
- return "", nil, errors.New("path required")
- }
- candidates := mappedPathCandidates(input)
- for _, candidate := range candidates {
- if info, err := os.Stat(candidate); err == nil {
- return filepath.Clean(candidate), info, nil
- }
- }
- abs, err := filepath.Abs(input)
- if err != nil {
- return "", nil, fmt.Errorf("invalid path: %w", err)
- }
- return "", nil, fmt.Errorf("path is not accessible: %s", abs)
-}
-
-func resolveMappedDestinationPath(path string) string {
- path = strings.TrimSpace(path)
- if path == "" {
- return ""
- }
- clean := filepath.Clean(path)
- if _, err := os.Stat(clean); err == nil {
- return clean
- }
- for _, candidate := range mappedPathCandidates(clean) {
- if candidate == clean {
- continue
- }
- return filepath.Clean(candidate)
- }
- return clean
-}
-
-func mappedPathCandidates(input string) []string {
- var candidates []string
- add := func(candidate string) {
- candidate = filepath.Clean(filepath.FromSlash(strings.TrimSpace(candidate)))
- if candidate == "" || candidate == "." {
- return
- }
- for _, existing := range candidates {
- if sameLibraryPath(existing, candidate) {
- return
- }
- }
- candidates = append(candidates, candidate)
- }
- clean := filepath.Clean(input)
- add(clean)
- for _, candidate := range dockerVolumePathCandidates(input) {
- add(candidate)
- }
- for _, candidate := range dockerVolumePathCandidates(clean) {
- add(candidate)
- }
- if slashClean := cleanPathForVolumeMapping(input); slashClean != "" {
- add(slashClean)
- }
- if abs, err := filepath.Abs(input); err == nil {
- add(abs)
- for _, candidate := range dockerVolumePathCandidates(abs) {
- add(candidate)
- }
- }
- return candidates
-}
-
-func isAccessibleDir(path string) bool {
- info, err := os.Stat(path)
- return err == nil && info.IsDir()
-}
-
-func dockerVolumePathCandidates(path string) []string {
- normalized := cleanPathForVolumeMapping(path)
- var candidates []string
- addCandidate := func(candidate string) {
- candidate = filepath.Clean(filepath.FromSlash(candidate))
- for _, existing := range candidates {
- if sameLibraryPath(existing, candidate) {
- return
- }
- }
- candidates = append(candidates, candidate)
- }
-
- for _, mapping := range []struct {
- env string
- container string
- }{
- {env: "MEDIASTATION_MEDIA_DIR", container: envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")},
- {env: "MEDIASTATION_DOWNLOAD_DIR", container: envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")},
- } {
- host := cleanPathForVolumeMapping(os.Getenv(mapping.env))
- if host == "." || host == "" || strings.HasPrefix(host, ".") {
- continue
- }
- if normalized == host {
- addCandidate(mapping.container)
- continue
- }
- if strings.HasPrefix(normalized, host+"/") {
- addCandidate(mapping.container + strings.TrimPrefix(normalized, host))
- }
- container := cleanPathForVolumeMapping(mapping.container)
- if container == "." || container == "" || strings.HasPrefix(container, ".") {
- continue
- }
- if normalized == container {
- addCandidate(host)
- continue
- }
- if strings.HasPrefix(normalized, container+"/") {
- addCandidate(host + strings.TrimPrefix(normalized, container))
- }
- }
-
- for _, marker := range []struct {
- part string
- container string
- }{
- {part: "/media", container: envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")},
- {part: "/downloads", container: envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")},
- } {
- part := strings.TrimRight(marker.part, "/")
- container := strings.TrimRight(filepath.ToSlash(marker.container), "/")
- markerPath := pathAfterWindowsDrivePrefix(normalized)
- if markerPath == part {
- addCandidate(container)
- continue
- }
- if strings.HasPrefix(markerPath, part+"/") {
- addCandidate(container + strings.TrimPrefix(markerPath, part))
- }
- }
-
- return candidates
-}
-
-func cleanPathForVolumeMapping(path string) string {
- path = strings.TrimSpace(path)
- if path == "" {
- return ""
- }
- path = strings.ReplaceAll(path, "\\", "/")
- path = trimEmbeddedWindowsDrive(path)
- return filepath.ToSlash(filepath.Clean(filepath.FromSlash(path)))
-}
-
-func pathAfterWindowsDrivePrefix(path string) string {
- if len(path) >= 3 && path[1] == ':' && path[2] == '/' && isASCIIAlpha(path[0]) {
- return path[2:]
- }
- return path
-}
-
-func trimEmbeddedWindowsDrive(path string) string {
- for i := 0; i+2 < len(path); i++ {
- if !isASCIIAlpha(path[i]) || path[i+1] != ':' || path[i+2] != '/' {
- continue
- }
- if i == 0 || path[i-1] == '/' {
- return path[i:]
- }
- }
- return path
-}
-
-func isASCIIAlpha(ch byte) bool {
- return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')
-}
-
-func sameLibraryPath(a, b string) bool {
- return filepath.Clean(a) == filepath.Clean(b)
-}
-
-func envOrDefault(key, fallback string) string {
- if value := strings.TrimSpace(os.Getenv(key)); value != "" {
- return value
- }
- return fallback
-}
-
// ListLibraries returns every library configured on the server.
func (s *MediaService) ListLibraries(ctx context.Context) ([]model.Library, error) {
return s.repo.Library.List(ctx)
@@ -440,138 +215,6 @@ func (s *MediaService) invalidateMediaCache(ctx context.Context) {
}
}
-func (s *MediaService) attachLibraryMetadata(ctx context.Context, items []model.Media) {
- if s == nil || s.repo == nil || s.repo.Library == nil || len(items) == 0 {
- return
- }
- libs, err := s.repo.Library.List(ctx)
- if err != nil {
- return
- }
- byID := make(map[string]model.Library, len(libs))
- for _, lib := range libs {
- byID[lib.ID] = lib
- }
- resolver := newMediaDisplayLibraryResolver(ctx, s.repo, libs)
- for i := range items {
- if lib, ok := byID[items[i].LibraryID]; ok {
- items[i].LibraryName = lib.Name
- items[i].LibraryPath = lib.Path
- }
- if lib, ok := resolver.DisplayLibraryForMedia(items[i]); ok {
- items[i].DisplayLibraryID = lib.ID
- items[i].DisplayLibraryName = lib.Name
- items[i].DisplayLibraryPath = lib.Path
- }
- }
-}
-
-type mediaDisplayLibraryResolver struct {
- byID map[string]model.Library
- displayByID map[string]model.Library
- displayByMergeKey map[string]model.Library
- displayLibraries []model.Library
-}
-
-func newMediaDisplayLibraryResolver(ctx context.Context, repo *repository.Container, libs []model.Library) mediaDisplayLibraryResolver {
- displayLibraries := FilterDisplayCloudLibraries(ctx, repo, append([]model.Library(nil), libs...))
- resolver := mediaDisplayLibraryResolver{
- byID: make(map[string]model.Library, len(libs)),
- displayByID: make(map[string]model.Library, len(displayLibraries)),
- displayByMergeKey: make(map[string]model.Library, len(displayLibraries)),
- displayLibraries: displayLibraries,
- }
- for _, lib := range libs {
- resolver.byID[lib.ID] = lib
- }
- for _, lib := range displayLibraries {
- resolver.displayByID[lib.ID] = lib
- if key, ok := CloudLibraryMergeKey(lib); ok {
- if _, exists := resolver.displayByMergeKey[key]; !exists {
- resolver.displayByMergeKey[key] = lib
- }
- }
- }
- return resolver
-}
-
-func (r mediaDisplayLibraryResolver) DisplayLibraryForMedia(media model.Media) (model.Library, bool) {
- if lib, ok := r.bestPathDisplayLibrary(media); ok {
- return lib, true
- }
- if lib, ok := r.displayByID[media.LibraryID]; ok {
- return lib, true
- }
- own, hasOwn := r.byID[media.LibraryID]
- if hasOwn {
- if key, ok := CloudLibraryMergeKey(own); ok {
- if lib, exists := r.displayByMergeKey[key]; exists {
- return lib, true
- }
- }
- return own, true
- }
- return model.Library{}, false
-}
-
-func (r mediaDisplayLibraryResolver) bestPathDisplayLibrary(media model.Media) (model.Library, bool) {
- if strings.HasPrefix(strings.ToLower(strings.TrimSpace(media.Path)), "cloud://") {
- mediaInfo, ok := ParseCloudLibraryMount(media.Path)
- if !ok {
- return model.Library{}, false
- }
- var best model.Library
- bestDepth := 0
- for _, lib := range r.displayLibraries {
- info, ok := ParseCloudLibraryMount(lib.Path)
- if !ok || info.Provider != mediaInfo.Provider || !lib.Enabled {
- continue
- }
- dir := strings.Trim(firstNonEmpty(info.DisplayDir, info.ScanDir), "/")
- if dir == "" {
- continue
- }
- mediaDir := strings.Trim(firstNonEmpty(mediaInfo.DisplayDir, mediaInfo.ScanDir), "/")
- if mediaDir != dir && !cloudMountAncestor(dir, mediaDir) {
- continue
- }
- depth := len(strings.Split(dir, "/"))
- if depth > bestDepth {
- best = lib
- bestDepth = depth
- }
- }
- if bestDepth > 0 {
- return best, true
- }
- return model.Library{}, false
- }
-
- mediaPath := cleanPathForVolumeMapping(media.Path)
- var best model.Library
- bestLen := 0
- for _, lib := range r.displayLibraries {
- if _, ok := ParseCloudLibraryMount(lib.Path); ok || !lib.Enabled {
- continue
- }
- libPath := cleanPathForVolumeMapping(lib.Path)
- if libPath == "" || libPath == "." {
- continue
- }
- if mediaPath != libPath && !strings.HasPrefix(mediaPath, strings.TrimRight(libPath, "/")+"/") {
- continue
- }
- if len(libPath) > bestLen {
- best = lib
- bestLen = len(libPath)
- }
- }
- if bestLen > 0 {
- return best, true
- }
- return model.Library{}, false
-}
-
func groupMediaVersions(items []model.Media) []MediaItem {
if len(items) == 0 {
return nil
diff --git a/internal/service/media_classifier.go b/internal/service/media_classifier.go
index a710cd9..dab6904 100644
--- a/internal/service/media_classifier.go
+++ b/internal/service/media_classifier.go
@@ -341,7 +341,7 @@ func (s *SubscriptionService) lookupSubscriptionMetadata(ctx context.Context, me
if candidate == "" {
continue
}
- match := s.scraper.lookup(ctx, lib, candidate, year)
+ match := s.scraper.lookup(ctx, lib, nil, candidate, year)
if match == nil || strings.TrimSpace(match.Title) == "" {
continue
}
diff --git a/internal/service/media_classifier_test.go b/internal/service/media_classifier_test.go
index 41d7114..b5164df 100644
--- a/internal/service/media_classifier_test.go
+++ b/internal/service/media_classifier_test.go
@@ -3,9 +3,7 @@ package service
import (
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -218,13 +216,7 @@ func TestNormalizeMediaTypeAcceptsChineseLibraryTypes(t *testing.T) {
}
func TestSubscriptionResolveClassifiedSavePath(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
t.Fatal(err)
diff --git a/internal/service/media_display_library.go b/internal/service/media_display_library.go
new file mode 100644
index 0000000..7523f12
--- /dev/null
+++ b/internal/service/media_display_library.go
@@ -0,0 +1,141 @@
+package service
+
+import (
+ "context"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func (s *MediaService) attachLibraryMetadata(ctx context.Context, items []model.Media) {
+ if s == nil || s.repo == nil || s.repo.Library == nil || len(items) == 0 {
+ return
+ }
+ libs, err := s.repo.Library.List(ctx)
+ if err != nil {
+ return
+ }
+ byID := make(map[string]model.Library, len(libs))
+ for _, lib := range libs {
+ byID[lib.ID] = lib
+ }
+ resolver := newMediaDisplayLibraryResolver(ctx, s.repo, libs)
+ for i := range items {
+ if lib, ok := byID[items[i].LibraryID]; ok {
+ items[i].LibraryName = lib.Name
+ items[i].LibraryPath = lib.Path
+ }
+ if lib, ok := resolver.DisplayLibraryForMedia(items[i]); ok {
+ items[i].DisplayLibraryID = lib.ID
+ items[i].DisplayLibraryName = lib.Name
+ items[i].DisplayLibraryPath = lib.Path
+ }
+ }
+}
+
+type mediaDisplayLibraryResolver struct {
+ byID map[string]model.Library
+ displayByID map[string]model.Library
+ displayByMergeKey map[string]model.Library
+ displayLibraries []model.Library
+}
+
+func newMediaDisplayLibraryResolver(ctx context.Context, repo *repository.Container, libs []model.Library) mediaDisplayLibraryResolver {
+ displayLibraries := FilterDisplayCloudLibraries(ctx, repo, append([]model.Library(nil), libs...))
+ resolver := mediaDisplayLibraryResolver{
+ byID: make(map[string]model.Library, len(libs)),
+ displayByID: make(map[string]model.Library, len(displayLibraries)),
+ displayByMergeKey: make(map[string]model.Library, len(displayLibraries)),
+ displayLibraries: displayLibraries,
+ }
+ for _, lib := range libs {
+ resolver.byID[lib.ID] = lib
+ }
+ for _, lib := range displayLibraries {
+ resolver.displayByID[lib.ID] = lib
+ if key, ok := CloudLibraryMergeKey(lib); ok {
+ if _, exists := resolver.displayByMergeKey[key]; !exists {
+ resolver.displayByMergeKey[key] = lib
+ }
+ }
+ }
+ return resolver
+}
+
+func (r mediaDisplayLibraryResolver) DisplayLibraryForMedia(media model.Media) (model.Library, bool) {
+ if lib, ok := r.bestPathDisplayLibrary(media); ok {
+ return lib, true
+ }
+ if lib, ok := r.displayByID[media.LibraryID]; ok {
+ return lib, true
+ }
+ own, hasOwn := r.byID[media.LibraryID]
+ if hasOwn {
+ if key, ok := CloudLibraryMergeKey(own); ok {
+ if lib, exists := r.displayByMergeKey[key]; exists {
+ return lib, true
+ }
+ }
+ return own, true
+ }
+ return model.Library{}, false
+}
+
+func (r mediaDisplayLibraryResolver) bestPathDisplayLibrary(media model.Media) (model.Library, bool) {
+ if strings.HasPrefix(strings.ToLower(strings.TrimSpace(media.Path)), "cloud://") {
+ mediaInfo, ok := ParseCloudLibraryMount(media.Path)
+ if !ok {
+ return model.Library{}, false
+ }
+ var best model.Library
+ bestDepth := 0
+ for _, lib := range r.displayLibraries {
+ info, ok := ParseCloudLibraryMount(lib.Path)
+ if !ok || info.Provider != mediaInfo.Provider || !lib.Enabled {
+ continue
+ }
+ dir := strings.Trim(firstNonEmpty(info.DisplayDir, info.ScanDir), "/")
+ if dir == "" {
+ continue
+ }
+ mediaDir := strings.Trim(firstNonEmpty(mediaInfo.DisplayDir, mediaInfo.ScanDir), "/")
+ if mediaDir != dir && !cloudMountAncestor(dir, mediaDir) {
+ continue
+ }
+ depth := len(strings.Split(dir, "/"))
+ if depth > bestDepth {
+ best = lib
+ bestDepth = depth
+ }
+ }
+ if bestDepth > 0 {
+ return best, true
+ }
+ return model.Library{}, false
+ }
+
+ mediaPath := cleanPathForVolumeMapping(media.Path)
+ var best model.Library
+ bestLen := 0
+ for _, lib := range r.displayLibraries {
+ if _, ok := ParseCloudLibraryMount(lib.Path); ok || !lib.Enabled {
+ continue
+ }
+ libPath := cleanPathForVolumeMapping(lib.Path)
+ if libPath == "" || libPath == "." {
+ continue
+ }
+ if mediaPath != libPath && !strings.HasPrefix(mediaPath, strings.TrimRight(libPath, "/")+"/") {
+ continue
+ }
+ if len(libPath) > bestLen {
+ best = lib
+ bestLen = len(libPath)
+ }
+ }
+ if bestLen > 0 {
+ return best, true
+ }
+ return model.Library{}, false
+}
diff --git a/internal/service/media_paths.go b/internal/service/media_paths.go
new file mode 100644
index 0000000..950b943
--- /dev/null
+++ b/internal/service/media_paths.go
@@ -0,0 +1,232 @@
+package service
+
+import (
+ "errors"
+ "fmt"
+ "os"
+ "path/filepath"
+ "strings"
+)
+
+func inferLibraryKind(name, path, requested string) string {
+ requested = normalizeOrganizeMediaType(requested)
+ text := strings.ToLower(name + " " + filepath.ToSlash(path))
+ switch {
+ case containsAnyText(text, "成人", "番号", "jav", "9kg", "adult", "nsfw"):
+ return "adult"
+ case containsAnyText(text, "综艺", "真人秀", "variety"):
+ return "variety"
+ case containsAnyText(text, "国漫", "日漫", "日番", "动漫", "动画", "anime", "bangumi") && !containsAnyText(text, "动画电影"):
+ return "anime"
+ case containsAnyText(text, "电视剧", "国产剧", "欧美剧", "日韩剧", "日剧", "韩剧", "剧集", "tv", "series"):
+ return "tv"
+ case containsAnyText(text, "电影", "movie", "film"):
+ return "movie"
+ }
+ if requested != "" {
+ return requested
+ }
+ return "movie"
+}
+
+func resolveAccessibleLibraryPath(path string) (string, error) {
+ input := strings.TrimSpace(path)
+ if input == "" {
+ return "", errors.New("path required")
+ }
+ for _, candidate := range mappedPathCandidates(input) {
+ if isAccessibleDir(candidate) {
+ return filepath.Clean(candidate), nil
+ }
+ }
+ abs, err := filepath.Abs(input)
+ if err != nil {
+ return "", fmt.Errorf("invalid path: %w", err)
+ }
+ return "", fmt.Errorf("path is not an accessible directory: %s", abs)
+}
+
+func resolveAccessibleMappedPath(path string) (string, os.FileInfo, error) {
+ input := strings.TrimSpace(path)
+ if input == "" {
+ return "", nil, errors.New("path required")
+ }
+ candidates := mappedPathCandidates(input)
+ for _, candidate := range candidates {
+ if info, err := os.Stat(candidate); err == nil {
+ return filepath.Clean(candidate), info, nil
+ }
+ }
+ abs, err := filepath.Abs(input)
+ if err != nil {
+ return "", nil, fmt.Errorf("invalid path: %w", err)
+ }
+ return "", nil, fmt.Errorf("path is not accessible: %s", abs)
+}
+
+func resolveMappedDestinationPath(path string) string {
+ path = strings.TrimSpace(path)
+ if path == "" {
+ return ""
+ }
+ clean := filepath.Clean(path)
+ if _, err := os.Stat(clean); err == nil {
+ return clean
+ }
+ for _, candidate := range mappedPathCandidates(clean) {
+ if candidate == clean {
+ continue
+ }
+ return filepath.Clean(candidate)
+ }
+ return clean
+}
+
+func mappedPathCandidates(input string) []string {
+ var candidates []string
+ add := func(candidate string) {
+ candidate = filepath.Clean(filepath.FromSlash(strings.TrimSpace(candidate)))
+ if candidate == "" || candidate == "." {
+ return
+ }
+ for _, existing := range candidates {
+ if sameLibraryPath(existing, candidate) {
+ return
+ }
+ }
+ candidates = append(candidates, candidate)
+ }
+ clean := filepath.Clean(input)
+ add(clean)
+ for _, candidate := range dockerVolumePathCandidates(input) {
+ add(candidate)
+ }
+ for _, candidate := range dockerVolumePathCandidates(clean) {
+ add(candidate)
+ }
+ if slashClean := cleanPathForVolumeMapping(input); slashClean != "" {
+ add(slashClean)
+ }
+ if abs, err := filepath.Abs(input); err == nil {
+ add(abs)
+ for _, candidate := range dockerVolumePathCandidates(abs) {
+ add(candidate)
+ }
+ }
+ return candidates
+}
+
+func isAccessibleDir(path string) bool {
+ info, err := os.Stat(path)
+ return err == nil && info.IsDir()
+}
+
+func dockerVolumePathCandidates(path string) []string {
+ normalized := cleanPathForVolumeMapping(path)
+ var candidates []string
+ addCandidate := func(candidate string) {
+ candidate = filepath.Clean(filepath.FromSlash(candidate))
+ for _, existing := range candidates {
+ if sameLibraryPath(existing, candidate) {
+ return
+ }
+ }
+ candidates = append(candidates, candidate)
+ }
+
+ for _, mapping := range []struct {
+ env string
+ container string
+ }{
+ {env: "MEDIASTATION_MEDIA_DIR", container: envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")},
+ {env: "MEDIASTATION_DOWNLOAD_DIR", container: envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")},
+ } {
+ host := cleanPathForVolumeMapping(os.Getenv(mapping.env))
+ if host == "." || host == "" || strings.HasPrefix(host, ".") {
+ continue
+ }
+ if normalized == host {
+ addCandidate(mapping.container)
+ continue
+ }
+ if strings.HasPrefix(normalized, host+"/") {
+ addCandidate(mapping.container + strings.TrimPrefix(normalized, host))
+ }
+ container := cleanPathForVolumeMapping(mapping.container)
+ if container == "." || container == "" || strings.HasPrefix(container, ".") {
+ continue
+ }
+ if normalized == container {
+ addCandidate(host)
+ continue
+ }
+ if strings.HasPrefix(normalized, container+"/") {
+ addCandidate(host + strings.TrimPrefix(normalized, container))
+ }
+ }
+
+ for _, marker := range []struct {
+ part string
+ container string
+ }{
+ {part: "/media", container: envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")},
+ {part: "/downloads", container: envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")},
+ } {
+ part := strings.TrimRight(marker.part, "/")
+ container := strings.TrimRight(filepath.ToSlash(marker.container), "/")
+ markerPath := pathAfterWindowsDrivePrefix(normalized)
+ if markerPath == part {
+ addCandidate(container)
+ continue
+ }
+ if strings.HasPrefix(markerPath, part+"/") {
+ addCandidate(container + strings.TrimPrefix(markerPath, part))
+ }
+ }
+
+ return candidates
+}
+
+func cleanPathForVolumeMapping(path string) string {
+ path = strings.TrimSpace(path)
+ if path == "" {
+ return ""
+ }
+ path = strings.ReplaceAll(path, "\\", "/")
+ path = trimEmbeddedWindowsDrive(path)
+ return filepath.ToSlash(filepath.Clean(filepath.FromSlash(path)))
+}
+
+func pathAfterWindowsDrivePrefix(path string) string {
+ if len(path) >= 3 && path[1] == ':' && path[2] == '/' && isASCIIAlpha(path[0]) {
+ return path[2:]
+ }
+ return path
+}
+
+func trimEmbeddedWindowsDrive(path string) string {
+ for i := 0; i+2 < len(path); i++ {
+ if !isASCIIAlpha(path[i]) || path[i+1] != ':' || path[i+2] != '/' {
+ continue
+ }
+ if i == 0 || path[i-1] == '/' {
+ return path[i:]
+ }
+ }
+ return path
+}
+
+func isASCIIAlpha(ch byte) bool {
+ return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')
+}
+
+func sameLibraryPath(a, b string) bool {
+ return filepath.Clean(a) == filepath.Clean(b)
+}
+
+func envOrDefault(key, fallback string) string {
+ if value := strings.TrimSpace(os.Getenv(key)); value != "" {
+ return value
+ }
+ return fallback
+}
diff --git a/internal/service/media_series.go b/internal/service/media_series.go
new file mode 100644
index 0000000..a097dce
--- /dev/null
+++ b/internal/service/media_series.go
@@ -0,0 +1,297 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "path/filepath"
+ "regexp"
+ "sort"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type SeriesCard struct {
+ Key string `json:"key"`
+ Rep model.Media `json:"rep"`
+ LinkMedia model.Media `json:"linkMedia"`
+ Count int `json:"count"`
+}
+
+func (s *MediaService) ListLibrarySeriesCards(ctx context.Context, libraryID string, visibility MediaVisibility) ([]SeriesCard, int64, error) {
+ rows, _, err := s.listAllMediaVisible(ctx, libraryID, visibility)
+ if err != nil {
+ return nil, 0, err
+ }
+ cards := groupMediaSeriesCards(rows)
+ return cards, int64(len(cards)), nil
+}
+
+func (s *MediaService) ListLibrarySeriesEpisodes(ctx context.Context, libraryID, key string, visibility MediaVisibility) ([]model.Media, error) {
+ rows, _, err := s.listAllMediaVisible(ctx, libraryID, visibility)
+ if err != nil {
+ return nil, err
+ }
+ out := make([]model.Media, 0)
+ for _, row := range rows {
+ if mediaSeriesKey(row) == key {
+ out = append(out, row)
+ }
+ }
+ sort.SliceStable(out, func(i, j int) bool {
+ if out[i].SeasonNum != out[j].SeasonNum {
+ return out[i].SeasonNum < out[j].SeasonNum
+ }
+ if out[i].EpisodeNum != out[j].EpisodeNum {
+ return out[i].EpisodeNum < out[j].EpisodeNum
+ }
+ return out[i].CreatedAt.Before(out[j].CreatedAt)
+ })
+ return out, nil
+}
+
+func (s *MediaService) listAllMediaVisible(ctx context.Context, libraryID string, visibility MediaVisibility) ([]model.Media, int64, error) {
+ const pageSize = 2000
+ var all []model.Media
+ var total int64
+ for page := 1; ; page++ {
+ rows, n, err := s.ListMediaVisible(ctx, libraryID, page, pageSize, visibility)
+ if err != nil {
+ return nil, 0, err
+ }
+ if page == 1 {
+ total = n
+ all = make([]model.Media, 0, minInt64(n, pageSize))
+ }
+ all = append(all, rows...)
+ if int64(len(all)) >= n || len(rows) < pageSize {
+ break
+ }
+ }
+ return all, total, nil
+}
+
+func groupMediaSeriesCards(items []model.Media) []SeriesCard {
+ if len(items) == 0 {
+ return nil
+ }
+ cards := make([]SeriesCard, 0)
+ byKey := make(map[string]int, len(items))
+ for _, item := range items {
+ key := mediaSeriesKey(item)
+ if key == "" {
+ continue
+ }
+ if idx, ok := byKey[key]; ok {
+ card := &cards[idx]
+ card.Count++
+ if betterSeriesLinkMedia(item, card.LinkMedia) {
+ card.LinkMedia = item
+ }
+ currentArtwork := seriesArtworkScore(item)
+ representativeArtwork := seriesArtworkScore(card.Rep)
+ if currentArtwork > representativeArtwork {
+ card.Rep = item
+ } else if currentArtwork == representativeArtwork {
+ cur := item.SeasonNum*10000 + item.EpisodeNum
+ rep := card.Rep.SeasonNum*10000 + card.Rep.EpisodeNum
+ if cur > 0 && (rep == 0 || cur < rep) {
+ card.Rep = item
+ }
+ }
+ continue
+ }
+ byKey[key] = len(cards)
+ cards = append(cards, SeriesCard{Key: key, Rep: item, LinkMedia: item, Count: 1})
+ }
+ return cards
+}
+
+var episodicPathRE = regexp.MustCompile(`(?i)[\\/](?:电视剧|剧集|国产剧|欧美剧|日韩剧|日剧|韩剧|综艺|纪录片|动漫|番剧|国漫|日番|儿童|tv|series|shows?|season[\s._-]*\d|s\d{1,2}(?:[\s._-]|[\\/])|specials?|sp|ova|oad|extra|extras|特别篇|特別篇|番外|特典)[\\/]`)
+
+func mediaSeriesKey(media model.Media) string {
+ return compactSeriesKey(mediaSeriesRawKey(media))
+}
+
+func mediaSeriesRawKey(media model.Media) string {
+ fromPath := seriesTitleFromMediaPath(media.Path)
+ if media.SeasonNum > 0 || media.EpisodeNum > 0 || episodicPathRE.MatchString(media.Path+" "+media.DisplayLibraryPath+" "+media.LibraryPath) {
+ if fromPath != "" {
+ return seriesFingerprint("library-path", mediaTargetLibraryID(media), fromPath)
+ }
+ if media.TMDbID > 0 {
+ return fmt.Sprintf("tmdb:%d", media.TMDbID)
+ }
+ if media.BangumiID > 0 {
+ return fmt.Sprintf("bgm:%d", media.BangumiID)
+ }
+ if strings.TrimSpace(media.DoubanID) != "" {
+ return "douban:" + strings.TrimSpace(media.DoubanID)
+ }
+ if strings.TrimSpace(media.TheTVDBID) != "" {
+ return "thetvdb:" + strings.TrimSpace(media.TheTVDBID)
+ }
+ if strings.TrimSpace(media.SeriesID) != "" {
+ return "series:" + strings.TrimSpace(media.SeriesID)
+ }
+ return seriesFingerprint("library-title", mediaTargetLibraryID(media), normalizeSeriesTitle(seriesDisplayTitle(media)))
+ }
+ if strings.TrimSpace(media.SeriesID) != "" {
+ return "series:" + strings.TrimSpace(media.SeriesID)
+ }
+ if media.TMDbID > 0 {
+ return fmt.Sprintf("tmdb:%d", media.TMDbID)
+ }
+ if media.BangumiID > 0 {
+ return fmt.Sprintf("bgm:%d", media.BangumiID)
+ }
+ if fromPath != "" {
+ return seriesFingerprint("library-path", media.LibraryID, fromPath)
+ }
+ return seriesFingerprint("library-title", media.LibraryID, normalizeSeriesTitle(media.Title))
+}
+
+func seriesFingerprint(parts ...string) string {
+ return strings.Join(parts, "\x1f")
+}
+
+func compactSeriesKey(raw string) string {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return ""
+ }
+ var hash uint32 = 2166136261
+ for _, b := range []byte(raw) {
+ hash ^= uint32(b)
+ hash *= 16777619
+ }
+ return fmt.Sprintf("series:%08x", hash)
+}
+
+var (
+ seriesYearRE = regexp.MustCompile(`\s*\((?:19|20)\d{2}\)\s*`)
+ seriesIDRE = regexp.MustCompile(`(?i)\s*\[(?:tmdb|tmdbid)[=-]\d+\]\s*`)
+ seriesBraceRE = regexp.MustCompile(`(?i)\s*\{(?:tmdb|tmdbid|douban|bangumi|bgm|thetvdb|tvdb)[\s:=#-]*[a-z0-9_-]+\}\s*`)
+ seriesSpacerRE = regexp.MustCompile(`[\s._-]+`)
+)
+
+func normalizeSeriesTitle(value string) string {
+ value = strings.ToLower(strings.TrimSpace(value))
+ value = seriesYearRE.ReplaceAllString(value, " ")
+ value = seriesIDRE.ReplaceAllString(value, " ")
+ value = seriesBraceRE.ReplaceAllString(value, " ")
+ value = seriesSpacerRE.ReplaceAllString(value, " ")
+ return strings.TrimSpace(value)
+}
+
+func seriesTitleFromMediaPath(path string) string {
+ if strings.TrimSpace(path) == "" {
+ return ""
+ }
+ parts := strings.FieldsFunc(path, func(r rune) bool { return r == '/' || r == '\\' })
+ if len(parts) < 2 {
+ return ""
+ }
+ dirIndex := len(parts) - 2
+ for dirIndex >= 0 && strictSeasonFolderMatched(filepath.Base(parts[dirIndex])) {
+ dirIndex--
+ }
+ if dirIndex < 0 {
+ return ""
+ }
+ return normalizeSeriesTitle(parts[dirIndex])
+}
+
+func seriesDisplayTitle(media model.Media) string {
+ if fromPath := seriesTitleFromMediaPath(media.Path); fromPath != "" {
+ return fromPath
+ }
+ if media.Title != "" {
+ return media.Title
+ }
+ if media.OriginalName != "" {
+ return media.OriginalName
+ }
+ return "未命名节目"
+}
+
+func mediaTargetLibraryID(media model.Media) string {
+ if strings.TrimSpace(media.DisplayLibraryID) != "" {
+ return media.DisplayLibraryID
+ }
+ return media.LibraryID
+}
+
+func betterSeriesLinkMedia(candidate, current model.Media) bool {
+ candidateScore := librarySpecificityScore(candidate)
+ currentScore := librarySpecificityScore(current)
+ if candidateScore != currentScore {
+ return candidateScore > currentScore
+ }
+ return seriesArtworkScore(candidate) > seriesArtworkScore(current)
+}
+
+func librarySpecificityScore(media model.Media) int {
+ rawPath := strings.TrimSpace(firstNonEmpty(media.DisplayLibraryPath, media.LibraryPath))
+ if rawPath == "" {
+ return 0
+ }
+ normalized := strings.TrimRight(strings.ReplaceAll(rawPath, "\\", "/"), "/")
+ lower := strings.ToLower(normalized)
+ if strings.HasPrefix(lower, "cloud://") {
+ rest := normalized[len("cloud://"):]
+ slash := strings.Index(rest, "/")
+ if slash < 0 || slash == len(rest)-1 {
+ return 0
+ }
+ return 100 + len(nonEmptySlashParts(rest[slash+1:]))
+ }
+ return 200 + len(nonEmptySlashParts(normalized))
+}
+
+func nonEmptySlashParts(value string) []string {
+ parts := strings.Split(value, "/")
+ out := parts[:0]
+ for _, part := range parts {
+ if strings.TrimSpace(part) != "" {
+ out = append(out, part)
+ }
+ }
+ return out
+}
+
+var (
+ posterArtworkRE = regexp.MustCompile(`(poster|folder|cover|movie|show|pl)(?:[._-]|\.[a-z0-9]+$|$)`)
+ badArtworkRE = regexp.MustCompile(`(actor|actress|cast|avatar|sample|screenshot|screen|still|scene|fanart|backdrop|background|landscape|banner|logo|disc)`)
+)
+
+func seriesArtworkScore(media model.Media) int {
+ poster := strings.ToLower(media.PosterURL)
+ backdrop := strings.ToLower(media.BackdropURL)
+ if poster == "" {
+ if backdrop != "" {
+ return 5
+ }
+ return 0
+ }
+ if posterArtworkRE.MatchString(poster) {
+ return 40
+ }
+ if badArtworkRE.MatchString(poster) {
+ return 10
+ }
+ if strings.Contains(poster, "thumb") {
+ return 20
+ }
+ return 30
+}
+
+func minInt64(a int64, b int) int {
+ if a <= 0 {
+ return 0
+ }
+ if a > int64(b) {
+ return b
+ }
+ return int(a)
+}
diff --git a/internal/service/media_series_test.go b/internal/service/media_series_test.go
new file mode 100644
index 0000000..a8d9aed
--- /dev/null
+++ b/internal/service/media_series_test.go
@@ -0,0 +1,29 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestMediaSeriesKeyCollapsesNestedSpecialFolders(t *testing.T) {
+ main := model.Media{
+ LibraryID: "lib-tv",
+ Path: `cloud://openlist/动漫/国漫/示例剧/Season 01/示例剧.S01E01.mkv`,
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ }
+ special := model.Media{
+ LibraryID: "lib-tv",
+ Path: `cloud://openlist/动漫/国漫/示例剧/Extras/Season 01/示例剧.SP01.mkv`,
+ }
+
+ if got, want := mediaSeriesKey(special), mediaSeriesKey(main); got != want {
+ t.Fatalf("special key=%q, want main key=%q", got, want)
+ }
+
+ cards := groupMediaSeriesCards([]model.Media{main, special})
+ if len(cards) != 1 || cards[0].Count != 2 {
+ t.Fatalf("cards=%#v, want one merged series card with two items", cards)
+ }
+}
diff --git a/internal/service/media_test.go b/internal/service/media_test.go
index 2ad0e3e..14482cb 100644
--- a/internal/service/media_test.go
+++ b/internal/service/media_test.go
@@ -7,7 +7,6 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
@@ -179,13 +178,7 @@ func TestResolveAccessibleMappedPathMapsWindowsDownloadVariants(t *testing.T) {
}
func TestDeleteCloudLibraryPurgesMountWithoutRecycleBin(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
lib := model.Library{Name: "OpenList · 剑来", Path: "cloud://openlist/Anime/JianLai", Type: "anime", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -266,13 +259,7 @@ func TestGroupMediaVersionsMergesEpisodeByExternalIDAcrossLibraries(t *testing.T
}
func TestUpdateMediaMetadataMarksManualMatch(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
lib := model.Library{Name: "自采集", Path: "/media/custom", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -309,13 +296,7 @@ func TestUpdateMediaMetadataMarksManualMatch(t *testing.T) {
}
func TestMediaUpsertBackfillsExternalIDsForPendingCloudRows(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
path := "cloud://openlist/国漫/折腰 (2025) {tmdb-296753}/Season 1/折腰.S01E01.mkv"
if err := repos.DB.Create(&model.Media{
@@ -350,13 +331,7 @@ func TestMediaUpsertBackfillsExternalIDsForPendingCloudRows(t *testing.T) {
}
func TestMediaUpsertCorrectsCloudExternalIDConflicts(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
path := "cloud://openlist/国产剧/折腰 (2025) {tmdb-296753}/Season 1/折腰.S01E01.mkv"
if err := repos.DB.Create(&model.Media{
@@ -392,13 +367,7 @@ func TestMediaUpsertCorrectsCloudExternalIDConflicts(t *testing.T) {
}
func TestRepairCloudPathMetadataBackfillsExistingPlaceholders(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
path := "cloud://openlist/动画电影/雄狮少年2 (2024) {tmdb-1154478}/雄狮少年2 (2024) - 2160p.WEB-DL.H.265.DDP 5.1-ADWeb.mp4"
if err := repos.DB.Create(&model.Media{
@@ -427,13 +396,7 @@ func TestRepairCloudPathMetadataBackfillsExistingPlaceholders(t *testing.T) {
}
func TestRepairCloudPathMetadataCorrectsConflictingMatchedID(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
path := "cloud://openlist/国产剧/折腰 (2025) {tmdb-296753}/Season 1/折腰.S01E01.mkv"
if err := repos.DB.Create(&model.Media{
@@ -465,13 +428,7 @@ func TestRepairCloudPathMetadataCorrectsConflictingMatchedID(t *testing.T) {
}
func TestSoftDeleteCloudMediaPurgesRecordWithoutRecycleBin(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
media := model.Media{
Base: model.Base{ID: "cloud-media"},
@@ -504,13 +461,7 @@ func TestSoftDeleteCloudMediaPurgesRecordWithoutRecycleBin(t *testing.T) {
}
func TestListRecycleBinPrunesOldRowsOverLimit(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
now := time.Now()
for i := 0; i < maxRecycleBinRecords+5; i++ {
@@ -553,13 +504,7 @@ func TestListRecycleBinPrunesOldRowsOverLimit(t *testing.T) {
}
func TestSoftDeleteInvalidatesMediaAndStatsCache(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{})
repos := repository.New(db)
media := model.Media{
Base: model.Base{ID: "local-media"},
diff --git a/internal/service/media_visibility_test.go b/internal/service/media_visibility_test.go
index cadc319..7200f5f 100644
--- a/internal/service/media_visibility_test.go
+++ b/internal/service/media_visibility_test.go
@@ -8,19 +8,11 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
)
func TestMediaVisibilityFiltersNSFWAndLibraries(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
@@ -90,14 +82,55 @@ func TestMediaVisibilityFiltersNSFWAndLibraries(t *testing.T) {
}
}
-func TestConfiguredAdultLibrariesDoNotHideSafeLibraryWithNSFWItems(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+func TestMediaVisibilityHidesDeprecatedNativeCloudLibraries(t *testing.T) {
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
+ repos := repository.New(db)
+ svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
+
+ legacy := model.Library{
+ Name: "旧云盘",
+ Path: BuildCloudLibraryPath(LegacyQuarkProvider, "archive", "archive"),
+ Type: "movie",
+ Enabled: true,
+ }
+ openList := model.Library{
+ Name: "OpenList",
+ Path: BuildCloudLibraryPath("openlist", "movies", "movies"),
+ Type: "movie",
+ Enabled: true,
+ }
+ if err := repos.Library.Create(t.Context(), &legacy); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &openList); err != nil {
+ t.Fatal(err)
+ }
+ if err := db.Create(&[]model.Media{
+ {LibraryID: legacy.ID, Title: "历史媒体", Path: "cloud://" + LegacyQuarkProvider + "/archive/old.mkv"},
+ {LibraryID: openList.ID, Title: "可见媒体", Path: "cloud://openlist/movies/new.mkv"},
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ items, err := svc.SearchMediaVisible(t.Context(), "媒体", 20, MediaVisibility{IncludeNSFW: true})
if err != nil {
t.Fatal(err)
}
- if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{}); err != nil {
+ if got := sortedMediaTitles(items); !slices.Equal(got, []string{"可见媒体"}) {
+ t.Fatalf("deprecated native cloud media should be hidden from search, got %#v", got)
+ }
+
+ listed, total, err := svc.ListMediaVisible(t.Context(), legacy.ID, 1, 20, MediaVisibility{IncludeNSFW: true})
+ if err != nil {
t.Fatal(err)
}
+ if total != 0 || len(listed) != 0 {
+ t.Fatalf("deprecated native cloud media should be hidden from direct list total=%d rows=%#v", total, sortedMediaTitles(listed))
+ }
+}
+
+func TestConfiguredAdultLibrariesDoNotHideSafeLibraryWithNSFWItems(t *testing.T) {
+ db := newServiceTestDB(t, &model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{})
repos := repository.New(db)
safe := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true}
@@ -142,13 +175,7 @@ func TestConfiguredAdultLibrariesDoNotHideSafeLibraryWithNSFWItems(t *testing.T)
}
func TestSearchMediaVisibleHonorsLargePosterWallLimit(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
lib := model.Library{Name: "海报墙", Path: "/media/all", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -177,13 +204,7 @@ func TestSearchMediaVisibleHonorsLargePosterWallLimit(t *testing.T) {
}
func TestSearchMediaVisibleCanReturnHugeLibraryResultsWhenRequested(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{})
repos := repository.New(db)
lib := model.Library{Name: "海量剧集", Path: "/media/huge", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
diff --git a/internal/service/nfo.go b/internal/service/nfo.go
index da91541..2375e56 100644
--- a/internal/service/nfo.go
+++ b/internal/service/nfo.go
@@ -129,9 +129,12 @@ func WriteMediaNFO(m *model.Media) (string, error) {
var doc any
if m.SeasonNum > 0 || m.EpisodeNum > 0 {
- title := m.OriginalName
+ title := strings.TrimSpace(m.EpisodeTitle)
+ if title == "" && m.EpisodeNum > 0 {
+ title = fmt.Sprintf("第 %d 集", m.EpisodeNum)
+ }
if title == "" {
- title = m.Title
+ title = strings.TrimSpace(m.Title)
}
doc = episodeNFO{
Title: title,
diff --git a/internal/service/nfo_test.go b/internal/service/nfo_test.go
index a58cc1c..6270b2c 100644
--- a/internal/service/nfo_test.go
+++ b/internal/service/nfo_test.go
@@ -3,6 +3,7 @@ package service
import (
"os"
"path/filepath"
+ "strings"
"testing"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -36,3 +37,34 @@ func TestWriteMediaNFOUsesMappedDestinationPath(t *testing.T) {
t.Fatal(err)
}
}
+
+func TestWriteMediaNFOUsesEpisodeTitleForEpisodeDetails(t *testing.T) {
+ root := t.TempDir()
+ mediaPath := filepath.Join(root, "剧集", "间谍过家家", "Season 02", "间谍过家家 - S02E01.mkv")
+ if err := os.MkdirAll(filepath.Dir(mediaPath), 0o755); err != nil {
+ t.Fatal(err)
+ }
+ if err := os.WriteFile(mediaPath, []byte("media"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+
+ got, err := WriteMediaNFO(&model.Media{
+ Title: "间谍过家家",
+ OriginalName: "SPY×FAMILY",
+ EpisodeTitle: "任务代号: 猫",
+ Path: mediaPath,
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ body, err := os.ReadFile(got)
+ if err != nil {
+ t.Fatal(err)
+ }
+ text := string(body)
+ if !strings.Contains(text, "任务代号: 猫") || !strings.Contains(text, "间谍过家家") {
+ t.Fatalf("episode nfo did not keep episode/show titles separate:\n%s", text)
+ }
+}
diff --git a/internal/service/notify_telegram.go b/internal/service/notify_telegram.go
index a4c07b3..764a842 100644
--- a/internal/service/notify_telegram.go
+++ b/internal/service/notify_telegram.go
@@ -4,10 +4,7 @@ package service
import (
"context"
"fmt"
- "net/url"
"regexp"
- "sort"
- "strconv"
"strings"
"time"
"unicode/utf8"
@@ -429,384 +426,3 @@ func telegramFieldIcon(key string) string {
return "•"
}
}
-
-func telegramFirstValue(data map[string]interface{}, keys ...string) string {
- for _, key := range keys {
- if value := telegramDataString(data, key); value != "" {
- return value
- }
- }
- return ""
-}
-
-func telegramMessageFieldValue(message string, keys ...string) string {
- if strings.TrimSpace(message) == "" {
- return ""
- }
- for _, line := range strings.Split(message, "\n") {
- key, value, ok := splitTelegramField(line)
- if !ok {
- continue
- }
- for _, want := range keys {
- if strings.EqualFold(strings.TrimSpace(key), strings.TrimSpace(want)) {
- return value
- }
- }
- }
- return ""
-}
-
-func telegramMediaCategory(data map[string]interface{}) string {
- if category := telegramFirstValue(data, "media_category", "category"); category != "" {
- return category
- }
- switch strings.ToLower(telegramFirstValue(data, "media_type")) {
- case "movie":
- return "电影"
- case "tv", "series", "show":
- return "剧集"
- case "anime":
- return "动漫"
- case "variety":
- return "综艺"
- case "documentary":
- return "纪录片"
- default:
- return telegramFirstValue(data, "media_type")
- }
-}
-
-func telegramLanguageName(raw string) string {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return ""
- }
- parts := strings.FieldsFunc(raw, func(r rune) bool {
- return r == ',' || r == ',' || r == '/' || r == '|' || r == '、'
- })
- if len(parts) == 0 {
- parts = []string{raw}
- }
- seen := map[string]struct{}{}
- out := []string{}
- for _, part := range parts {
- part = strings.TrimSpace(strings.Trim(part, "[]"))
- if part == "" {
- continue
- }
- lower := strings.ToLower(strings.ReplaceAll(part, "_", "-"))
- name := part
- switch {
- case strings.HasPrefix(lower, "zh") || lower == "cn" || lower == "cmn":
- name = "中文"
- case lower == "en" || strings.HasPrefix(lower, "en-"):
- name = "英语"
- case lower == "ja" || lower == "jp" || strings.HasPrefix(lower, "ja-"):
- name = "日语"
- case lower == "ko" || lower == "kr" || strings.HasPrefix(lower, "ko-"):
- name = "韩语"
- case lower == "fr" || strings.HasPrefix(lower, "fr-"):
- name = "法语"
- case lower == "de" || strings.HasPrefix(lower, "de-"):
- name = "德语"
- case lower == "es" || strings.HasPrefix(lower, "es-"):
- name = "西班牙语"
- case lower == "it" || strings.HasPrefix(lower, "it-"):
- name = "意大利语"
- case lower == "ru" || strings.HasPrefix(lower, "ru-"):
- name = "俄语"
- case lower == "th" || strings.HasPrefix(lower, "th-"):
- name = "泰语"
- }
- if _, ok := seen[name]; ok {
- continue
- }
- seen[name] = struct{}{}
- out = append(out, name)
- }
- return strings.Join(out, "、")
-}
-
-func telegramSeasonEpisodeValue(event NotifyEvent) string {
- if value := telegramFirstValue(event.Data, "season_episode", "episode_tag"); value != "" {
- return strings.ToUpper(value)
- }
- for _, raw := range []string{
- telegramFirstValue(event.Data, "resource_title", "torrent_title", "release_title"),
- telegramFirstValue(event.Data, "title", "name"),
- event.Message,
- } {
- if value := telegramExtractSeasonEpisode(raw); value != "" {
- return value
- }
- }
- season := telegramFirstValue(event.Data, "season")
- episode := telegramFirstValue(event.Data, "episode")
- if season != "" && episode != "" {
- return fmt.Sprintf("S%02dE%02d", telegramEpisodeNumber(season), telegramEpisodeNumber(episode))
- }
- return ""
-}
-
-func telegramEpisodeNumber(raw string) int {
- raw = strings.TrimSpace(strings.TrimLeft(strings.ToUpper(raw), "SE"))
- raw = strings.TrimLeft(raw, "0")
- if raw == "" {
- return 0
- }
- n, _ := strconv.Atoi(raw)
- return n
-}
-
-func telegramExtractSeasonEpisode(raw string) string {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return ""
- }
- return strings.ToUpper(telegramSeasonEpisodePattern.FindString(raw))
-}
-
-func telegramSizeValue(data map[string]interface{}) string {
- size := telegramFirstValue(data, "size")
- bitrate := telegramFirstValue(data, "bitrate")
- if size != "" && bitrate != "" {
- return size + " / " + bitrate
- }
- if size != "" {
- return size
- }
- return bitrate
-}
-
-func telegramVersionValue(event NotifyEvent, seasonEpisode string) string {
- if version := telegramFirstValue(event.Data, "version", "release_group"); version != "" && !strings.EqualFold(version, "best") {
- return version
- }
- return telegramVersionFromResourceTitle(
- telegramFirstValue(event.Data, "resource_title", "torrent_title", "release_title"),
- seasonEpisode,
- telegramFirstValue(event.Data, "year", "release_year"),
- )
-}
-
-func telegramVersionFromResourceTitle(raw, seasonEpisode, year string) string {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return ""
- }
- tail := ""
- if seasonEpisode != "" {
- upperRaw := strings.ToUpper(raw)
- upperEpisode := strings.ToUpper(seasonEpisode)
- if idx := strings.Index(upperRaw, upperEpisode); idx >= 0 {
- tail = raw[idx+len(seasonEpisode):]
- }
- }
- if tail == "" && year != "" {
- if idx := strings.LastIndex(raw, year); idx >= 0 {
- tail = raw[idx+len(year):]
- }
- }
- tail = strings.Trim(tail, " \t\r\n._-[]()【】")
- if tail == "" {
- return ""
- }
- tail = strings.TrimSuffix(tail, ".torrent")
- tail = strings.TrimSuffix(tail, ".mkv")
- tail = strings.TrimSuffix(tail, ".mp4")
- tail = strings.Join(strings.Fields(tail), ".")
- if len([]rune(tail)) > 72 {
- tail = string([]rune(tail)[:72]) + "..."
- }
- return tail
-}
-
-func telegramGenresValue(raw string) string {
- raw = strings.TrimSpace(strings.Trim(raw, "[]"))
- if raw == "" {
- return ""
- }
- parts := strings.FieldsFunc(raw, func(r rune) bool {
- return r == ',' || r == ',' || r == '/' || r == '|' || r == '、'
- })
- if len(parts) <= 1 {
- return raw
- }
- out := make([]string, 0, len(parts))
- seen := map[string]struct{}{}
- for _, part := range parts {
- part = strings.TrimSpace(part)
- if part == "" {
- continue
- }
- if _, ok := seen[part]; ok {
- continue
- }
- seen[part] = struct{}{}
- out = append(out, part)
- }
- return strings.Join(out, "、")
-}
-
-type telegramDataField struct {
- key string
- value string
-}
-
-func telegramDisplayData(data map[string]interface{}) []telegramDataField {
- if len(data) == 0 {
- return nil
- }
- keys := make([]string, 0, len(data))
- for key := range data {
- if telegramHiddenDataKey(key) {
- continue
- }
- keys = append(keys, key)
- }
- sort.Strings(keys)
- fields := make([]telegramDataField, 0, len(keys))
- for _, key := range keys {
- value := strings.TrimSpace(fmt.Sprint(data[key]))
- if value == "" || value == "" {
- continue
- }
- fields = append(fields, telegramDataField{key: key, value: value})
- }
- return fields
-}
-
-func telegramHiddenDataKey(key string) bool {
- switch strings.ToLower(strings.TrimSpace(key)) {
- case "photo_url", "poster_url", "poster", "image_url", "backdrop_url",
- "tmdb_url", "imdb_url", "douban_url", "detail_url", "external_url",
- "resource_title", "torrent_title", "release_title":
- return true
- default:
- return false
- }
-}
-
-func telegramFieldLabel(key string) string {
- switch strings.ToLower(strings.TrimSpace(key)) {
- case "title", "name":
- return "标题"
- case "original_title":
- return "原始片名"
- case "original_language":
- return "原始语言"
- case "year", "release_year":
- return "发行年份"
- case "save_path":
- return "保存路径"
- case "hash":
- return "Hash"
- case "media_type":
- return "媒体类型"
- case "media_category":
- return "类别"
- case "season_episode":
- return "季集"
- case "size", "bitrate":
- return "大小"
- case "version", "release_group":
- return "版本"
- case "rating":
- return "评分"
- case "genres":
- return "类型"
- case "overview":
- return "简介"
- case "subscription":
- return "订阅"
- case "queued":
- return "新增资源"
- default:
- return strings.TrimSpace(key)
- }
-}
-
-func telegramExternalLinks(data map[string]interface{}) string {
- if len(data) == 0 {
- return ""
- }
- links := []string{}
- for _, item := range []struct {
- key string
- name string
- }{
- {key: "tmdb_url", name: "TMDB"},
- {key: "imdb_url", name: "IMDB"},
- {key: "douban_url", name: "豆瓣"},
- } {
- value := telegramDataString(data, item.key)
- if isTelegramRemotePhotoURL(value) {
- links = append(links, fmt.Sprintf(`%s`, escapeHTML(value), escapeHTML(item.name)))
- }
- }
- if len(links) == 0 {
- return ""
- }
- return "🔗 外链:" + strings.Join(links, " / ")
-}
-
-func telegramEventPhotoURL(event NotifyEvent) string {
- for _, key := range []string{"photo_url", "poster_url", "poster", "image_url", "backdrop_url"} {
- value := telegramDataString(event.Data, key)
- if isTelegramRemotePhotoURL(value) {
- return value
- }
- }
- return ""
-}
-
-func telegramDataString(data map[string]interface{}, key string) string {
- if len(data) == 0 {
- return ""
- }
- for k, value := range data {
- if strings.EqualFold(strings.TrimSpace(k), key) {
- return telegramValueString(value)
- }
- }
- return ""
-}
-
-func telegramValueString(value interface{}) string {
- switch v := value.(type) {
- case nil:
- return ""
- case string:
- return strings.TrimSpace(v)
- case []string:
- return strings.TrimSpace(strings.Join(v, ","))
- case []interface{}:
- out := make([]string, 0, len(v))
- for _, item := range v {
- if s := telegramValueString(item); s != "" {
- out = append(out, s)
- }
- }
- return strings.Join(out, ",")
- case float32:
- return strings.TrimRight(strings.TrimRight(fmt.Sprintf("%.1f", v), "0"), ".")
- case float64:
- return strings.TrimRight(strings.TrimRight(fmt.Sprintf("%.1f", v), "0"), ".")
- default:
- text := strings.TrimSpace(fmt.Sprint(value))
- if text == "" {
- return ""
- }
- return text
- }
-}
-
-func isTelegramRemotePhotoURL(raw string) bool {
- raw = strings.TrimSpace(raw)
- if raw == "" {
- return false
- }
- u, err := url.Parse(raw)
- return err == nil && (u.Scheme == "http" || u.Scheme == "https") && u.Host != ""
-}
diff --git a/internal/service/notify_telegram_fields.go b/internal/service/notify_telegram_fields.go
new file mode 100644
index 0000000..43a760d
--- /dev/null
+++ b/internal/service/notify_telegram_fields.go
@@ -0,0 +1,390 @@
+package service
+
+import (
+ "fmt"
+ "net/url"
+ "sort"
+ "strconv"
+ "strings"
+)
+
+func telegramFirstValue(data map[string]interface{}, keys ...string) string {
+ for _, key := range keys {
+ if value := telegramDataString(data, key); value != "" {
+ return value
+ }
+ }
+ return ""
+}
+
+func telegramMessageFieldValue(message string, keys ...string) string {
+ if strings.TrimSpace(message) == "" {
+ return ""
+ }
+ for _, line := range strings.Split(message, "\n") {
+ key, value, ok := splitTelegramField(line)
+ if !ok {
+ continue
+ }
+ for _, want := range keys {
+ if strings.EqualFold(strings.TrimSpace(key), strings.TrimSpace(want)) {
+ return value
+ }
+ }
+ }
+ return ""
+}
+
+func telegramMediaCategory(data map[string]interface{}) string {
+ if category := telegramFirstValue(data, "media_category", "category"); category != "" {
+ return category
+ }
+ switch strings.ToLower(telegramFirstValue(data, "media_type")) {
+ case "movie":
+ return "电影"
+ case "tv", "series", "show":
+ return "剧集"
+ case "anime":
+ return "动漫"
+ case "variety":
+ return "综艺"
+ case "documentary":
+ return "纪录片"
+ default:
+ return telegramFirstValue(data, "media_type")
+ }
+}
+
+func telegramLanguageName(raw string) string {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return ""
+ }
+ parts := strings.FieldsFunc(raw, func(r rune) bool {
+ return r == ',' || r == ',' || r == '/' || r == '|' || r == '、'
+ })
+ if len(parts) == 0 {
+ parts = []string{raw}
+ }
+ seen := map[string]struct{}{}
+ out := []string{}
+ for _, part := range parts {
+ part = strings.TrimSpace(strings.Trim(part, "[]"))
+ if part == "" {
+ continue
+ }
+ lower := strings.ToLower(strings.ReplaceAll(part, "_", "-"))
+ name := part
+ switch {
+ case strings.HasPrefix(lower, "zh") || lower == "cn" || lower == "cmn":
+ name = "中文"
+ case lower == "en" || strings.HasPrefix(lower, "en-"):
+ name = "英语"
+ case lower == "ja" || lower == "jp" || strings.HasPrefix(lower, "ja-"):
+ name = "日语"
+ case lower == "ko" || lower == "kr" || strings.HasPrefix(lower, "ko-"):
+ name = "韩语"
+ case lower == "fr" || strings.HasPrefix(lower, "fr-"):
+ name = "法语"
+ case lower == "de" || strings.HasPrefix(lower, "de-"):
+ name = "德语"
+ case lower == "es" || strings.HasPrefix(lower, "es-"):
+ name = "西班牙语"
+ case lower == "it" || strings.HasPrefix(lower, "it-"):
+ name = "意大利语"
+ case lower == "ru" || strings.HasPrefix(lower, "ru-"):
+ name = "俄语"
+ case lower == "th" || strings.HasPrefix(lower, "th-"):
+ name = "泰语"
+ }
+ if _, ok := seen[name]; ok {
+ continue
+ }
+ seen[name] = struct{}{}
+ out = append(out, name)
+ }
+ return strings.Join(out, "、")
+}
+
+func telegramSeasonEpisodeValue(event NotifyEvent) string {
+ if value := telegramFirstValue(event.Data, "season_episode", "episode_tag"); value != "" {
+ return strings.ToUpper(value)
+ }
+ for _, raw := range []string{
+ telegramFirstValue(event.Data, "resource_title", "torrent_title", "release_title"),
+ telegramFirstValue(event.Data, "title", "name"),
+ event.Message,
+ } {
+ if value := telegramExtractSeasonEpisode(raw); value != "" {
+ return value
+ }
+ }
+ season := telegramFirstValue(event.Data, "season")
+ episode := telegramFirstValue(event.Data, "episode")
+ if season != "" && episode != "" {
+ return fmt.Sprintf("S%02dE%02d", telegramEpisodeNumber(season), telegramEpisodeNumber(episode))
+ }
+ return ""
+}
+
+func telegramEpisodeNumber(raw string) int {
+ raw = strings.TrimSpace(strings.TrimLeft(strings.ToUpper(raw), "SE"))
+ raw = strings.TrimLeft(raw, "0")
+ if raw == "" {
+ return 0
+ }
+ n, _ := strconv.Atoi(raw)
+ return n
+}
+
+func telegramExtractSeasonEpisode(raw string) string {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return ""
+ }
+ return strings.ToUpper(telegramSeasonEpisodePattern.FindString(raw))
+}
+
+func telegramSizeValue(data map[string]interface{}) string {
+ size := telegramFirstValue(data, "size")
+ bitrate := telegramFirstValue(data, "bitrate")
+ if size != "" && bitrate != "" {
+ return size + " / " + bitrate
+ }
+ if size != "" {
+ return size
+ }
+ return bitrate
+}
+
+func telegramVersionValue(event NotifyEvent, seasonEpisode string) string {
+ if version := telegramFirstValue(event.Data, "version", "release_group"); version != "" && !strings.EqualFold(version, "best") {
+ return version
+ }
+ return telegramVersionFromResourceTitle(
+ telegramFirstValue(event.Data, "resource_title", "torrent_title", "release_title"),
+ seasonEpisode,
+ telegramFirstValue(event.Data, "year", "release_year"),
+ )
+}
+
+func telegramVersionFromResourceTitle(raw, seasonEpisode, year string) string {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return ""
+ }
+ tail := ""
+ if seasonEpisode != "" {
+ upperRaw := strings.ToUpper(raw)
+ upperEpisode := strings.ToUpper(seasonEpisode)
+ if idx := strings.Index(upperRaw, upperEpisode); idx >= 0 {
+ tail = raw[idx+len(seasonEpisode):]
+ }
+ }
+ if tail == "" && year != "" {
+ if idx := strings.LastIndex(raw, year); idx >= 0 {
+ tail = raw[idx+len(year):]
+ }
+ }
+ tail = strings.Trim(tail, " \t\r\n._-[]()【】")
+ if tail == "" {
+ return ""
+ }
+ tail = strings.TrimSuffix(tail, ".torrent")
+ tail = strings.TrimSuffix(tail, ".mkv")
+ tail = strings.TrimSuffix(tail, ".mp4")
+ tail = strings.Join(strings.Fields(tail), ".")
+ if len([]rune(tail)) > 72 {
+ tail = string([]rune(tail)[:72]) + "..."
+ }
+ return tail
+}
+
+func telegramGenresValue(raw string) string {
+ raw = strings.TrimSpace(strings.Trim(raw, "[]"))
+ if raw == "" {
+ return ""
+ }
+ parts := strings.FieldsFunc(raw, func(r rune) bool {
+ return r == ',' || r == ',' || r == '/' || r == '|' || r == '、'
+ })
+ if len(parts) <= 1 {
+ return raw
+ }
+ out := make([]string, 0, len(parts))
+ seen := map[string]struct{}{}
+ for _, part := range parts {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ continue
+ }
+ if _, ok := seen[part]; ok {
+ continue
+ }
+ seen[part] = struct{}{}
+ out = append(out, part)
+ }
+ return strings.Join(out, "、")
+}
+
+type telegramDataField struct {
+ key string
+ value string
+}
+
+func telegramDisplayData(data map[string]interface{}) []telegramDataField {
+ if len(data) == 0 {
+ return nil
+ }
+ keys := make([]string, 0, len(data))
+ for key := range data {
+ if telegramHiddenDataKey(key) {
+ continue
+ }
+ keys = append(keys, key)
+ }
+ sort.Strings(keys)
+ fields := make([]telegramDataField, 0, len(keys))
+ for _, key := range keys {
+ value := strings.TrimSpace(fmt.Sprint(data[key]))
+ if value == "" || value == "" {
+ continue
+ }
+ fields = append(fields, telegramDataField{key: key, value: value})
+ }
+ return fields
+}
+
+func telegramHiddenDataKey(key string) bool {
+ switch strings.ToLower(strings.TrimSpace(key)) {
+ case "photo_url", "poster_url", "poster", "image_url", "backdrop_url",
+ "tmdb_url", "imdb_url", "douban_url", "detail_url", "external_url",
+ "resource_title", "torrent_title", "release_title":
+ return true
+ default:
+ return false
+ }
+}
+
+func telegramFieldLabel(key string) string {
+ switch strings.ToLower(strings.TrimSpace(key)) {
+ case "title", "name":
+ return "标题"
+ case "original_title":
+ return "原始片名"
+ case "original_language":
+ return "原始语言"
+ case "year", "release_year":
+ return "发行年份"
+ case "save_path":
+ return "保存路径"
+ case "hash":
+ return "Hash"
+ case "media_type":
+ return "媒体类型"
+ case "media_category":
+ return "类别"
+ case "season_episode":
+ return "季集"
+ case "size", "bitrate":
+ return "大小"
+ case "version", "release_group":
+ return "版本"
+ case "rating":
+ return "评分"
+ case "genres":
+ return "类型"
+ case "overview":
+ return "简介"
+ case "subscription":
+ return "订阅"
+ case "queued":
+ return "新增资源"
+ default:
+ return strings.TrimSpace(key)
+ }
+}
+
+func telegramExternalLinks(data map[string]interface{}) string {
+ if len(data) == 0 {
+ return ""
+ }
+ links := []string{}
+ for _, item := range []struct {
+ key string
+ name string
+ }{
+ {key: "tmdb_url", name: "TMDB"},
+ {key: "imdb_url", name: "IMDB"},
+ {key: "douban_url", name: "豆瓣"},
+ } {
+ value := telegramDataString(data, item.key)
+ if isTelegramRemotePhotoURL(value) {
+ links = append(links, fmt.Sprintf(`%s`, escapeHTML(value), escapeHTML(item.name)))
+ }
+ }
+ if len(links) == 0 {
+ return ""
+ }
+ return "🔗 外链:" + strings.Join(links, " / ")
+}
+
+func telegramEventPhotoURL(event NotifyEvent) string {
+ for _, key := range []string{"photo_url", "poster_url", "poster", "image_url", "backdrop_url"} {
+ value := telegramDataString(event.Data, key)
+ if isTelegramRemotePhotoURL(value) {
+ return value
+ }
+ }
+ return ""
+}
+
+func telegramDataString(data map[string]interface{}, key string) string {
+ if len(data) == 0 {
+ return ""
+ }
+ for k, value := range data {
+ if strings.EqualFold(strings.TrimSpace(k), key) {
+ return telegramValueString(value)
+ }
+ }
+ return ""
+}
+
+func telegramValueString(value interface{}) string {
+ switch v := value.(type) {
+ case nil:
+ return ""
+ case string:
+ return strings.TrimSpace(v)
+ case []string:
+ return strings.TrimSpace(strings.Join(v, ","))
+ case []interface{}:
+ out := make([]string, 0, len(v))
+ for _, item := range v {
+ if s := telegramValueString(item); s != "" {
+ out = append(out, s)
+ }
+ }
+ return strings.Join(out, ",")
+ case float32:
+ return strings.TrimRight(strings.TrimRight(fmt.Sprintf("%.1f", v), "0"), ".")
+ case float64:
+ return strings.TrimRight(strings.TrimRight(fmt.Sprintf("%.1f", v), "0"), ".")
+ default:
+ text := strings.TrimSpace(fmt.Sprint(value))
+ if text == "" {
+ return ""
+ }
+ return text
+ }
+}
+
+func isTelegramRemotePhotoURL(raw string) bool {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ return false
+ }
+ u, err := url.Parse(raw)
+ return err == nil && (u.Scheme == "http" || u.Scheme == "https") && u.Host != ""
+}
diff --git a/internal/service/organize_pipeline.go b/internal/service/organize_pipeline.go
index 980f819..21708b0 100644
--- a/internal/service/organize_pipeline.go
+++ b/internal/service/organize_pipeline.go
@@ -149,6 +149,7 @@ func (p *OrganizePipelineService) Run(ctx context.Context, req OrganizePipelineR
zap.String("dest", res.DestPath),
zap.Int("organized", res.Organized),
zap.Int("replaced", res.Replaced),
+ zap.Int("reclassified", res.Reclassified),
zap.Int("skipped", res.Skipped))
}
@@ -161,6 +162,7 @@ func (p *OrganizePipelineService) Run(ctx context.Context, req OrganizePipelineR
zap.String("dest", firstNonEmpty(res.DestPath, filepath.Dir(path))),
zap.Int("organized", res.Organized),
zap.Int("replaced", res.Replaced),
+ zap.Int("reclassified", res.Reclassified),
zap.Int("skipped", res.Skipped),
zap.Int("scrapes", len(res.Scrapes)),
zap.Int("errors", len(res.Errors)))
@@ -172,7 +174,7 @@ func organizeFatalResultError(res *OrganizeResult) error {
if res == nil || len(res.Errors) == 0 {
return nil
}
- if res.Organized > 0 || res.Replaced > 0 {
+ if res.Organized > 0 || res.Replaced > 0 || res.Reclassified > 0 {
return nil
}
samples := organizeErrorSamples(res.Errors, 3)
@@ -212,6 +214,7 @@ func (p *OrganizePipelineService) logOrganizeProblem(req OrganizePipelineRequest
zap.String("dest", res.DestPath),
zap.Int("organized", res.Organized),
zap.Int("replaced", res.Replaced),
+ zap.Int("reclassified", res.Reclassified),
zap.Int("skipped", res.Skipped),
zap.Int("errors", len(res.Errors)),
zap.Strings("error_samples", organizeErrorSamples(res.Errors, 5)),
@@ -260,6 +263,7 @@ func organizeAuditDetail(req OrganizePipelineRequest, res *OrganizeResult, err e
fmt.Sprintf("dest=%s", res.DestPath),
fmt.Sprintf("organized=%d", res.Organized),
fmt.Sprintf("replaced=%d", res.Replaced),
+ fmt.Sprintf("reclassified=%d", res.Reclassified),
fmt.Sprintf("skipped=%d", res.Skipped),
fmt.Sprintf("errors=%d", len(res.Errors)),
)
@@ -355,7 +359,7 @@ func organizeScanRoot(res *OrganizeResult, path string) string {
func organizeItemNeedsVisibilitySync(item OrganizePreviewItem) bool {
switch item.Action {
- case "organize", "replace":
+ case "organize", "replace", "reclassify", "cleanup":
return true
case "skip":
switch item.Reason {
diff --git a/internal/service/organizer.go b/internal/service/organizer.go
index 039d601..6aacedf 100644
--- a/internal/service/organizer.go
+++ b/internal/service/organizer.go
@@ -17,10 +17,6 @@ import (
"context"
"errors"
"fmt"
- "io"
- "os"
- "path/filepath"
- "strings"
"go.uber.org/zap"
@@ -52,55 +48,6 @@ func (o *OrganizerService) SetProbe(p *FFprobeService) { o.probe = p }
// metadata before it decides the final folder and filename.
func (o *OrganizerService) SetScraper(s *ScraperService) { o.scraper = s }
-// OrganizeResult reports what happened.
-type OrganizeResult struct {
- Organized int `json:"organized"`
- Skipped int `json:"skipped"`
- Replaced int `json:"replaced,omitempty"`
- Errors []string `json:"errors,omitempty"`
- SourcePath string `json:"source_path,omitempty"`
- DestPath string `json:"dest_path,omitempty"`
- DryRun bool `json:"dry_run,omitempty"`
- Items []OrganizePreviewItem `json:"items,omitempty"`
- Scans []OrganizeScanSummary `json:"scans,omitempty"`
- Scrapes []OrganizeScrapeSummary `json:"scrapes,omitempty"`
-}
-
-type OrganizePreviewItem struct {
- Source string `json:"source"`
- Target string `json:"target,omitempty"`
- Action string `json:"action"` // organize / skip / replace / error
- Reason string `json:"reason,omitempty"`
- MediaType string `json:"media_type,omitempty"`
- Category string `json:"category,omitempty"`
- Title string `json:"title,omitempty"`
-}
-
-// OrganizeOptions carries per-request overrides for an organize operation.
-// 空值表示沿用系统设置中的默认值。
-//
-// 整理是「从源目录整理到目的地目录」:SourcePath 指定待整理文件所在的源目录,
-// DestPath 指定整理输出的目的地目录。两者相互独立,不再混用同一个目录。
-type OrganizeOptions struct {
- // SourcePath 本次整理的源目录(待整理文件所在目录),覆盖 organize.source_dir
- // 设置与媒体库路径。仅整理位于该目录下的媒体;留空表示整个媒体库。
- SourcePath string
- // DestPath 本次整理的目的地根路径(整理输出到哪里),覆盖 organize.target_dir 设置。
- // 留空则使用设置中的默认目的地目录,再退回媒体库路径。
- DestPath string
- // TransferMode 本次整理的转移方式,覆盖 organize.transfer_mode 设置。
- TransferMode TransferMode
- // MediaType 手动整理时由 UI 指定的媒体类型。空值时按文件名/目录推断。
- MediaType string
- // MediaCategory 由订阅/下载任务或 UI 指定的分类。空值时按目录/NFO/规则推断。
- MediaCategory string
- // DryRun 仅生成整理预览,不实际移动/复制/硬链接文件。
- DryRun bool
- // AllowReplaceExisting 允许用本次来源替换目标库中已存在的同一媒体。
- // 默认 false:只去重不洗版,避免未开启洗版的订阅/手动整理留下或替换出多份版本。
- AllowReplaceExisting bool
-}
-
// OrganizeMedia moves a single media file into the target library directory.
// It auto-detects whether the media is a movie or TV episode based on the
// parsed season/episode numbers and builds the destination path accordingly.
@@ -112,197 +59,24 @@ func (o *OrganizerService) OrganizeMedia(ctx context.Context, mediaID string) (s
// OrganizeMediaWithOptions is OrganizeMedia with per-request overrides for the
// target path and transfer mode.
func (o *OrganizerService) OrganizeMediaWithOptions(ctx context.Context, mediaID string, opts OrganizeOptions) (string, error) {
- m, err := o.repo.Media.FindByID(ctx, mediaID)
- if err != nil || m == nil {
- return "", errors.New("media not found")
+ req, err := o.resolveOrganizeMediaRequest(ctx, mediaID, opts)
+ if err != nil {
+ return "", err
}
- lib, err := o.repo.Library.FindByID(ctx, m.LibraryID)
- if err != nil || lib == nil {
- return "", errors.New("library not found")
- }
- if _, ok := ParseCloudLibraryMount(lib.Path); ok {
- return "", errors.New("local organize cannot use cloud libraries directly; use external storage scan/mount for cloud media or enable cloud transfer to write to cloud")
- }
- baseRoot := redirectOrganizeStagingRoot(o.resolveBaseRoot(ctx, lib, opts.DestPath))
- if _, ok := ParseCloudLibraryMount(baseRoot); ok {
- return "", errors.New("organize destination must be a local writable media directory; enable cloud transfer in external storage when writing to cloud")
- }
- if !opts.DryRun {
- if err := ensureOrganizeDestinationWritable(baseRoot); err != nil {
- return "", err
- }
- }
- mode := o.resolveTransferMode(ctx, opts.TransferMode)
- if isSeriesLibraryType(lib.Type) {
- if err := o.refreshEpisodeIdentity(m, lib); err != nil {
- return "", err
- }
- }
- ext := filepath.Ext(m.Path)
- title := sanitizeFilename(m.Title)
- if title == "" {
- title = "Unknown"
- }
-
- // Determine category folder (if smart classify is enabled)
- category := o.SmartClassify(ctx, m)
-
- var dst string
- if isSeriesLibraryType(lib.Type) {
- root := o.organizeRoot(baseRoot, lib.Type, category)
- target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
- Root: categoryRoot(root, category),
- MediaType: lib.Type,
- Category: category,
- Title: title,
- Source: m.Path,
- Ext: ext,
- Year: m.Year,
- Season: m.SeasonNum,
- Episode: m.EpisodeNum,
- Series: true,
- })
- if err != nil {
- return "", err
- }
- dst = target.Path
- } else {
- root := o.organizeRoot(baseRoot, lib.Type, category)
- target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
- Root: categoryRoot(root, category),
- MediaType: lib.Type,
- Category: category,
- Title: title,
- Source: m.Path,
- Ext: ext,
- Year: m.Year,
- })
- if err != nil {
- return "", err
- }
- dst = target.Path
+ dst, err := o.buildOrganizeMediaDestination(ctx, req)
+ if err != nil {
+ return "", err
}
// Skip if already in place.
- if m.Path == dst {
- return dst, nil
+ if req.media.Path == dst.path {
+ return dst.path, nil
+ }
+ if req.dryRun {
+ return dst.path, nil
}
- // Refuse to overwrite an existing different file. 当多个 release(如
- // 不同字幕组、不同源)刮削后被统一改名,原本不重复的文件会被映射到
- // 同一个目标路径,盲目 move 会导致后者覆盖前者,造成数据丢失。
- if _, err := os.Stat(dst); err == nil {
- o.log.Warn("organize skipped: destination already exists",
- zap.String("media", m.ID),
- zap.String("from", m.Path),
- zap.String("to", dst))
- return dst, nil
- }
-
- // Create directories.
- if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
- return "", err
- }
-
- // Transfer the file according to the resolved mode. move 删除源;
- // copy/hardlink/symlink 保留源文件,从而让下载器可继续做种。
- if err := transferFile(m.Path, dst, mode); err != nil {
- return "", err
- }
-
- // Update the database row.
- if err := o.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Where("id = ?", m.ID).
- Updates(map[string]any{
- "path": dst,
- "season_num": m.SeasonNum,
- "episode_num": m.EpisodeNum,
- }).Error; err != nil {
- return dst, err
- }
- if err := transferSidecarNFO(m.Path, dst, mode); err != nil {
- o.log.Warn("organize sidecar nfo failed",
- zap.String("media", m.ID),
- zap.String("from", nfoPath(m.Path)),
- zap.String("to", nfoPath(dst)),
- zap.Error(err))
- }
- o.log.Info("organized",
- zap.String("media", m.ID),
- zap.String("from", m.Path),
- zap.String("to", dst),
- zap.String("category", category),
- zap.String("mode", string(mode)),
- )
- return dst, nil
-}
-
-// resolveBaseRoot picks the organize destination root (目的地目录): a
-// per-request override wins, then the organize.target_dir setting, then the
-// library's own path.
-func (o *OrganizerService) resolveBaseRoot(ctx context.Context, lib *model.Library, override string) string {
- if r := strings.TrimSpace(override); r != "" {
- return r
- }
- if o.repo != nil && o.repo.Setting != nil {
- if v, err := o.repo.Setting.Get(ctx, "organize.target_dir"); err == nil && strings.TrimSpace(v) != "" {
- return strings.TrimSpace(v)
- }
- }
- return lib.Path
-}
-
-// resolveSourceRoot picks the organize source root (源目录,待整理文件所在目录):
-// a per-request override wins, then the organize.source_dir setting, then the
-// library's own path. Library organize only touches media located under this
-// root, so operators can point at a specific download/staging folder.
-func (o *OrganizerService) resolveSourceRoot(ctx context.Context, lib *model.Library, override string) string {
- if r := strings.TrimSpace(override); r != "" {
- return r
- }
- if o.repo != nil && o.repo.Setting != nil {
- if v, err := o.repo.Setting.Get(ctx, "organize.source_dir"); err == nil && strings.TrimSpace(v) != "" {
- return strings.TrimSpace(v)
- }
- }
- return lib.Path
-}
-
-// resolveTransferMode picks the transfer mode: a per-request override wins,
-// otherwise the organize.transfer_mode setting (default move). When the
-// effective mode is move and 做种保种 (organize.keep_seeding) is enabled, it is
-// upgraded to hardlink so the source stays in place for the torrent client.
-func (o *OrganizerService) resolveTransferMode(ctx context.Context, override TransferMode) TransferMode {
- mode := override
- if mode == "" {
- mode = TransferMove
- if o.repo != nil && o.repo.Setting != nil {
- if v, err := o.repo.Setting.Get(ctx, "organize.transfer_mode"); err == nil && strings.TrimSpace(v) != "" {
- mode = parseTransferMode(v)
- }
- }
- }
- if mode == TransferMove && o.keepSeedingEnabled(ctx) {
- // 移动会删除源文件导致 qBittorrent 停止做种;保种开启时改用硬链接
- // 既规范命名又保留源文件继续做种上传。硬链接失败时会报错,避免静默
- // 退化复制后占用双份磁盘空间。
- return TransferHardlink
- }
- return mode
-}
-
-// keepSeedingEnabled reports whether 做种保种 is on. Defaults to true so an
-// unconfigured instance never silently breaks seeding on organize.
-func (o *OrganizerService) keepSeedingEnabled(ctx context.Context) bool {
- if o.repo == nil || o.repo.Setting == nil {
- return true
- }
- v, err := o.repo.Setting.Get(ctx, "organize.keep_seeding")
- if err != nil || strings.TrimSpace(v) == "" {
- return true
- }
- return v == "true" || v == "1" || v == "on"
+ return o.applyOrganizeMedia(ctx, req, dst)
}
// OrganizeLibrary organizes every media row in a library whose file is
@@ -344,6 +118,12 @@ func (o *OrganizerService) OrganizeLibraryWithOptions(ctx context.Context, libra
}
res := &OrganizeResult{SourcePath: sourceRoot, DestPath: baseRoot, DryRun: opts.DryRun}
for i := range rows {
+ if changed, err := o.reclassifyScannedMedia(ctx, rows[i], *lib, opts.DryRun, res); err != nil {
+ res.Errors = append(res.Errors, fmt.Sprintf("%s: %s", rows[i].Title, err.Error()))
+ continue
+ } else if changed {
+ continue
+ }
// 不在源目录内的文件跳过(不属于本次「从源目录整理」的范围)。
if !pathWithin(rows[i].Path, sourceRoot) {
res.Skipped++
@@ -395,295 +175,3 @@ func (o *OrganizerService) refreshEpisodeIdentity(m *model.Media, lib *model.Lib
m.EpisodeNum = episode
return nil
}
-
-// moveFile tries os.Rename first (instant on same fs), then falls back
-// to copy + remove for cross-device moves.
-//
-// 重要:如果 dst 已经存在,moveFile 会直接报错而不是覆盖。OrganizeMedia
-// 已经在调用前做过 stat 检查,这里是第二道防线。
-func moveFile(src, dst string) error {
- if _, err := os.Stat(dst); err == nil {
- return fmt.Errorf("destination already exists: %s", dst)
- }
- if err := os.Rename(src, dst); err == nil {
- return nil
- }
- // Cross-device: stream copy → remove. This can temporarily consume the
- // destination file size while copying, but the source is removed after the
- // copy succeeds.
- in, err := os.Open(src) // #nosec G304 -- src is selected from configured media/download roots by the organizer.
- if err != nil {
- return err
- }
- defer in.Close()
- // O_EXCL 保证不会覆盖已存在的目标。
- f, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) // #nosec G304,G302 -- dst is organizer-generated; media files must remain readable by local players.
- if err != nil {
- return err
- }
- if _, werr := io.Copy(f, in); werr != nil {
- _ = f.Close()
- _ = os.Remove(dst)
- return werr
- }
- if cerr := f.Close(); cerr != nil {
- return cerr
- }
- return os.Remove(src)
-}
-
-// transferSidecarNFO moves/copies/links the .nfo sidecar alongside its media
-// using the same transfer mode, so metadata follows the organized file.
-func transferSidecarNFO(srcMedia, dstMedia string, mode TransferMode) error {
- src := nfoPath(srcMedia)
- dst := nfoPath(dstMedia)
- if src == dst {
- return nil
- }
- if _, err := os.Stat(src); err != nil {
- if os.IsNotExist(err) {
- return nil
- }
- return err
- }
- if _, err := os.Stat(dst); err == nil {
- return nil
- }
- if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { // #nosec G301 -- sidecar media directories must remain readable by NAS/player users.
- return err
- }
- return transferFile(src, dst, mode)
-}
-
-// sanitizeFilename removes characters not safe for filesystem names.
-func sanitizeFilename(s string) string {
- r := strings.NewReplacer(
- "/", " ", "\\", " ", ":", " ", "*", "", "?", "",
- "\"", "", "<", "", ">", "", "|", "",
- )
- return strings.TrimSpace(r.Replace(s))
-}
-
-func (o *OrganizerService) organizeRoot(libraryPath, mediaType, category string) string {
- typeDir := o.mediaTypeRootDirForCategory(mediaType, category)
- if typeDir == "" || pathAlreadyEndsWith(libraryPath, typeDir) {
- return libraryPath
- }
- if isGenericMediaRoot(libraryPath) {
- return filepath.Join(libraryPath, typeDir)
- }
- return libraryPath
-}
-
-func (o *OrganizerService) mediaTypeRootDirForCategory(mediaType, category string) string {
- if root := o.categoryPhysicalRootDir(category); root != "" {
- return root
- }
- return mediaTypeRootDir(mediaType)
-}
-
-func (o *OrganizerService) categoryPhysicalRootDir(category string) string {
- key := normalizeOrganizeCategoryKey(category)
- if key == "" {
- return ""
- }
- categories := o.categoryMap()
- match := func(values ...string) bool {
- for _, value := range values {
- if key == normalizeOrganizeCategoryKey(value) {
- return true
- }
- }
- return false
- }
- switch {
- case match(
- categoryName(categories, "cn_anime", "国漫"),
- categoryName(categories, "jp_anime", "日番"),
- categoryName(categories, "children", "儿童"),
- "国漫", "国产动漫", "日番", "番剧", "日漫", "日本动漫", "日本动画", "儿童", "少儿",
- ):
- return "动漫"
- case match(
- categoryName(categories, "domestic_tv", "国产剧"),
- categoryName(categories, "euus_tv", "欧美剧"),
- categoryName(categories, "jk_tv", "日韩剧"),
- categoryName(categories, "variety", "综艺"),
- categoryName(categories, "documentary", "纪录片"),
- categoryName(categories, "uncategorized_tv", "未分类"),
- "国产剧", "欧美剧", "日韩剧", "日剧", "韩剧", "综艺", "真人秀", "纪录片", "纪录", "未分类",
- ):
- return "电视剧"
- case match(
- categoryName(categories, "animation_movie", "动画电影"),
- categoryName(categories, "chinese_movie", "华语电影"),
- categoryName(categories, "foreign_movie", "外语电影"),
- categoryName(categories, "euus_movie", "欧美电影"),
- categoryName(categories, "jk_movie", "日韩电影"),
- "动画电影", "动漫电影", "华语电影", "国产电影", "外语电影", "欧美电影", "日韩电影",
- ):
- return "电影"
- case match(categoryName(categories, "adult", "成人"), categoryName(categories, "adult_9kg", "9KG"), categoryName(categories, "adult_jav", "番号"), "成人", "9kg", "番号", "jav"):
- return "成人"
- default:
- return ""
- }
-}
-
-func categoryRoot(root, category string) string {
- if strings.TrimSpace(category) == "" || pathAlreadyEndsWith(root, category) {
- return root
- }
- return filepath.Join(root, category)
-}
-
-func pathWithin(path, root string) bool {
- cleanPath := filepath.Clean(path)
- cleanRoot := filepath.Clean(root)
- if strings.EqualFold(cleanPath, cleanRoot) {
- return true
- }
- rel, err := filepath.Rel(cleanRoot, cleanPath)
- if err != nil {
- return false
- }
- return rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator))
-}
-
-func mediaTypeRootDir(mediaType string) string {
- switch normalizeMediaType(mediaType, "", "") {
- case "movie":
- return "电影"
- case "anime":
- return "动漫"
- case "tv", "variety":
- return "电视剧"
- case "adult":
- return "成人"
- default:
- return ""
- }
-}
-
-func isGenericMediaRoot(path string) bool {
- base := strings.ToLower(strings.TrimSpace(filepath.Base(filepath.Clean(path))))
- switch base {
- case "media", "medias", "library", "libraries", "organized", "整理":
- return true
- default:
- return false
- }
-}
-
-// organizeStagingFolderNames 列出"手动整理"类暂存目录名。这些目录只是修正错误
-// 入库时的中转工作区,不能作为一级分类目录留存。整理时若目标根落在这类目录
-// 内,应重定向到其父级媒体根,让媒体真正归入 电影/电视剧/动漫/成人 的二级分类
-// 目录中(如 媒体/电影/华语电影/片名),而不是停留在暂存目录下。
-func organizeStagingFolderNames() map[string]struct{} {
- return map[string]struct{}{
- "手动整理": {}, "手动整理入库": {}, "待整理": {}, "待分类": {},
- "manual": {}, "manual_organize": {}, "manualorganize": {}, "staging": {}, "inbox": {},
- }
-}
-
-func isOrganizeStagingDir(path string) bool {
- base := strings.ToLower(strings.TrimSpace(filepath.Base(filepath.Clean(path))))
- if base == "" {
- return false
- }
- _, ok := organizeStagingFolderNames()[base]
- return ok
-}
-
-// redirectOrganizeStagingRoot 把"手动整理"类暂存目录的目标根重定向到父级媒体根。
-// 连续多层暂存目录(如 .../media/手动整理/待整理)会被逐层上提到真正的媒体根,
-// 随后分类逻辑会补上 电影/电视剧/动漫/成人 等顶层与二级分类。
-func redirectOrganizeStagingRoot(root string) string {
- cleaned := filepath.Clean(strings.TrimSpace(root))
- if cleaned == "" || cleaned == "." {
- return root
- }
- for isOrganizeStagingDir(cleaned) {
- parent := filepath.Dir(cleaned)
- if parent == cleaned || parent == "." || parent == string(filepath.Separator) {
- break
- }
- cleaned = parent
- }
- return cleaned
-}
-
-func pathAlreadyEndsWith(path, suffix string) bool {
- base := strings.TrimSpace(filepath.Base(filepath.Clean(path)))
- return strings.EqualFold(base, suffix)
-}
-
-func isSeriesLibraryType(mediaType string) bool {
- switch normalizeMediaType(mediaType, "", "") {
- case "tv", "anime", "variety":
- return true
- default:
- return false
- }
-}
-
-// isSmartClassifyEnabled checks if smart classify is enabled.
-// It first checks the database setting, then falls back to config.yaml.
-func (o *OrganizerService) isSmartClassifyEnabled(ctx context.Context) bool {
- // Try database first
- if o.repo != nil && o.repo.Setting != nil {
- val, err := o.repo.Setting.Get(ctx, "organizer.smart_classify")
- if err == nil && val != "" {
- return val == "true" || val == "1" || val == "on"
- }
- }
- // Fallback to config.yaml
- if o == nil || o.cfg == nil {
- return false
- }
- return o.cfg.Organizer.SmartClassify
-}
-
-// SmartClassify determines the subcategory folder based on media metadata.
-// It returns the category folder name (e.g., "华语电影", "欧美剧", "日番").
-// Returns empty string if smart classify is disabled or metadata is insufficient.
-func (o *OrganizerService) SmartClassify(ctx context.Context, m *model.Media) string {
- // Check if smart classify is enabled (from database first, then config)
- smartClassify := o.isSmartClassifyEnabled(ctx)
- if !smartClassify {
- return ""
- }
-
- // Determine media type from library
- lib, err := o.repo.Library.FindByID(ctx, m.LibraryID)
- if err != nil || lib == nil {
- return ""
- }
- return o.classifyMedia(ctx, m, lib.Type)
-}
-
-// parseCommaList splits a comma-separated string into a slice of trimmed strings.
-func parseCommaList(s string) []string {
- if s == "" {
- return nil
- }
- parts := strings.Split(s, ",")
- result := make([]string, 0, len(parts))
- for _, p := range parts {
- trimmed := strings.TrimSpace(p)
- if trimmed != "" {
- result = append(result, trimmed)
- }
- }
- return result
-}
-
-// contains checks if a string slice contains a specific string.
-func contains(slice []string, s string) bool {
- for _, v := range slice {
- if v == s {
- return true
- }
- }
- return false
-}
diff --git a/internal/service/organizer_classify.go b/internal/service/organizer_classify.go
new file mode 100644
index 0000000..25016a3
--- /dev/null
+++ b/internal/service/organizer_classify.go
@@ -0,0 +1,60 @@
+package service
+
+import (
+ "context"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// isSmartClassifyEnabled checks database settings first, then config.yaml.
+func (o *OrganizerService) isSmartClassifyEnabled(ctx context.Context) bool {
+ if o.repo != nil && o.repo.Setting != nil {
+ val, err := o.repo.Setting.Get(ctx, "organizer.smart_classify")
+ if err == nil && val != "" {
+ return val == "true" || val == "1" || val == "on"
+ }
+ }
+ if o == nil || o.cfg == nil {
+ return false
+ }
+ return o.cfg.Organizer.SmartClassify
+}
+
+// SmartClassify determines the subcategory folder based on media metadata.
+// It returns values such as "华语电影", "欧美剧", or "日番".
+func (o *OrganizerService) SmartClassify(ctx context.Context, m *model.Media) string {
+ if !o.isSmartClassifyEnabled(ctx) {
+ return ""
+ }
+ lib, err := o.repo.Library.FindByID(ctx, m.LibraryID)
+ if err != nil || lib == nil {
+ return ""
+ }
+ return o.classifyMedia(ctx, m, lib.Type)
+}
+
+// parseCommaList splits a comma-separated string into trimmed non-empty values.
+func parseCommaList(s string) []string {
+ if s == "" {
+ return nil
+ }
+ parts := strings.Split(s, ",")
+ result := make([]string, 0, len(parts))
+ for _, p := range parts {
+ trimmed := strings.TrimSpace(p)
+ if trimmed != "" {
+ result = append(result, trimmed)
+ }
+ }
+ return result
+}
+
+func contains(slice []string, s string) bool {
+ for _, v := range slice {
+ if v == s {
+ return true
+ }
+ }
+ return false
+}
diff --git a/internal/service/organizer_directory.go b/internal/service/organizer_directory.go
index 2435f51..bddf685 100644
--- a/internal/service/organizer_directory.go
+++ b/internal/service/organizer_directory.go
@@ -18,16 +18,14 @@ import (
"context"
"errors"
"fmt"
- "os"
"path/filepath"
"strings"
- "unicode"
"go.uber.org/zap"
-
- "github.com/ShukeBta/MediaStationGo/internal/model"
)
+var ErrUnsupportedOrganizeSource = errors.New("source is not a supported video file")
+
const (
organizeSkipAlreadyOrganized = "already organized"
organizeSkipDuplicateLibrary = "duplicate in library"
@@ -35,83 +33,6 @@ const (
organizeSkipSampleClip = "sample/trailer clip"
)
-// OrganizeSourceCandidate is a selectable organize source directory surfaced to
-// the UI so operators can organize an arbitrary directory (such as the download
-// directory) and not only registered libraries.
-type OrganizeSourceCandidate struct {
- Label string `json:"label"`
- Path string `json:"path"`
- Kind string `json:"kind"` // "download" | "media"
-}
-
-// OrganizeSourceCandidates returns the configured directories that are valid
-// organize sources (download dir + media dir). It uses the container-visible
-// paths; in NAS direct-read mode those equal the host paths the operator sees.
-func (o *OrganizerService) OrganizeSourceCandidates(ctx context.Context) []OrganizeSourceCandidate {
- out := []OrganizeSourceCandidate{}
- seen := map[string]struct{}{}
- add := func(label, path, kind string) {
- path = strings.TrimSpace(path)
- if path == "" || path == "." || strings.HasPrefix(path, ".") {
- return
- }
- clean := filepath.Clean(path)
- if !isAccessibleDir(clean) {
- return
- }
- if _, ok := seen[clean]; ok {
- return
- }
- seen[clean] = struct{}{}
- out = append(out, OrganizeSourceCandidate{Label: label, Path: clean, Kind: kind})
- }
- add("默认整理源", o.settingValue(ctx, "organize.source_dir"), "source")
- add("下载器保存目录", o.settingValue(ctx, "qbittorrent.savepath"), "download")
- add("下载目录", envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads"), "download")
- add("媒体目录", envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media"), "media")
- return out
-}
-
-func (o *OrganizerService) settingValue(ctx context.Context, key string) string {
- if o.repo == nil || o.repo.Setting == nil {
- return ""
- }
- if v, err := o.repo.Setting.Get(ctx, key); err == nil {
- return strings.TrimSpace(v)
- }
- return ""
-}
-
-// defaultSourceRoot resolves the source root for a directory organize:
-// explicit override → organize.source_dir setting → qB default save path →
-// download container dir.
-func (o *OrganizerService) defaultSourceRoot(ctx context.Context, override string) string {
- if r := strings.TrimSpace(override); r != "" {
- return r
- }
- if v := o.settingValue(ctx, "organize.source_dir"); v != "" {
- return v
- }
- if v := o.settingValue(ctx, "qbittorrent.savepath"); v != "" {
- return v
- }
- return envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")
-}
-
-// defaultDestRoot resolves the destination root for a directory organize:
-// explicit override → organize.target_dir setting → media container dir.
-func (o *OrganizerService) defaultDestRoot(ctx context.Context, override string) string {
- if r := strings.TrimSpace(override); r != "" {
- return r
- }
- if o.repo != nil && o.repo.Setting != nil {
- if v, err := o.repo.Setting.Get(ctx, "organize.target_dir"); err == nil && strings.TrimSpace(v) != "" {
- return strings.TrimSpace(v)
- }
- }
- return envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")
-}
-
// OrganizeDirectory organizes every video file found under opts.SourcePath into
// the destination root, applying dedup + 洗版 (resolution replacement).
func (o *OrganizerService) OrganizeDirectory(ctx context.Context, opts OrganizeOptions) (*OrganizeResult, error) {
@@ -142,7 +63,7 @@ func (o *OrganizerService) OrganizeDirectory(ctx context.Context, opts OrganizeO
if !info.IsDir() {
ext := strings.ToLower(filepath.Ext(source))
if _, ok := videoExtensions[ext]; !ok {
- return nil, fmt.Errorf("source is not a supported video file: %s", source)
+ return nil, fmt.Errorf("%w: %s", ErrUnsupportedOrganizeSource, source)
}
if skipped, reason := shouldSkipOrganizeSourceVideo(source, filepath.Dir(source)); skipped {
res.Skipped++
@@ -150,7 +71,18 @@ func (o *OrganizerService) OrganizeDirectory(ctx context.Context, opts OrganizeO
o.logOrganizeDirectoryResult("organize file finished", res, mode)
return res, nil
}
- if err := o.organizeSourceFile(ctx, source, filepath.Dir(source), dest, mode, opts.MediaType, opts.MediaCategory, opts.DryRun, opts.AllowReplaceExisting, metadataCache, res); err != nil {
+ if err := o.organizeSourceFile(ctx, organizeSourceFileRequest{
+ Source: source,
+ SourceRoot: filepath.Dir(source),
+ DestRoot: dest,
+ Mode: mode,
+ MediaTypeOverride: opts.MediaType,
+ MediaCategoryOverride: opts.MediaCategory,
+ DryRun: opts.DryRun,
+ AllowReplaceExisting: opts.AllowReplaceExisting,
+ MetadataCache: metadataCache,
+ Result: res,
+ }); err != nil {
res.Errors = append(res.Errors, fmt.Sprintf("%s: %s", filepath.Base(source), err.Error()))
res.Items = append(res.Items, OrganizePreviewItem{Source: source, Action: "error", Reason: err.Error()})
}
@@ -170,7 +102,18 @@ func (o *OrganizerService) OrganizeDirectory(ctx context.Context, opts OrganizeO
res.Items = append(res.Items, OrganizePreviewItem{Source: path, Action: "skip", Reason: reason})
return nil
}
- if err := o.organizeSourceFile(ctx, path, source, dest, mode, opts.MediaType, opts.MediaCategory, opts.DryRun, opts.AllowReplaceExisting, metadataCache, res); err != nil {
+ if err := o.organizeSourceFile(ctx, organizeSourceFileRequest{
+ Source: path,
+ SourceRoot: source,
+ DestRoot: dest,
+ Mode: mode,
+ MediaTypeOverride: opts.MediaType,
+ MediaCategoryOverride: opts.MediaCategory,
+ DryRun: opts.DryRun,
+ AllowReplaceExisting: opts.AllowReplaceExisting,
+ MetadataCache: metadataCache,
+ Result: res,
+ }); err != nil {
res.Errors = append(res.Errors, fmt.Sprintf("%s: %s", filepath.Base(path), err.Error()))
res.Items = append(res.Items, OrganizePreviewItem{Source: path, Action: "error", Reason: err.Error()})
}
@@ -193,6 +136,7 @@ func (o *OrganizerService) logOrganizeDirectoryResult(message string, res *Organ
zap.String("mode", string(mode)),
zap.Int("organized", res.Organized),
zap.Int("replaced", res.Replaced),
+ zap.Int("reclassified", res.Reclassified),
zap.Int("skipped", res.Skipped),
zap.Int("errors", len(res.Errors)),
zap.Any("skip_reasons", OrganizeSkipReasonCounts(res)),
@@ -204,1150 +148,3 @@ func (o *OrganizerService) logOrganizeDirectoryResult(message string, res *Organ
}
o.log.Info(message, fields...)
}
-
-func ensureOrganizeDestinationWritable(dest string) error {
- dest = strings.TrimSpace(dest)
- if dest == "" || dest == "." {
- return errors.New("destination path required")
- }
- if _, ok := ParseCloudLibraryMount(dest); ok {
- return errors.New("organize destination must be a local writable media directory; enable cloud transfer in external storage when writing to cloud")
- }
- if err := os.MkdirAll(dest, 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
- return fmt.Errorf("destination path is not a writable directory: %s: %w", dest, err)
- }
- probe, err := os.CreateTemp(dest, ".mediastation-write-test-*") // #nosec G304 -- dest is operator-configured organize root.
- if err != nil {
- return fmt.Errorf("destination path is not writable: %s: %w", dest, err)
- }
- name := probe.Name()
- if closeErr := probe.Close(); closeErr != nil {
- _ = os.Remove(name)
- return fmt.Errorf("destination path write probe failed: %s: %w", dest, closeErr)
- }
- if err := os.Remove(name); err != nil {
- return fmt.Errorf("destination path cleanup probe failed: %s: %w", dest, err)
- }
- return nil
-}
-
-type organizeDirectoryLayout struct {
- MediaType string
- Category string
-}
-
-// organizeSourceFile organizes a single video file from the source directory
-// into destRoot, applying dedup + 洗版.
-func (o *OrganizerService) organizeSourceFile(ctx context.Context, src, sourceRoot, destRoot string, mode TransferMode, mediaTypeOverride, mediaCategoryOverride string, dryRun bool, allowReplaceExisting bool, metadataCache map[string]*Match, res *OrganizeResult) error {
- ext := filepath.Ext(src)
- season, episode := ParseEpisode(src)
- title, year := CleanQuery(src)
- if organizeWeakFileTitle(title) {
- if folderTitle, folderYear := organizeTitleFromParentFolder(src, sourceRoot, season > 0 || episode > 0); folderTitle != "" {
- title = folderTitle
- if year <= 0 {
- year = folderYear
- }
- } else {
- title = strings.TrimSuffix(filepath.Base(src), ext)
- }
- }
- // CleanQuery lowercases the parsed title; title-case it so organized output
- // matches typical library casing (and stays consistent for dedup).
- parsedTitle := title
- title = sanitizeFilename(titleCaseWords(title))
- if title == "" {
- title = "Unknown"
- }
- pathLayout := o.inferOrganizeDirectoryLayout(src, sourceRoot)
- layout := pathLayout
- forcedType := normalizeOrganizeMediaType(mediaTypeOverride)
- inferredType := o.inferMediaTypeForSourceFile(src, title, season, episode)
- if forcedType != "" {
- if layout.Category != "" && layout.MediaType != "" && layout.MediaType != forcedType {
- layout.Category = ""
- }
- layout.MediaType = forcedType
- } else if inferredType != "" {
- if inferredType == "tv" && layout.MediaType == "movie" {
- // 文件名中明确有季/集信息时,目录名只能作为弱提示;否则
- // 下载到错误的“电影/外语电影”等目录会把剧集按电影入库。
- layout = organizeDirectoryLayout{MediaType: inferredType}
- } else if layout.MediaType == "" {
- layout.MediaType = inferredType
- }
- }
- var metadataMatch *Match
- if match := o.lookupOrganizeMetadata(ctx, src, sourceRoot, layout.MediaType, title, year, season, episode, metadataCache); match != nil {
- metadataMatch = match
- if matchedTitle := sanitizeFilename(strings.TrimSpace(match.Title)); matchedTitle != "" {
- title = matchedTitle
- parsedTitle = strings.TrimSpace(match.Title)
- }
- if match.Year > 0 {
- year = match.Year
- }
- }
- if category := strings.TrimSpace(mediaCategoryOverride); category != "" {
- layout.Category = sanitizeFilename(category)
- } else if category := o.smartClassifySourceFile(ctx, src, sourceRoot, layout.MediaType, title, parsedTitle, metadataMatch); category != "" {
- // 智能分类以识别后的元数据为主,下载/源目录只作为
- // 兜底提示。这里即使源目录已有二级分类,也允许 TMDb/Bangumi/NFO
- // 识别结果修正到真正的分类,避免错误目录导致错误入库。
- layout.Category = category
- }
- if forcedType == "" {
- if impliedType, normalizedCategory := o.mediaTypeForDirectoryCategory(layout.Category); impliedType != "" {
- layout.Category = normalizedCategory
- if layout.MediaType == "" || layout.MediaType == "tv" || layout.MediaType == "anime" || pathLayout.Category != layout.Category {
- layout.MediaType = impliedType
- }
- }
- }
- layoutRoot, matchedLibrary := o.organizeLibraryRootForLayout(ctx, destRoot, layout.MediaType, layout.Category)
- if !matchedLibrary && layout.MediaType != "" {
- layoutRoot = o.organizeRoot(destRoot, layout.MediaType, layout.Category)
- }
- if !matchedLibrary && layout.Category != "" {
- layoutRoot = categoryRoot(layoutRoot, sanitizeFilename(layout.Category))
- }
- if !matchedLibrary && !dryRun {
- o.ensureOrganizeLibraryForRoot(ctx, layoutRoot, layout.MediaType, layout.Category)
- }
-
- var destDir, dst, episodeTag string
- isSeries := season > 0 || episode > 0
- if layout.MediaType != "" {
- isSeries = isSeriesLibraryType(layout.MediaType) && (season > 0 || episode > 0)
- }
- target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
- Root: layoutRoot,
- MediaType: layout.MediaType,
- Category: layout.Category,
- Title: title,
- Source: src,
- Ext: ext,
- Year: year,
- Season: season,
- Episode: episode,
- Series: isSeries,
- })
- if err != nil {
- return err
- }
- destDir = target.Dir
- dst = target.Path
- episodeTag = target.EpisodeTag
-
- // 源文件已经位于目标位置:无需处理。
- if filepath.Clean(src) == filepath.Clean(dst) {
- res.Skipped++
- res.Items = append(res.Items, OrganizePreviewItem{
- Source: src, Target: dst, Action: "skip", Reason: organizeSkipAlreadyOrganized,
- MediaType: layout.MediaType, Category: layout.Category, Title: title,
- })
- return nil
- }
-
- // 去重候选:合并「目的地媒体库已扫描入库的同一媒体(按标题/年份/季集匹配,
- // 不受目录大小写或布局影响)」与「目标文件夹内已存在的同名视频文件」。
- externalExisting := o.existingByExternalIdentity(ctx, destRoot, metadataMatch, season, episode)
- identityExisting := o.existingByIdentity(ctx, destRoot, parsedTitle, year, season, episode)
- folderExisting := o.existingByFolder(destDir, episodeTag)
- existing := mergeExistingVersionPaths(externalExisting, identityExisting, folderExisting)
- if len(existing) > 0 {
- srcArea := o.resolutionArea(ctx, src)
- bestArea := 0
- for _, e := range existing {
- if a := o.resolutionArea(ctx, e); a > bestArea {
- bestArea = a
- }
- }
- // 洗版:仅当来源与已存在版本的分辨率都可判定、且来源更高时才替换;
- // 任一方分辨率未知时保守跳过,绝不删除无法判定的已存在文件。
- if allowReplaceExisting && srcArea > 0 && bestArea > 0 && srcArea > bestArea {
- res.Items = append(res.Items, OrganizePreviewItem{
- Source: src, Target: dst, Action: "replace", Reason: "higher resolution",
- MediaType: layout.MediaType, Category: layout.Category, Title: title,
- })
- if dryRun {
- res.Replaced++
- return nil
- }
- if err := o.replaceVersions(ctx, src, existing, dst, mode); err != nil {
- return err
- }
- o.log.Info("organize replaced lower-resolution media",
- zap.String("from", src),
- zap.String("to", dst),
- zap.Int("src_area", srcArea),
- zap.Int("existing_area", bestArea),
- )
- res.Replaced++
- return nil
- }
- // 去重:目的地已存在同一媒体且不低于来源分辨率,跳过不再整理过去。
- reason := organizeSkipTargetExists
- if len(externalExisting) > 0 || len(identityExisting) > 0 || o.allExistingPathsInDB(ctx, existing) {
- reason = organizeSkipDuplicateLibrary
- }
- o.log.Debug("organize skip duplicate",
- zap.String("src", src), zap.String("dest_dir", destDir), zap.String("reason", reason))
- res.Skipped++
- res.Items = append(res.Items, OrganizePreviewItem{
- Source: src, Target: dst, Action: "skip", Reason: reason,
- MediaType: layout.MediaType, Category: layout.Category, Title: title,
- })
- return nil
- }
-
- res.Items = append(res.Items, OrganizePreviewItem{
- Source: src, Target: dst, Action: "organize",
- MediaType: layout.MediaType, Category: layout.Category, Title: title,
- })
- if dryRun {
- res.Organized++
- return nil
- }
- if err := os.MkdirAll(destDir, 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
- return err
- }
- if _, err := os.Stat(dst); err == nil {
- res.Skipped++
- if len(res.Items) > 0 {
- res.Items[len(res.Items)-1].Action = "skip"
- res.Items[len(res.Items)-1].Reason = organizeSkipTargetExists
- }
- return nil
- }
- if err := transferFile(src, dst, mode); err != nil {
- return err
- }
- if err := transferSidecarNFO(src, dst, mode); err != nil {
- o.log.Warn("organize sidecar nfo failed",
- zap.String("from", src), zap.String("to", dst), zap.Error(err))
- }
- res.Organized++
- return nil
-}
-
-func organizeTitleFromParentFolder(src, sourceRoot string, seriesLike bool) (string, int) {
- if !seriesLike {
- return "", 0
- }
- raw := seriesFolderTitle(src, sourceRoot)
- if strings.TrimSpace(raw) == "" {
- return "", 0
- }
- title, year := CleanQuery(raw)
- if title == "" {
- title = strings.TrimSpace(raw)
- }
- return title, year
-}
-
-func shouldSkipOrganizeSourceVideo(path, sourceRoot string) (bool, string) {
- cleanPath := filepath.Clean(path)
- cleanRoot := filepath.Clean(sourceRoot)
- if rel, err := filepath.Rel(cleanRoot, cleanPath); err == nil && rel != "." && !strings.HasPrefix(rel, "..") {
- dir := filepath.Dir(rel)
- if dir != "." {
- for _, part := range strings.Split(dir, string(os.PathSeparator)) {
- switch normalizeOrganizeCategoryKey(part) {
- case "sample", "samples", "trailer", "trailers", "preview", "previews", "teaser", "teasers":
- return true, organizeSkipSampleClip
- }
- }
- }
- }
- base := strings.ToLower(strings.TrimSuffix(filepath.Base(cleanPath), filepath.Ext(cleanPath)))
- normalized := strings.NewReplacer("_", " ", "-", " ", ".", " ").Replace(base)
- fields := strings.Fields(normalized)
- if len(fields) == 0 {
- return false, ""
- }
- if len(fields) == 1 && strings.HasPrefix(fields[0], "sample") {
- return true, organizeSkipSampleClip
- }
- last := fields[len(fields)-1]
- switch last {
- case "sample", "trailer", "preview", "teaser":
- return true, organizeSkipSampleClip
- }
- return false, ""
-}
-
-func normalizeOrganizeMediaType(mediaType string) string {
- switch strings.ToLower(strings.TrimSpace(mediaType)) {
- case "movie", "film":
- return "movie"
- case "tv", "series", "show", "drama":
- return "tv"
- case "anime", "animation":
- return "anime"
- case "variety":
- return "variety"
- case "adult", "nsfw":
- return "adult"
- default:
- return ""
- }
-}
-
-func organizeWeakFileTitle(title string) bool {
- title = strings.TrimSpace(title)
- if title == "" {
- return true
- }
- fields := strings.Fields(strings.ToLower(title))
- if len(fields) == 0 {
- return true
- }
- meaningful := 0
- for _, field := range fields {
- if _, ok := noiseTokenSet[field]; ok {
- continue
- }
- if _, ok := releaseBoundaryTokenSet[field]; ok {
- continue
- }
- if len(field) == 4 && strings.HasPrefix(field, "20") {
- continue
- }
- meaningful++
- }
- return meaningful == 0
-}
-
-func (o *OrganizerService) inferMediaTypeForSourceFile(src, title string, season, episode int) string {
- if season > 0 || episode > 0 {
- return "tv"
- }
- return normalizeMediaType("", title, src)
-}
-
-func (o *OrganizerService) lookupOrganizeMetadata(ctx context.Context, src, sourceRoot, mediaType, title string, year, season, episode int, cache map[string]*Match) *Match {
- seriesLike := isSeriesLibraryType(mediaType) || season > 0 || episode > 0
- if local, err := ReadLocalMetadata(src, sourceRoot, seriesLike); err == nil && local != nil {
- if match := organizeMatchFromLocalMetadata(local); match != nil {
- return match
- }
- } else if err != nil && o.log != nil {
- o.log.Debug("organize read local metadata before rename failed", zap.String("path", src), zap.Error(err))
- }
- if match := o.lookupOrganizeAdultMetadata(ctx, src, mediaType, title); match != nil {
- return match
- }
- if o == nil || o.scraper == nil || !o.scraper.AnyEnabled() {
- return nil
- }
- libType := normalizeOrganizeMediaType(mediaType)
- if libType == "" {
- libType = organizeLibraryModelType(mediaType)
- }
- lib := &model.Library{Path: sourceRoot, Type: libType, Enabled: true}
- media := &model.Media{
- Title: title,
- Year: year,
- Path: src,
- SeasonNum: season,
- EpisodeNum: episode,
- }
- for _, candidate := range scrapeQueryCandidates(media, lib) {
- key := organizeMetadataCacheKey(lib.Type, candidate, year)
- if cache != nil {
- if cached, ok := cache[key]; ok {
- if cached != nil {
- return cached
- }
- continue
- }
- }
- match := o.scraper.lookup(ctx, lib, candidate, year)
- if match != nil && strings.TrimSpace(match.Title) != "" {
- if !organizeMetadataMatchTrusted(candidate, year, match) {
- if cache != nil {
- cache[key] = nil
- }
- if o.log != nil {
- o.log.Warn("organize metadata match rejected before rename",
- zap.String("source", src),
- zap.String("query", candidate),
- zap.String("title", match.Title),
- zap.Int("source_year", year),
- zap.Int("match_year", match.Year),
- zap.Int("tmdb_id", match.TMDbID),
- zap.Int("bangumi_id", match.BangumiID),
- zap.String("douban_id", match.DoubanID),
- zap.String("thetvdb_id", match.TheTVDBID))
- }
- continue
- }
- if cache != nil {
- cache[key] = match
- }
- if o.log != nil {
- o.log.Info("organize metadata matched before rename",
- zap.String("source", src),
- zap.String("query", candidate),
- zap.String("title", match.Title),
- zap.Int("year", match.Year),
- zap.Int("tmdb_id", match.TMDbID),
- zap.Int("bangumi_id", match.BangumiID),
- zap.String("douban_id", match.DoubanID),
- zap.String("thetvdb_id", match.TheTVDBID))
- }
- return match
- }
- if cache != nil {
- cache[key] = nil
- }
- }
- return nil
-}
-
-func (o *OrganizerService) lookupOrganizeAdultMetadata(ctx context.Context, src, mediaType, title string) *Match {
- if o == nil || o.scraper == nil || o.scraper.adult == nil || !o.scraper.adult.Enabled() {
- return nil
- }
- isAdult := normalizeOrganizeMediaType(mediaType) == "adult"
- candidates := []string{src, filepath.Base(src), title}
- outCodes := make([]string, 0, len(candidates))
- seen := map[string]struct{}{}
- for _, candidate := range candidates {
- code := normalizeAdultCode(candidate)
- if code == "" {
- continue
- }
- if _, ok := seen[code]; ok {
- continue
- }
- seen[code] = struct{}{}
- outCodes = append(outCodes, code)
- }
- if !isAdult && len(outCodes) == 0 {
- return nil
- }
- for _, code := range outCodes {
- match, err := o.scraper.adult.Search(ctx, code)
- if err != nil {
- if o.log != nil {
- o.log.Debug("organize adult metadata search failed", zap.String("source", src), zap.String("code", code), zap.Error(err))
- }
- continue
- }
- if match != nil && strings.TrimSpace(match.Title) != "" {
- if o.log != nil {
- o.log.Info("organize adult metadata matched before rename",
- zap.String("source", src),
- zap.String("code", code),
- zap.String("title", match.Title))
- }
- return match
- }
- }
- return nil
-}
-
-func organizeMetadataMatchTrusted(query string, sourceYear int, match *Match) bool {
- if match == nil || strings.TrimSpace(match.Title) == "" {
- return false
- }
- if sourceYear > 0 && match.Year > 0 {
- diff := sourceYear - match.Year
- if diff < 0 {
- diff = -diff
- }
- if diff > 1 {
- return false
- }
- }
- return true
-}
-
-func organizeMatchFromLocalMetadata(local *LocalMetadata) *Match {
- if local == nil || strings.TrimSpace(local.Title) == "" {
- return nil
- }
- match := &Match{
- Title: strings.TrimSpace(local.Title),
- OriginalName: strings.TrimSpace(local.OriginalName),
- Overview: local.Overview,
- PosterURL: local.PosterURL,
- BackdropURL: local.BackdropURL,
- Year: local.Year,
- Rating: local.Rating,
- TMDbID: local.TMDbID,
- DoubanID: local.DoubanID,
- TheTVDBID: local.TheTVDBID,
- NSFW: local.NSFW,
- }
- if local.Genres != "" {
- match.Genres = splitNFOList(local.Genres)
- }
- if local.Countries != "" {
- match.Countries = splitNFOList(local.Countries)
- }
- if local.Languages != "" {
- match.Languages = splitNFOList(local.Languages)
- }
- return match
-}
-
-func organizeMetadataCacheKey(mediaType, query string, year int) string {
- return strings.ToLower(strings.TrimSpace(mediaType)) + "|" + fmt.Sprint(year) + "|" + strings.ToLower(strings.TrimSpace(query))
-}
-
-func (o *OrganizerService) smartClassifySourceFile(ctx context.Context, src, sourceRoot, mediaType, title, parsedTitle string, metadataMatch *Match) string {
- if o == nil || !o.isSmartClassifyEnabled(ctx) {
- return ""
- }
- seriesLike := isSeriesLibraryType(mediaType)
- input := mediaClassifyInput{
- MediaType: mediaType,
- Title: strings.Join([]string{title, parsedTitle, filepath.Base(src)}, " "),
- Category: strings.Join(organizeDirectoryCategoryCandidates(src, sourceRoot), " "),
- }
- if metadataMatch != nil {
- input.Title = strings.Join([]string{
- metadataMatch.OriginalName,
- title,
- parsedTitle,
- filepath.Base(src),
- }, " ")
- input.Languages = metadataMatch.Languages
- input.Countries = metadataMatch.Countries
- input.Genres = metadataMatch.Genres
- if metadataMatch.NSFW {
- input.MediaType = "adult"
- }
- }
- if meta, err := ReadLocalMetadata(src, sourceRoot, seriesLike); err == nil && meta != nil && meta.HasNFO {
- input.Title = strings.Join([]string{meta.Title, meta.OriginalName, title, parsedTitle, filepath.Base(src)}, " ")
- input.Languages = parseCommaList(meta.Languages)
- input.Countries = parseCommaList(meta.Countries)
- input.Genres = parseCommaList(meta.Genres)
- if meta.NSFW {
- input.MediaType = "adult"
- }
- }
- return sanitizeFilename(classifyMediaCategory(input, o.categoryMap()))
-}
-
-func (o *OrganizerService) inferOrganizeDirectoryLayout(src, sourceRoot string) organizeDirectoryLayout {
- for _, name := range organizeDirectoryCategoryCandidates(src, sourceRoot) {
- if mediaType, category := o.mediaTypeForDirectoryCategory(name); mediaType != "" && category != "" {
- return organizeDirectoryLayout{MediaType: mediaType, Category: category}
- }
- }
- return organizeDirectoryLayout{}
-}
-
-func organizeDirectoryCategoryCandidates(src, sourceRoot string) []string {
- var out []string
- seen := map[string]struct{}{}
- add := func(value string) {
- value = strings.TrimSpace(value)
- if value == "" || value == "." || value == string(os.PathSeparator) {
- return
- }
- key := strings.ToLower(value)
- if _, ok := seen[key]; ok {
- return
- }
- seen[key] = struct{}{}
- out = append(out, value)
- }
-
- cleanSourceRoot := filepath.Clean(sourceRoot)
- for _, part := range organizePathNameParts(cleanSourceRoot) {
- add(part)
- }
- rel, err := filepath.Rel(cleanSourceRoot, filepath.Clean(src))
- if err != nil || rel == "." || strings.HasPrefix(rel, "..") {
- return out
- }
- dir := filepath.Dir(rel)
- if dir == "." {
- return out
- }
- for _, part := range strings.Split(dir, string(os.PathSeparator)) {
- add(part)
- }
- return out
-}
-
-func organizePathNameParts(path string) []string {
- clean := filepath.Clean(strings.TrimSpace(path))
- if clean == "" || clean == "." {
- return nil
- }
- volume := filepath.VolumeName(clean)
- if volume != "" {
- clean = strings.TrimPrefix(clean, volume)
- }
- clean = strings.Trim(clean, string(os.PathSeparator))
- if clean == "" {
- base := filepath.Base(filepath.Clean(path))
- if base == "." || base == string(os.PathSeparator) {
- return nil
- }
- return []string{base}
- }
- parts := strings.Split(clean, string(os.PathSeparator))
- out := make([]string, 0, len(parts))
- for _, part := range parts {
- part = strings.TrimSpace(part)
- if part != "" && part != "." {
- out = append(out, part)
- }
- }
- return out
-}
-
-func (o *OrganizerService) mediaTypeForDirectoryCategory(name string) (string, string) {
- key := strings.ToLower(strings.TrimSpace(name))
- if key == "" {
- return "", ""
- }
- if hit, ok := o.directoryCategoryTypes()[key]; ok {
- return hit.MediaType, hit.Category
- }
- return "", ""
-}
-
-func (o *OrganizerService) directoryCategoryTypes() map[string]organizeDirectoryLayout {
- categories := o.categoryMap()
- out := map[string]organizeDirectoryLayout{}
- add := func(category, mediaType string) {
- category = strings.TrimSpace(category)
- if category == "" {
- return
- }
- out[strings.ToLower(category)] = organizeDirectoryLayout{
- MediaType: mediaType,
- Category: category,
- }
- }
- addConfigured := func(key, fallback, mediaType string) {
- add(fallback, mediaType)
- add(categoryName(categories, key, fallback), mediaType)
- }
- addConfigured("animation_movie", "动画电影", "movie")
- addConfigured("chinese_movie", "华语电影", "movie")
- addConfigured("jk_movie", "日韩电影", "movie")
- addConfigured("euus_movie", "欧美电影", "movie")
- addConfigured("foreign_movie", "外语电影", "movie")
- addConfigured("domestic_tv", "国产剧", "tv")
- addConfigured("euus_tv", "欧美剧", "tv")
- addConfigured("jk_tv", "日韩剧", "tv")
- addConfigured("cn_anime", "国漫", "anime")
- addConfigured("jp_anime", "日番", "anime")
- addConfigured("variety", "综艺", "variety")
- addConfigured("documentary", "纪录片", "tv")
- addConfigured("children", "儿童", "tv")
- addConfigured("uncategorized_tv", "未分类", "tv")
- addConfigured("adult", "成人", "adult")
- addConfigured("adult_9kg", "9KG", "adult")
- addConfigured("adult_jav", "番号", "adult")
- return out
-}
-
-func (o *OrganizerService) organizeLibraryRootForLayout(ctx context.Context, destRoot, mediaType, category string) (string, bool) {
- if o == nil || o.repo == nil || o.repo.Library == nil {
- return "", false
- }
- libraries, err := o.repo.Library.List(ctx)
- if err != nil {
- if o.log != nil {
- o.log.Debug("list libraries for organize target failed", zap.Error(err))
- }
- return "", false
- }
- destRoot = filepath.Clean(strings.TrimSpace(destRoot))
- mediaType = normalizeOrganizeMediaType(mediaType)
- aliases := o.organizeCategoryAliases(mediaType, category)
-
- bestPath := ""
- bestScore := -1
- bestDepth := -1
- for _, lib := range libraries {
- if !lib.Enabled || strings.TrimSpace(lib.Path) == "" {
- continue
- }
- if _, ok := ParseCloudLibraryMount(lib.Path); ok {
- continue
- }
- if isOrganizeStagingDir(lib.Path) {
- // "手动整理"等暂存库不作为入库目标,避免把媒体留在暂存目录里。
- continue
- }
- if destRoot != "" && destRoot != "." && !pathWithin(lib.Path, destRoot) && !pathWithin(destRoot, lib.Path) {
- continue
- }
- categoryMatch := len(aliases) > 0 && libraryMatchesOrganizeCategory(lib, aliases)
- typeScore := organizeLibraryTypeScore(mediaType, lib.Type)
- if len(aliases) > 0 {
- if !categoryMatch {
- continue
- }
- } else if typeScore <= 0 {
- continue
- }
- score := typeScore
- if categoryMatch {
- score += 20
- }
- depth := pathDepth(lib.Path)
- if score > bestScore || (score == bestScore && depth > bestDepth) {
- bestScore = score
- bestDepth = depth
- bestPath = lib.Path
- }
- }
- if bestPath == "" {
- return "", false
- }
- return filepath.Clean(bestPath), true
-}
-
-func (o *OrganizerService) ensureOrganizeLibraryForRoot(ctx context.Context, root, mediaType, category string) {
- if o == nil || o.repo == nil || o.repo.Library == nil {
- return
- }
- root = filepath.Clean(strings.TrimSpace(root))
- if root == "" || root == "." {
- return
- }
- if _, ok := ParseCloudLibraryMount(root); ok {
- return
- }
- libraries, err := o.repo.Library.List(ctx)
- if err != nil {
- if o.log != nil {
- o.log.Debug("list libraries before organize auto-create failed", zap.Error(err))
- }
- return
- }
- for _, lib := range libraries {
- if !lib.Enabled || strings.TrimSpace(lib.Path) == "" {
- continue
- }
- if _, ok := ParseCloudLibraryMount(lib.Path); ok {
- continue
- }
- if pathWithin(root, lib.Path) {
- return
- }
- }
- name := strings.TrimSpace(category)
- if name == "" {
- name = filepath.Base(root)
- }
- if name == "" || name == "." || name == string(os.PathSeparator) {
- name = organizeLibraryTypeName(mediaType)
- }
- lib := model.Library{
- Name: name,
- Path: root,
- Type: organizeLibraryModelType(mediaType),
- Enabled: true,
- }
- if err := o.repo.Library.Create(ctx, &lib); err != nil {
- if o.log != nil {
- o.log.Warn("organize auto-create library failed",
- zap.String("path", root),
- zap.String("type", lib.Type),
- zap.String("name", lib.Name),
- zap.Error(err))
- }
- return
- }
- if o.log != nil {
- o.log.Info("organize auto-created missing library",
- zap.String("path", root),
- zap.String("type", lib.Type),
- zap.String("name", lib.Name))
- }
-}
-
-func organizeLibraryModelType(mediaType string) string {
- switch normalizeOrganizeMediaType(mediaType) {
- case "tv", "anime", "variety":
- return "tv"
- case "adult", "movie":
- return "movie"
- default:
- return "movie"
- }
-}
-
-func organizeLibraryTypeName(mediaType string) string {
- switch normalizeOrganizeMediaType(mediaType) {
- case "tv":
- return "电视剧"
- case "anime":
- return "动漫"
- case "variety":
- return "综艺"
- case "adult":
- return "成人"
- default:
- return "电影"
- }
-}
-
-func (o *OrganizerService) organizeCategoryAliases(mediaType, category string) map[string]struct{} {
- aliases := map[string]struct{}{}
- add := func(values ...string) {
- for _, value := range values {
- key := normalizeOrganizeCategoryKey(value)
- if key != "" {
- aliases[key] = struct{}{}
- }
- }
- }
- categories := o.categoryMap()
- add(category)
- switch normalizeOrganizeCategoryKey(category) {
- case normalizeOrganizeCategoryKey(categoryName(categories, "jp_anime", "日番")), "日番", "日漫", "日本动漫", "日本動畫", "日本动画":
- add("日番", "日漫", "日本动漫", "日本动画")
- case normalizeOrganizeCategoryKey(categoryName(categories, "cn_anime", "国漫")), "国漫", "国产动漫", "國漫":
- add("国漫", "国产动漫")
- case normalizeOrganizeCategoryKey(categoryName(categories, "domestic_tv", "国产剧")), "国产剧", "国剧", "大陆剧", "国产电视剧":
- add("国产剧", "国剧", "大陆剧", "国产电视剧")
- case normalizeOrganizeCategoryKey(categoryName(categories, "euus_tv", "欧美剧")), "欧美剧", "欧美电视剧":
- add("欧美剧", "欧美电视剧")
- case normalizeOrganizeCategoryKey(categoryName(categories, "jk_tv", "日韩剧")), "日韩剧", "日剧", "韩剧":
- add("日韩剧", "日剧", "韩剧")
- case normalizeOrganizeCategoryKey(categoryName(categories, "variety", "综艺")), "综艺", "真人秀":
- add("综艺", "真人秀")
- case normalizeOrganizeCategoryKey(categoryName(categories, "documentary", "纪录片")), "纪录片", "纪录":
- add("纪录片", "纪录")
- case normalizeOrganizeCategoryKey(categoryName(categories, "children", "儿童")), "儿童", "少儿":
- add("儿童", "少儿")
- case normalizeOrganizeCategoryKey(categoryName(categories, "chinese_movie", "华语电影")), "华语电影", "国产电影", "大陆电影":
- add("华语电影", "国产电影", "大陆电影")
- case normalizeOrganizeCategoryKey(categoryName(categories, "foreign_movie", "外语电影")), "外语电影":
- add("外语电影")
- case normalizeOrganizeCategoryKey(categoryName(categories, "animation_movie", "动画电影")), "动画电影", "动漫电影":
- add("动画电影", "动漫电影")
- case normalizeOrganizeCategoryKey(categoryName(categories, "adult", "成人")), "成人":
- add("成人")
- case normalizeOrganizeCategoryKey(categoryName(categories, "adult_9kg", "9KG")), "9kg":
- add("9KG")
- case normalizeOrganizeCategoryKey(categoryName(categories, "adult_jav", "番号")), "番号", "jav":
- add("番号", "JAV")
- }
- return aliases
-}
-
-func libraryMatchesOrganizeCategory(lib model.Library, aliases map[string]struct{}) bool {
- for _, value := range []string{lib.Name, filepath.Base(filepath.Clean(lib.Path))} {
- if _, ok := aliases[normalizeOrganizeCategoryKey(value)]; ok {
- return true
- }
- }
- return false
-}
-
-func organizeLibraryTypeScore(mediaType, libraryType string) int {
- libraryType = normalizeOrganizeMediaType(libraryType)
- if mediaType == "" || libraryType == "" {
- return 1
- }
- if mediaType == libraryType {
- return 8
- }
- if mediaType == "anime" && libraryType == "tv" {
- return 5
- }
- if mediaType == "variety" && libraryType == "tv" {
- return 5
- }
- return 0
-}
-
-func normalizeOrganizeCategoryKey(value string) string {
- value = strings.ToLower(strings.TrimSpace(value))
- value = strings.ReplaceAll(value, " ", "")
- value = strings.ReplaceAll(value, "_", "")
- value = strings.ReplaceAll(value, "-", "")
- return value
-}
-
-func pathDepth(path string) int {
- path = filepath.Clean(path)
- if path == "." || path == string(os.PathSeparator) {
- return 0
- }
- return len(strings.Split(path, string(os.PathSeparator)))
-}
-
-// existingVersionPaths returns existing destination files that represent the
-// same media, combining two strategies and de-duplicating by path:
-//
-// 1. DB identity: media rows already scanned into the destination root whose
-// title (case-insensitive) + year [or + season/episode] match the source.
-// This is robust to directory case/layout differences.
-// 2. Filesystem: video files inside the computed destination folder (matching
-// the SxxExx tag for episodes). Covers destinations that were not scanned.
-func (o *OrganizerService) existingVersionPaths(ctx context.Context, destRoot, destDir, title, episodeTag string, year, season, episode int) []string {
- return mergeExistingVersionPaths(
- o.existingByIdentity(ctx, destRoot, title, year, season, episode),
- o.existingByFolder(destDir, episodeTag),
- )
-}
-
-func mergeExistingVersionPaths(groups ...[]string) []string {
- seen := map[string]struct{}{}
- var out []string
- add := func(p string) {
- if p == "" {
- return
- }
- c := filepath.Clean(p)
- if _, ok := seen[c]; ok {
- return
- }
- if _, err := os.Stat(c); err != nil {
- return
- }
- seen[c] = struct{}{}
- out = append(out, c)
- }
- for _, group := range groups {
- for _, p := range group {
- add(p)
- }
- }
- return out
-}
-
-func (o *OrganizerService) allExistingPathsInDB(ctx context.Context, paths []string) bool {
- if o == nil || o.repo == nil || o.repo.DB == nil || len(paths) == 0 {
- return false
- }
- cleaned := make([]string, 0, len(paths))
- seen := map[string]struct{}{}
- for _, path := range paths {
- path = filepath.Clean(strings.TrimSpace(path))
- if path == "" || path == "." {
- continue
- }
- if _, ok := seen[path]; ok {
- continue
- }
- seen[path] = struct{}{}
- cleaned = append(cleaned, path)
- }
- if len(cleaned) == 0 {
- return false
- }
- var count int64
- if err := o.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Where("path IN ?", cleaned).
- Count(&count).Error; err != nil {
- return false
- }
- return count == int64(len(cleaned))
-}
-
-// existingByIdentity finds scanned destination media matching the parsed
-// identity (case-insensitive title + year for movies; title + season/episode
-// for episodes), located under destRoot.
-func (o *OrganizerService) existingByIdentity(ctx context.Context, destRoot, title string, year, season, episode int) []string {
- if o.repo == nil || o.repo.DB == nil {
- return nil
- }
- title = strings.TrimSpace(title)
- if title == "" {
- return nil
- }
- q := o.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("deleted_at IS NULL").
- Where("LOWER(title) = ?", strings.ToLower(title))
- if season > 0 || episode > 0 {
- q = q.Where("season_num = ? AND episode_num = ?", season, episode)
- } else if year > 0 {
- q = q.Where("year = ?", year)
- }
- var rows []model.Media
- if err := q.Find(&rows).Error; err != nil {
- return nil
- }
- var out []string
- for _, r := range rows {
- if r.Path != "" && pathWithin(r.Path, destRoot) {
- out = append(out, r.Path)
- }
- }
- return out
-}
-
-func (o *OrganizerService) existingByExternalIdentity(ctx context.Context, destRoot string, match *Match, season, episode int) []string {
- if o.repo == nil || o.repo.DB == nil || match == nil {
- return nil
- }
- var conds []string
- var args []any
- if match.TMDbID > 0 {
- conds = append(conds, "tm_db_id = ?")
- args = append(args, match.TMDbID)
- }
- if match.BangumiID > 0 {
- conds = append(conds, "bangumi_id = ?")
- args = append(args, match.BangumiID)
- }
- if strings.TrimSpace(match.DoubanID) != "" {
- conds = append(conds, "douban_id = ?")
- args = append(args, strings.TrimSpace(match.DoubanID))
- }
- if strings.TrimSpace(match.TheTVDBID) != "" {
- conds = append(conds, "thetvdb_id = ?")
- args = append(args, strings.TrimSpace(match.TheTVDBID))
- }
- if len(conds) == 0 {
- return nil
- }
- q := o.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Where("deleted_at IS NULL").
- Where("("+strings.Join(conds, " OR ")+")", args...)
- if season > 0 || episode > 0 {
- q = q.Where("season_num = ? AND episode_num = ?", season, episode)
- }
- var rows []model.Media
- if err := q.Find(&rows).Error; err != nil {
- return nil
- }
- var out []string
- for _, row := range rows {
- if row.Path != "" && pathWithin(row.Path, destRoot) {
- out = append(out, row.Path)
- }
- }
- return out
-}
-
-// existingByFolder returns video files already present in destDir that
-// represent the same media. For an episode (episodeTag != "") it matches files
-// carrying the same SxxExx tag; for a movie it matches every video file in the
-// movie folder.
-func (o *OrganizerService) existingByFolder(destDir, episodeTag string) []string {
- entries, err := os.ReadDir(destDir)
- if err != nil {
- return nil
- }
- tag := strings.ToLower(episodeTag)
- var out []string
- for _, e := range entries {
- if e.IsDir() {
- continue
- }
- name := e.Name()
- if _, ok := videoExtensions[strings.ToLower(filepath.Ext(name))]; !ok {
- continue
- }
- if tag != "" && !strings.Contains(strings.ToLower(name), tag) {
- continue
- }
- out = append(out, filepath.Join(destDir, name))
- }
- return out
-}
-
-// titleCaseWords upper-cases the first letter of each ASCII word; CJK and other
-// non-ASCII leading characters are left untouched. Roman numerals (ii, iii, iv,
-// …) are fully upper-cased so sequels like "Wandering Earth II" keep their
-// canonical casing instead of becoming "Ii".
-func titleCaseWords(s string) string {
- fields := strings.Fields(s)
- for i, w := range fields {
- if isRomanNumeral(w) {
- fields[i] = strings.ToUpper(w)
- continue
- }
- r := []rune(w)
- if len(r) > 0 && r[0] < 128 {
- r[0] = unicode.ToUpper(r[0])
- fields[i] = string(r)
- }
- }
- return strings.Join(fields, " ")
-}
-
-// sequelNumerals is a conservative whitelist of multi-letter Roman numerals
-// used for movie/series sequels. A whitelist avoids false positives on normal
-// English words that happen to be valid numerals (e.g. "mix", "civ", "mi").
-var sequelNumerals = map[string]struct{}{
- "ii": {}, "iii": {}, "iv": {}, "vi": {}, "vii": {}, "viii": {},
- "ix": {}, "xi": {}, "xii": {}, "xiii": {}, "xiv": {}, "xv": {},
-}
-
-func isRomanNumeral(w string) bool {
- _, ok := sequelNumerals[strings.ToLower(w)]
- return ok
-}
-
-// replaceVersions removes the existing lower-resolution files (and their NFO
-// sidecars + DB rows) and transfers src into dst.
-func (o *OrganizerService) replaceVersions(ctx context.Context, src string, existing []string, dst string, mode TransferMode) error {
- for _, e := range existing {
- if err := os.Remove(e); err != nil && !os.IsNotExist(err) {
- return fmt.Errorf("remove existing %s: %w", e, err)
- }
- if nfo := nfoPath(e); nfo != "" {
- _ = os.Remove(nfo)
- }
- if o.repo != nil && o.repo.DB != nil {
- _ = o.repo.DB.WithContext(ctx).Where("path = ?", e).Delete(&model.Media{}).Error
- }
- }
- if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
- return err
- }
- if err := transferFile(src, dst, mode); err != nil {
- return err
- }
- if err := transferSidecarNFO(src, dst, mode); err != nil {
- o.log.Warn("organize sidecar nfo failed",
- zap.String("from", src), zap.String("to", dst), zap.Error(err))
- }
- return nil
-}
-
-// resolutionArea returns the pixel area (width*height) of a video file for 洗版
-// comparison. It prefers ffprobe; when unavailable it falls back to a
-// resolution token in the filename (2160p/1080p/720p). Returns 0 when the
-// resolution cannot be determined, in which case the caller treats the file as
-// "unknown" and never performs a destructive replace.
-func (o *OrganizerService) resolutionArea(ctx context.Context, path string) int {
- // Prefer a scanned media row's stored dimensions. The destination library
- // is normally scanned with ffprobe, so its files have accurate Width/Height
- // even after organize stripped the resolution token from the filename.
- if o.repo != nil && o.repo.DB != nil {
- var m model.Media
- if err := o.repo.DB.WithContext(ctx).
- Select("width", "height").
- Where("path = ?", path).
- Limit(1).Take(&m).Error; err == nil && m.Width > 0 && m.Height > 0 {
- return m.Width * m.Height
- }
- }
- if o.probe != nil {
- if pr, err := o.probe.Probe(ctx, path); err == nil && pr != nil && pr.Width > 0 && pr.Height > 0 {
- return pr.Width * pr.Height
- }
- }
- switch detectResolutionScore(strings.ToLower(filepath.Base(path))) {
- case 4:
- return 3840 * 2160
- case 3:
- return 1920 * 1080
- case 2:
- return 1280 * 720
- default:
- return 0
- }
-}
diff --git a/internal/service/organizer_directory_classification_test.go b/internal/service/organizer_directory_classification_test.go
new file mode 100644
index 0000000..cd91ab7
--- /dev/null
+++ b/internal/service/organizer_directory_classification_test.go
@@ -0,0 +1,272 @@
+package service
+
+import (
+ "os"
+ "path/filepath"
+ "testing"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestOrganizeDirectoryUsesDownloadCategoryLayout(t *testing.T) {
+ root := t.TempDir()
+ src := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ writeOrgFile(t, filepath.Join(src, "国产剧", "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "kuangbiao-e01")
+ writeOrgFile(t, filepath.Join(src, "华语电影", "流浪地球2.2023.2160p.WEB-DL.H265.mkv"), "wandering-earth-2")
+
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), newOrganizerTestRepo(t))
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: src,
+ DestPath: dest,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ if res.Organized != 2 || res.Replaced != 0 || res.Skipped != 0 {
+ t.Fatalf("expected organized=2 replaced=0 skipped=0, got %+v", res)
+ }
+
+ tv := filepath.Join(dest, "电视剧", "国产剧", "狂飙", "Season 01", "狂飙 - S01E01.mkv")
+ if _, err := os.Stat(tv); err != nil {
+ t.Fatalf("expected TV episode organized at %q: %v", tv, err)
+ }
+ movie := filepath.Join(dest, "电影", "华语电影", "流浪地球2 (2023)", "流浪地球2 (2023).mkv")
+ if _, err := os.Stat(movie); err != nil {
+ t.Fatalf("expected movie organized at %q: %v", movie, err)
+ }
+}
+
+func TestOrganizeDirectoryUsesExplicitCategoryLibraryRoot(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)
+ libraryRoot := filepath.Join(dest, "电视剧", "国产剧")
+ wrongType := model.Library{Name: "国产剧", Path: libraryRoot, Type: "movie", Enabled: true}
+ rightType := model.Library{Name: "国产剧", Path: libraryRoot, Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &wrongType); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &rightType); err != nil {
+ t.Fatal(err)
+ }
+
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: src,
+ DestPath: dest,
+ MediaType: "tv",
+ MediaCategory: "国产剧",
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize explicit category: %v", err)
+ }
+ if res.Organized != 1 || len(res.Items) != 1 {
+ t.Fatalf("result = %+v, want one organized item", res)
+ }
+ if !pathWithin(res.Items[0].Target, libraryRoot) {
+ t.Fatalf("target = %q, want under %q", res.Items[0].Target, libraryRoot)
+ }
+ if pathWithin(res.Items[0].Target, filepath.Join(dest, "电视剧")) && !pathWithin(res.Items[0].Target, libraryRoot) {
+ t.Fatalf("target landed outside category library: %q", res.Items[0].Target)
+ }
+}
+
+func TestOrganizeDirectoryCreatesMissingCategoryLibraryForVisibility(t *testing.T) {
+ root := t.TempDir()
+ srcRoot := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ source := filepath.Join(srcRoot, "Gourd.Brothers.S01E01.2026.1080p.mkv")
+ target := filepath.Join(dest, "电视剧", "未分类", "Gourd Brothers", "Season 01", "Gourd Brothers - S01E01.mkv")
+ writeOrgFile(t, source, "source")
+ writeOrgFile(t, target, "already-there")
+
+ repos := newOrganizerTestRepo(t)
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: srcRoot,
+ DestPath: dest,
+ MediaType: "tv",
+ MediaCategory: "未分类",
+ TransferMode: TransferCopy,
+ AllowReplaceExisting: false,
+ })
+ if err != nil {
+ t.Fatalf("organize missing category: %v", err)
+ }
+ if res.Organized != 0 || res.Skipped != 1 || len(res.Items) != 1 || res.Items[0].Reason != organizeSkipTargetExists {
+ t.Fatalf("result = %+v, want skipped target exists", res)
+ }
+
+ var lib model.Library
+ if err := repos.DB.Where("path = ?", filepath.Join(dest, "电视剧", "未分类")).First(&lib).Error; err != nil {
+ t.Fatalf("missing auto-created category library: %v", err)
+ }
+ if lib.Name != "未分类" || lib.Type != "tv" || !lib.Enabled {
+ t.Fatalf("auto-created library = %+v, want enabled tv 未分类", lib)
+ }
+
+ scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
+ scans := scanner.ScanLibrariesForPath(t.Context(), res.DestPath, "")
+ added := 0
+ for _, scan := range scans {
+ if scan.Error != "" {
+ t.Fatalf("scan failed: %#v", scan)
+ }
+ added += scan.Added
+ }
+ if added != 1 {
+ t.Fatalf("scan added = %d, want 1; scans=%#v", added, scans)
+ }
+}
+
+func TestOrganizeDirectorySmartClassifiesUncategorizedSources(t *testing.T) {
+ root := t.TempDir()
+ src := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ writeOrgFile(t, filepath.Join(src, "流浪地球2.2023.2160p.WEB-DL.mkv"), "cn-movie")
+ writeOrgFile(t, filepath.Join(src, "Dune.2021.2160p.WEB-DL.mkv"), "foreign-movie")
+ writeOrgFile(t, filepath.Join(src, "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "cn-tv")
+ writeOrgFile(t, filepath.Join(src, "The.Last.of.Us.S01E01.2023.1080p.WEB-DL.mkv"), "western-tv")
+
+ repos := newOrganizerTestRepo(t)
+ if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
+ t.Fatal(err)
+ }
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: src,
+ DestPath: dest,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ if res.Organized != 4 {
+ t.Fatalf("organized = %d, want 4; result=%+v", res.Organized, res)
+ }
+
+ for _, want := range []string{
+ filepath.Join(dest, "电影", "华语电影", "流浪地球2 (2023)", "流浪地球2 (2023).mkv"),
+ filepath.Join(dest, "电影", "外语电影", "Dune (2021)", "Dune (2021).mkv"),
+ filepath.Join(dest, "电视剧", "国产剧", "狂飙", "Season 01", "狂飙 - S01E01.mkv"),
+ filepath.Join(dest, "电视剧", "未分类", "The Last Of Us", "Season 01", "The Last Of Us - S01E01.mkv"),
+ } {
+ if _, err := os.Stat(want); err != nil {
+ t.Fatalf("expected smart classified file at %q: %v; items=%+v", want, err, res.Items)
+ }
+ }
+}
+
+func TestOrganizeDirectorySmartClassifiesWithLocalNFO(t *testing.T) {
+ root := t.TempDir()
+ src := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ writeOrgFile(t, filepath.Join(src, "Some.Show.S01E01.2024.1080p.mkv"), "jp-anime")
+ writeOrgFile(t, filepath.Join(src, "tvshow.nfo"), `
+ Some Show
+ Animation
+ JP
+ ja
+`)
+
+ repos := newOrganizerTestRepo(t)
+ if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
+ t.Fatal(err)
+ }
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: src,
+ DestPath: dest,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ if res.Organized != 1 {
+ t.Fatalf("organized = %d, want 1", res.Organized)
+ }
+ want := filepath.Join(dest, "动漫", "日番", "Some Show", "Season 01", "Some Show - S01E01.mkv")
+ if _, err := os.Stat(want); err != nil {
+ t.Fatalf("expected NFO classified episode at %q: %v", want, err)
+ }
+}
+
+func TestOrganizeDirectoryScanAfterRecursesNestedDownloadFolders(t *testing.T) {
+ root := t.TempDir()
+ src := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ writeOrgFile(t, filepath.Join(src, "国产剧", "子目录", "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "kuangbiao-e01")
+ writeOrgFile(t, filepath.Join(src, "华语电影", "更深", "流浪地球2.2023.2160p.WEB-DL.H265.mkv"), "wandering-earth-2")
+
+ repos := newOrganizerTestRepo(t)
+ tvLib := model.Library{Name: "国产剧", Path: filepath.Join(dest, "电视剧", "国产剧"), Type: "tv", Enabled: true}
+ movieLib := model.Library{Name: "华语电影", Path: filepath.Join(dest, "电影", "华语电影"), Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &tvLib); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &movieLib); err != nil {
+ t.Fatal(err)
+ }
+
+ org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
+ res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: src,
+ DestPath: dest,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ if res.Organized != 2 {
+ t.Fatalf("organized = %d, want 2", res.Organized)
+ }
+
+ scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
+ scans := scanner.ScanLibrariesForPath(t.Context(), res.DestPath, "")
+ if len(scans) != 2 {
+ t.Fatalf("scans = %#v, want two matching libraries", scans)
+ }
+ added := 0
+ for _, scan := range scans {
+ if scan.Error != "" {
+ t.Fatalf("scan failed: %#v", scan)
+ }
+ added += scan.Added
+ }
+ if added != 2 {
+ t.Fatalf("scan added = %d, want 2", added)
+ }
+ var count int64
+ if err := repos.DB.Model(&model.Media{}).Count(&count).Error; err != nil {
+ t.Fatal(err)
+ }
+ if count != 2 {
+ t.Fatalf("media rows = %d, want 2", count)
+ }
+}
+
+func TestSelectOrganizeScanTargetsDedupesSamePathByPathType(t *testing.T) {
+ root := t.TempDir()
+ path := filepath.Join(root, "media", "电视剧", "国产剧")
+ libraries := []model.Library{
+ {Name: "国产剧", Path: path, Type: "movie", Enabled: true},
+ {Name: "国产剧", Path: path, Type: "tv", Enabled: true},
+ }
+
+ targets := selectOrganizeScanTargets(libraries, filepath.Join(root, "media"), "")
+ if len(targets) != 1 {
+ t.Fatalf("targets = %#v, want one deduped target", targets)
+ }
+ if targets[0].Type != "tv" {
+ t.Fatalf("target type = %q, want tv", targets[0].Type)
+ }
+}
diff --git a/internal/service/organizer_directory_layout.go b/internal/service/organizer_directory_layout.go
new file mode 100644
index 0000000..36638a2
--- /dev/null
+++ b/internal/service/organizer_directory_layout.go
@@ -0,0 +1,159 @@
+package service
+
+import (
+ "os"
+ "path/filepath"
+ "strings"
+ "unicode"
+)
+
+func (o *OrganizerService) inferOrganizeDirectoryLayout(src, sourceRoot string) organizeDirectoryLayout {
+ for _, name := range organizeDirectoryCategoryCandidates(src, sourceRoot) {
+ if mediaType, category := o.mediaTypeForDirectoryCategory(name); mediaType != "" && category != "" {
+ return organizeDirectoryLayout{MediaType: mediaType, Category: category}
+ }
+ }
+ return organizeDirectoryLayout{}
+}
+
+func organizeDirectoryCategoryCandidates(src, sourceRoot string) []string {
+ var out []string
+ seen := map[string]struct{}{}
+ add := func(value string) {
+ value = strings.TrimSpace(value)
+ if value == "" || value == "." || value == string(os.PathSeparator) {
+ return
+ }
+ key := strings.ToLower(value)
+ if _, ok := seen[key]; ok {
+ return
+ }
+ seen[key] = struct{}{}
+ out = append(out, value)
+ }
+
+ cleanSourceRoot := filepath.Clean(sourceRoot)
+ for _, part := range organizePathNameParts(cleanSourceRoot) {
+ add(part)
+ }
+ rel, err := filepath.Rel(cleanSourceRoot, filepath.Clean(src))
+ if err != nil || rel == "." || strings.HasPrefix(rel, "..") {
+ return out
+ }
+ dir := filepath.Dir(rel)
+ if dir == "." {
+ return out
+ }
+ for _, part := range strings.Split(dir, string(os.PathSeparator)) {
+ add(part)
+ }
+ return out
+}
+
+func organizePathNameParts(path string) []string {
+ clean := filepath.Clean(strings.TrimSpace(path))
+ if clean == "" || clean == "." {
+ return nil
+ }
+ volume := filepath.VolumeName(clean)
+ if volume != "" {
+ clean = strings.TrimPrefix(clean, volume)
+ }
+ clean = strings.Trim(clean, string(os.PathSeparator))
+ if clean == "" {
+ base := filepath.Base(filepath.Clean(path))
+ if base == "." || base == string(os.PathSeparator) {
+ return nil
+ }
+ return []string{base}
+ }
+ parts := strings.Split(clean, string(os.PathSeparator))
+ out := make([]string, 0, len(parts))
+ for _, part := range parts {
+ part = strings.TrimSpace(part)
+ if part != "" && part != "." {
+ out = append(out, part)
+ }
+ }
+ return out
+}
+
+func (o *OrganizerService) mediaTypeForDirectoryCategory(name string) (string, string) {
+ key := strings.ToLower(strings.TrimSpace(name))
+ if key == "" {
+ return "", ""
+ }
+ if hit, ok := o.directoryCategoryTypes()[key]; ok {
+ return hit.MediaType, hit.Category
+ }
+ return "", ""
+}
+
+func (o *OrganizerService) directoryCategoryTypes() map[string]organizeDirectoryLayout {
+ categories := o.categoryMap()
+ out := map[string]organizeDirectoryLayout{}
+ add := func(category, mediaType string) {
+ category = strings.TrimSpace(category)
+ if category == "" {
+ return
+ }
+ out[strings.ToLower(category)] = organizeDirectoryLayout{
+ MediaType: mediaType,
+ Category: category,
+ }
+ }
+ addConfigured := func(key, fallback, mediaType string) {
+ add(fallback, mediaType)
+ add(categoryName(categories, key, fallback), mediaType)
+ }
+ addConfigured("animation_movie", "动画电影", "movie")
+ addConfigured("chinese_movie", "华语电影", "movie")
+ addConfigured("jk_movie", "日韩电影", "movie")
+ addConfigured("euus_movie", "欧美电影", "movie")
+ addConfigured("foreign_movie", "外语电影", "movie")
+ addConfigured("domestic_tv", "国产剧", "tv")
+ addConfigured("euus_tv", "欧美剧", "tv")
+ addConfigured("jk_tv", "日韩剧", "tv")
+ addConfigured("cn_anime", "国漫", "anime")
+ addConfigured("jp_anime", "日番", "anime")
+ addConfigured("variety", "综艺", "variety")
+ addConfigured("documentary", "纪录片", "tv")
+ addConfigured("children", "儿童", "tv")
+ addConfigured("uncategorized_tv", "未分类", "tv")
+ addConfigured("adult", "成人", "adult")
+ addConfigured("adult_9kg", "9KG", "adult")
+ addConfigured("adult_jav", "番号", "adult")
+ return out
+}
+
+// titleCaseWords upper-cases the first letter of each ASCII word; CJK and other
+// non-ASCII leading characters are left untouched. Roman numerals (ii, iii, iv,
+// etc.) are fully upper-cased so sequels keep their canonical casing.
+func titleCaseWords(s string) string {
+ fields := strings.Fields(s)
+ for i, w := range fields {
+ if isRomanNumeral(w) {
+ fields[i] = strings.ToUpper(w)
+ continue
+ }
+ r := []rune(w)
+ if len(r) > 0 && r[0] < 128 {
+ r[0] = unicode.ToUpper(r[0])
+ fields[i] = string(r)
+ }
+ }
+ return strings.Join(fields, " ")
+}
+
+// sequelNumerals is a conservative whitelist of multi-letter Roman numerals
+// used for movie/series sequels. A whitelist avoids false positives on normal
+// English words that happen to be valid numerals (e.g. "mix", "civ", "mi").
+var sequelNumerals = map[string]struct{}{
+ "ii": {}, "iii": {}, "iv": {}, "vi": {}, "vii": {}, "viii": {},
+ "ix": {}, "xi": {}, "xii": {}, "xiii": {}, "xiv": {}, "xv": {},
+}
+
+func isRomanNumeral(w string) bool {
+ _, ok := sequelNumerals[strings.ToLower(w)]
+ return ok
+}
diff --git a/internal/service/organizer_directory_libraries.go b/internal/service/organizer_directory_libraries.go
new file mode 100644
index 0000000..c7687cb
--- /dev/null
+++ b/internal/service/organizer_directory_libraries.go
@@ -0,0 +1,335 @@
+package service
+
+import (
+ "context"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (o *OrganizerService) organizeLibraryRootForLayout(ctx context.Context, destRoot, mediaType, category string) (string, bool) {
+ lib, ok := o.organizeLibraryForLayout(ctx, destRoot, mediaType, category)
+ if !ok {
+ return "", false
+ }
+ return filepath.Clean(lib.Path), true
+}
+
+func (o *OrganizerService) organizeLibraryForLayout(ctx context.Context, destRoot, mediaType, category string) (model.Library, bool) {
+ if o == nil || o.repo == nil || o.repo.Library == nil {
+ return model.Library{}, false
+ }
+ libraries, err := o.repo.Library.List(ctx)
+ if err != nil {
+ if o.log != nil {
+ o.log.Debug("list libraries for organize target failed", zap.Error(err))
+ }
+ return model.Library{}, false
+ }
+ destRoot = filepath.Clean(strings.TrimSpace(destRoot))
+ mediaType = normalizeOrganizeMediaType(mediaType)
+ aliases := o.organizeCategoryAliases(mediaType, category)
+
+ var best model.Library
+ bestScore := -1
+ bestDepth := -1
+ for _, lib := range libraries {
+ if !lib.Enabled || strings.TrimSpace(lib.Path) == "" {
+ continue
+ }
+ if _, ok := ParseCloudLibraryMount(lib.Path); ok {
+ continue
+ }
+ if isOrganizeStagingDir(lib.Path) {
+ // "手动整理"等暂存库不作为入库目标,避免把媒体留在暂存目录里。
+ continue
+ }
+ scopeScore, inScope := o.organizeLibraryTargetScopeScore(lib.Path, destRoot, mediaType, category)
+ if !inScope {
+ continue
+ }
+ categoryMatch := len(aliases) > 0 && libraryMatchesOrganizeCategory(lib, aliases)
+ typeScore := organizeLibraryTypeScore(mediaType, lib.Type)
+ if len(aliases) > 0 {
+ if !categoryMatch {
+ continue
+ }
+ } else if typeScore <= 0 {
+ continue
+ }
+ score := typeScore
+ if categoryMatch {
+ score += 20
+ }
+ score += scopeScore
+ depth := pathDepth(lib.Path)
+ if score > bestScore || (score == bestScore && depth > bestDepth) {
+ bestScore = score
+ bestDepth = depth
+ best = lib
+ }
+ }
+ if strings.TrimSpace(best.Path) == "" {
+ return model.Library{}, false
+ }
+ best.Path = filepath.Clean(best.Path)
+ return best, true
+}
+
+func (o *OrganizerService) organizeLibraryTargetScopeScore(libPath, destRoot, mediaType, category string) (int, bool) {
+ libPath = filepath.Clean(strings.TrimSpace(libPath))
+ destRoot = filepath.Clean(strings.TrimSpace(destRoot))
+ if destRoot == "" || destRoot == "." {
+ return 0, true
+ }
+ if libPath == "" || libPath == "." {
+ return 0, false
+ }
+ if pathWithin(libPath, destRoot) || pathWithin(destRoot, libPath) {
+ return 8, true
+ }
+ if _, destCategory := o.mediaTypeForDirectoryCategory(filepath.Base(destRoot)); destCategory == "" {
+ if root := organizeMediaCollectionRoot(destRoot); root != "" && pathWithin(libPath, root) {
+ return o.organizeLibraryPhysicalRootScore(libPath, root, mediaType, category), true
+ }
+ return 0, false
+ }
+ if strings.EqualFold(filepath.Dir(libPath), filepath.Dir(destRoot)) {
+ return 4, true
+ }
+ if root := organizeMediaCollectionRoot(destRoot); root != "" && pathWithin(libPath, root) {
+ return o.organizeLibraryPhysicalRootScore(libPath, root, mediaType, category), true
+ }
+ return 0, false
+}
+
+func (o *OrganizerService) organizeLibraryPhysicalRootScore(libPath, collectionRoot, mediaType, category string) int {
+ physicalRoot := o.categoryPhysicalRootDir(category)
+ if physicalRoot == "" {
+ physicalRoot = mediaTypeRootDir(mediaType)
+ }
+ if physicalRoot == "" {
+ return 1
+ }
+ if pathHasDirectChild(libPath, collectionRoot, physicalRoot) {
+ return 14
+ }
+ return 1
+}
+
+func organizeMediaCollectionRoot(path string) string {
+ clean := filepath.Clean(strings.TrimSpace(path))
+ if clean == "" || clean == "." {
+ return ""
+ }
+ if isGenericMediaRoot(clean) {
+ return clean
+ }
+ base := filepath.Base(clean)
+ parent := filepath.Dir(clean)
+ if isPhysicalMediaRootDir(base) {
+ return parent
+ }
+ if isPhysicalMediaRootDir(filepath.Base(parent)) {
+ return filepath.Dir(parent)
+ }
+ return ""
+}
+
+func isPhysicalMediaRootDir(name string) bool {
+ switch normalizeOrganizeCategoryKey(name) {
+ case normalizeOrganizeCategoryKey("电影"), normalizeOrganizeCategoryKey("电视剧"), normalizeOrganizeCategoryKey("动漫"), normalizeOrganizeCategoryKey("成人"):
+ return true
+ default:
+ return false
+ }
+}
+
+func pathHasDirectChild(path, root, child string) bool {
+ rel, err := filepath.Rel(filepath.Clean(root), filepath.Clean(path))
+ if err != nil || rel == "." || strings.HasPrefix(rel, "..") {
+ return false
+ }
+ parts := strings.Split(rel, string(os.PathSeparator))
+ if len(parts) == 0 {
+ return false
+ }
+ return strings.EqualFold(parts[0], child)
+}
+
+func (o *OrganizerService) ensureOrganizeLibraryForRoot(ctx context.Context, root, mediaType, category string) {
+ if o == nil || o.repo == nil || o.repo.Library == nil {
+ return
+ }
+ root = filepath.Clean(strings.TrimSpace(root))
+ if root == "" || root == "." {
+ return
+ }
+ if _, ok := ParseCloudLibraryMount(root); ok {
+ return
+ }
+ libraries, err := o.repo.Library.List(ctx)
+ if err != nil {
+ if o.log != nil {
+ o.log.Debug("list libraries before organize auto-create failed", zap.Error(err))
+ }
+ return
+ }
+ for _, lib := range libraries {
+ if !lib.Enabled || strings.TrimSpace(lib.Path) == "" {
+ continue
+ }
+ if _, ok := ParseCloudLibraryMount(lib.Path); ok {
+ continue
+ }
+ if pathWithin(root, lib.Path) {
+ return
+ }
+ }
+ name := strings.TrimSpace(category)
+ if name == "" {
+ name = filepath.Base(root)
+ }
+ if name == "" || name == "." || name == string(os.PathSeparator) {
+ name = organizeLibraryTypeName(mediaType)
+ }
+ lib := model.Library{
+ Name: name,
+ Path: root,
+ Type: organizeLibraryModelType(mediaType),
+ Enabled: true,
+ }
+ if err := o.repo.Library.Create(ctx, &lib); err != nil {
+ if o.log != nil {
+ o.log.Warn("organize auto-create library failed",
+ zap.String("path", root),
+ zap.String("type", lib.Type),
+ zap.String("name", lib.Name),
+ zap.Error(err))
+ }
+ return
+ }
+ if o.log != nil {
+ o.log.Info("organize auto-created missing library",
+ zap.String("path", root),
+ zap.String("type", lib.Type),
+ zap.String("name", lib.Name))
+ }
+}
+
+func organizeLibraryModelType(mediaType string) string {
+ switch normalizeOrganizeMediaType(mediaType) {
+ case "tv", "anime", "variety":
+ return "tv"
+ case "adult", "movie":
+ return "movie"
+ default:
+ return "movie"
+ }
+}
+
+func organizeLibraryTypeName(mediaType string) string {
+ switch normalizeOrganizeMediaType(mediaType) {
+ case "tv":
+ return "电视剧"
+ case "anime":
+ return "动漫"
+ case "variety":
+ return "综艺"
+ case "adult":
+ return "成人"
+ default:
+ return "电影"
+ }
+}
+
+func (o *OrganizerService) organizeCategoryAliases(mediaType, category string) map[string]struct{} {
+ aliases := map[string]struct{}{}
+ add := func(values ...string) {
+ for _, value := range values {
+ key := normalizeOrganizeCategoryKey(value)
+ if key != "" {
+ aliases[key] = struct{}{}
+ }
+ }
+ }
+ categories := o.categoryMap()
+ add(category)
+ switch normalizeOrganizeCategoryKey(category) {
+ case normalizeOrganizeCategoryKey(categoryName(categories, "jp_anime", "日番")), "日番", "日漫", "日本动漫", "日本動畫", "日本动画":
+ add("日番", "日漫", "日本动漫", "日本动画")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "cn_anime", "国漫")), "国漫", "国产动漫", "國漫":
+ add("国漫", "国产动漫")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "domestic_tv", "国产剧")), "国产剧", "国剧", "大陆剧", "国产电视剧":
+ add("国产剧", "国剧", "大陆剧", "国产电视剧")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "euus_tv", "欧美剧")), "欧美剧", "欧美电视剧":
+ add("欧美剧", "欧美电视剧")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "jk_tv", "日韩剧")), "日韩剧", "日剧", "韩剧":
+ add("日韩剧", "日剧", "韩剧")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "variety", "综艺")), "综艺", "真人秀":
+ add("综艺", "真人秀")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "documentary", "纪录片")), "纪录片", "纪录":
+ add("纪录片", "纪录")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "children", "儿童")), "儿童", "少儿":
+ add("儿童", "少儿")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "chinese_movie", "华语电影")), "华语电影", "国产电影", "大陆电影":
+ add("华语电影", "国产电影", "大陆电影")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "foreign_movie", "外语电影")), "外语电影":
+ add("外语电影")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "animation_movie", "动画电影")), "动画电影", "动漫电影":
+ add("动画电影", "动漫电影")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "adult", "成人")), "成人":
+ add("成人")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "adult_9kg", "9KG")), "9kg":
+ add("9KG")
+ case normalizeOrganizeCategoryKey(categoryName(categories, "adult_jav", "番号")), "番号", "jav":
+ add("番号", "JAV")
+ }
+ return aliases
+}
+
+func libraryMatchesOrganizeCategory(lib model.Library, aliases map[string]struct{}) bool {
+ for _, value := range []string{lib.Name, filepath.Base(filepath.Clean(lib.Path))} {
+ if _, ok := aliases[normalizeOrganizeCategoryKey(value)]; ok {
+ return true
+ }
+ }
+ return false
+}
+
+func organizeLibraryTypeScore(mediaType, libraryType string) int {
+ libraryType = normalizeOrganizeMediaType(libraryType)
+ if mediaType == "" || libraryType == "" {
+ return 1
+ }
+ if mediaType == libraryType {
+ return 8
+ }
+ if mediaType == "anime" && libraryType == "tv" {
+ return 5
+ }
+ if mediaType == "variety" && libraryType == "tv" {
+ return 5
+ }
+ return 0
+}
+
+func normalizeOrganizeCategoryKey(value string) string {
+ value = strings.ToLower(strings.TrimSpace(value))
+ value = strings.ReplaceAll(value, " ", "")
+ value = strings.ReplaceAll(value, "_", "")
+ value = strings.ReplaceAll(value, "-", "")
+ return value
+}
+
+func pathDepth(path string) int {
+ path = filepath.Clean(path)
+ if path == "." || path == string(os.PathSeparator) {
+ return 0
+ }
+ return len(strings.Split(path, string(os.PathSeparator)))
+}
diff --git a/internal/service/organizer_directory_metadata.go b/internal/service/organizer_directory_metadata.go
new file mode 100644
index 0000000..2c3a341
--- /dev/null
+++ b/internal/service/organizer_directory_metadata.go
@@ -0,0 +1,274 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (o *OrganizerService) lookupOrganizeMetadata(ctx context.Context, src, sourceRoot, mediaType, title string, year, season, episode int, cache map[string]*Match) *Match {
+ seriesLike := isSeriesLibraryType(mediaType) || season > 0 || episode > 0
+ if local, err := ReadLocalMetadata(src, sourceRoot, seriesLike); err == nil && local != nil {
+ if match := organizeMatchFromLocalMetadata(local); match != nil {
+ return match
+ }
+ } else if err != nil && o.log != nil {
+ o.log.Debug("organize read local metadata before rename failed", zap.String("path", src), zap.Error(err))
+ }
+ if match := o.lookupOrganizeAdultMetadata(ctx, src, mediaType, title); match != nil {
+ return match
+ }
+ if o == nil || o.scraper == nil || !o.scraper.AnyEnabled() {
+ return nil
+ }
+ libType := normalizeOrganizeMediaType(mediaType)
+ if libType == "" {
+ libType = organizeLibraryModelType(mediaType)
+ }
+ lib := &model.Library{Path: sourceRoot, Type: libType, Enabled: true}
+ media := &model.Media{
+ Title: title,
+ Year: year,
+ Path: src,
+ SeasonNum: season,
+ EpisodeNum: episode,
+ }
+ for _, candidate := range scrapeQueryCandidates(media, lib) {
+ key := organizeMetadataCacheKey(lib.Type, candidate, year)
+ if cache != nil {
+ if cached, ok := cache[key]; ok {
+ if cached != nil {
+ return cached
+ }
+ continue
+ }
+ }
+ match := o.scraper.lookup(ctx, lib, media, candidate, year)
+ if match != nil && strings.TrimSpace(match.Title) != "" {
+ if !organizeMetadataMatchTrusted(candidate, year, match) {
+ if cache != nil {
+ cache[key] = nil
+ }
+ if o.log != nil {
+ o.log.Warn("organize metadata match rejected before rename",
+ zap.String("source", src),
+ zap.String("query", candidate),
+ zap.String("title", match.Title),
+ zap.Int("source_year", year),
+ zap.Int("match_year", match.Year),
+ zap.Int("tmdb_id", match.TMDbID),
+ zap.Int("bangumi_id", match.BangumiID),
+ zap.String("douban_id", match.DoubanID),
+ zap.String("thetvdb_id", match.TheTVDBID))
+ }
+ continue
+ }
+ if cache != nil {
+ cache[key] = match
+ }
+ if o.log != nil {
+ o.log.Info("organize metadata matched before rename",
+ zap.String("source", src),
+ zap.String("query", candidate),
+ zap.String("title", match.Title),
+ zap.Int("year", match.Year),
+ zap.Int("tmdb_id", match.TMDbID),
+ zap.Int("bangumi_id", match.BangumiID),
+ zap.String("douban_id", match.DoubanID),
+ zap.String("thetvdb_id", match.TheTVDBID))
+ }
+ return match
+ }
+ if cache != nil {
+ cache[key] = nil
+ }
+ }
+ return nil
+}
+
+func (o *OrganizerService) lookupOrganizeAdultMetadata(ctx context.Context, src, mediaType, title string) *Match {
+ if o == nil || o.scraper == nil || o.scraper.adult == nil || !o.scraper.adult.Enabled() {
+ return nil
+ }
+ isAdult := normalizeOrganizeMediaType(mediaType) == "adult"
+ candidates := []string{src, filepath.Base(src), title}
+ outCodes := make([]string, 0, len(candidates))
+ seen := map[string]struct{}{}
+ for _, candidate := range candidates {
+ code := normalizeAdultCode(candidate)
+ if code == "" {
+ continue
+ }
+ if _, ok := seen[code]; ok {
+ continue
+ }
+ seen[code] = struct{}{}
+ outCodes = append(outCodes, code)
+ }
+ if !isAdult && len(outCodes) == 0 {
+ return nil
+ }
+ for _, code := range outCodes {
+ match, err := o.scraper.adult.Search(ctx, code)
+ if err != nil {
+ if o.log != nil {
+ o.log.Debug("organize adult metadata search failed", zap.String("source", src), zap.String("code", code), zap.Error(err))
+ }
+ continue
+ }
+ if match != nil && strings.TrimSpace(match.Title) != "" {
+ if o.log != nil {
+ o.log.Info("organize adult metadata matched before rename",
+ zap.String("source", src),
+ zap.String("code", code),
+ zap.String("title", match.Title))
+ }
+ return match
+ }
+ }
+ return nil
+}
+
+func organizeMetadataMatchTrusted(query string, sourceYear int, match *Match) bool {
+ if match == nil || strings.TrimSpace(match.Title) == "" {
+ return false
+ }
+ if sourceYear > 0 && match.Year > 0 {
+ diff := sourceYear - match.Year
+ if diff < 0 {
+ diff = -diff
+ }
+ if diff > 1 {
+ return false
+ }
+ }
+ return true
+}
+
+func organizeMatchFromLocalMetadata(local *LocalMetadata) *Match {
+ if local == nil || strings.TrimSpace(local.Title) == "" {
+ return nil
+ }
+ match := &Match{
+ Title: strings.TrimSpace(local.Title),
+ OriginalName: strings.TrimSpace(local.OriginalName),
+ Overview: local.Overview,
+ PosterURL: local.PosterURL,
+ BackdropURL: local.BackdropURL,
+ Year: local.Year,
+ Rating: local.Rating,
+ TMDbID: local.TMDbID,
+ DoubanID: local.DoubanID,
+ TheTVDBID: local.TheTVDBID,
+ NSFW: local.NSFW,
+ }
+ if local.Genres != "" {
+ match.Genres = splitNFOList(local.Genres)
+ }
+ if local.Countries != "" {
+ match.Countries = splitNFOList(local.Countries)
+ }
+ if local.Languages != "" {
+ match.Languages = splitNFOList(local.Languages)
+ }
+ return match
+}
+
+func (o *OrganizerService) lookupOrganizeSourceMedia(ctx context.Context, path string) *model.Media {
+ if o == nil || o.repo == nil || o.repo.DB == nil {
+ return nil
+ }
+ path = filepath.Clean(strings.TrimSpace(path))
+ if path == "" || path == "." {
+ return nil
+ }
+ var media model.Media
+ if err := o.repo.DB.WithContext(ctx).
+ Where("path = ? AND deleted_at IS NULL", path).
+ Limit(1).
+ Take(&media).Error; err != nil {
+ return nil
+ }
+ return &media
+}
+
+func organizeMatchFromMedia(media *model.Media) *Match {
+ if media == nil || strings.TrimSpace(media.Title) == "" {
+ return nil
+ }
+ return &Match{
+ TMDbID: media.TMDbID,
+ BangumiID: media.BangumiID,
+ DoubanID: strings.TrimSpace(media.DoubanID),
+ TheTVDBID: strings.TrimSpace(media.TheTVDBID),
+ Title: strings.TrimSpace(media.Title),
+ OriginalName: strings.TrimSpace(media.OriginalName),
+ Overview: media.Overview,
+ PosterURL: media.PosterURL,
+ BackdropURL: media.BackdropURL,
+ Year: media.Year,
+ Rating: media.Rating,
+ Languages: parseCommaList(media.Languages),
+ Countries: parseCommaList(media.Countries),
+ Genres: parseCommaList(media.Genres),
+ NSFW: media.NSFW,
+ }
+}
+
+func applyOrganizeMetadataMatch(match *Match, title, parsedTitle *string, year *int) {
+ if match == nil {
+ return
+ }
+ if matchedTitle := sanitizeFilename(strings.TrimSpace(match.Title)); matchedTitle != "" {
+ *title = matchedTitle
+ *parsedTitle = strings.TrimSpace(match.Title)
+ }
+ if match.Year > 0 {
+ *year = match.Year
+ }
+}
+
+func organizeMetadataCacheKey(mediaType, query string, year int) string {
+ return strings.ToLower(strings.TrimSpace(mediaType)) + "|" + fmt.Sprint(year) + "|" + strings.ToLower(strings.TrimSpace(query))
+}
+
+func (o *OrganizerService) smartClassifySourceFile(ctx context.Context, src, sourceRoot, mediaType, title, parsedTitle string, metadataMatch *Match) string {
+ if o == nil || !o.isSmartClassifyEnabled(ctx) {
+ return ""
+ }
+ seriesLike := isSeriesLibraryType(mediaType)
+ input := mediaClassifyInput{
+ MediaType: mediaType,
+ Title: strings.Join([]string{title, parsedTitle, filepath.Base(src)}, " "),
+ Category: strings.Join(organizeDirectoryCategoryCandidates(src, sourceRoot), " "),
+ }
+ if metadataMatch != nil {
+ input.Title = strings.Join([]string{
+ metadataMatch.OriginalName,
+ title,
+ parsedTitle,
+ filepath.Base(src),
+ }, " ")
+ input.Languages = metadataMatch.Languages
+ input.Countries = metadataMatch.Countries
+ input.Genres = metadataMatch.Genres
+ if metadataMatch.NSFW {
+ input.MediaType = "adult"
+ }
+ }
+ if meta, err := ReadLocalMetadata(src, sourceRoot, seriesLike); err == nil && meta != nil && meta.HasNFO {
+ input.Title = strings.Join([]string{meta.Title, meta.OriginalName, title, parsedTitle, filepath.Base(src)}, " ")
+ input.Languages = parseCommaList(meta.Languages)
+ input.Countries = parseCommaList(meta.Countries)
+ input.Genres = parseCommaList(meta.Genres)
+ if meta.NSFW {
+ input.MediaType = "adult"
+ }
+ }
+ return sanitizeFilename(classifyMediaCategory(input, o.categoryMap()))
+}
diff --git a/internal/service/organizer_directory_reclassify.go b/internal/service/organizer_directory_reclassify.go
new file mode 100644
index 0000000..f7938b7
--- /dev/null
+++ b/internal/service/organizer_directory_reclassify.go
@@ -0,0 +1,256 @@
+package service
+
+import (
+ "context"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type organizeExistingReclassifyRequest struct {
+ Source string
+ Target string
+ DestRoot string
+ TargetLibraryID string
+ Existing []string
+ DryRun bool
+ MediaType string
+ Category string
+ Title string
+ Year int
+ Season int
+ Episode int
+ Result *OrganizeResult
+}
+
+func (o *OrganizerService) reclassifyExistingMedia(ctx context.Context, req organizeExistingReclassifyRequest) (bool, error) {
+ if req.Result == nil || len(req.Existing) == 0 {
+ return false, nil
+ }
+ if strings.TrimSpace(req.Category) == "" {
+ return false, nil
+ }
+ target := filepath.Clean(strings.TrimSpace(req.Target))
+ if target == "" || target == "." {
+ return false, nil
+ }
+ candidates := reclassifyExistingCandidates(req.Existing, target, req.DestRoot)
+ if len(candidates) == 0 {
+ return false, nil
+ }
+ if organizeFileExists(target) {
+ cleaned, err := o.cleanupReclassifiedDuplicates(ctx, req, target, candidates)
+ return cleaned > 0, err
+ }
+ if len(candidates) != 1 || o.mediaPathExists(ctx, target) {
+ return false, nil
+ }
+ oldPath := candidates[0]
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: oldPath, Target: target, Action: "reclassify",
+ MediaType: req.MediaType, Category: req.Category, Title: req.Title,
+ Reason: "metadata category changed",
+ })
+ if req.DryRun {
+ req.Result.Reclassified++
+ return true, nil
+ }
+ if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
+ return false, err
+ }
+ if err := moveFile(oldPath, target); err != nil {
+ return false, err
+ }
+ if err := moveSidecarNFO(oldPath, target); err != nil && o != nil && o.log != nil {
+ o.log.Warn("organize reclassify sidecar nfo failed",
+ zap.String("from", nfoPath(oldPath)),
+ zap.String("to", nfoPath(target)),
+ zap.Error(err))
+ }
+ if err := o.updateReclassifiedMediaRow(ctx, oldPath, target, req); err != nil {
+ return false, err
+ }
+ cleanupEmptyMediaDirs(filepath.Dir(oldPath), req.DestRoot)
+ if o != nil && o.log != nil {
+ o.log.Info("organize reclassified existing media",
+ zap.String("from", oldPath),
+ zap.String("to", target),
+ zap.String("category", req.Category),
+ zap.String("media_type", req.MediaType))
+ }
+ req.Result.Reclassified++
+ return true, nil
+}
+
+func reclassifyExistingCandidates(existing []string, target, destRoot string) []string {
+ target = filepath.Clean(target)
+ destRoot = filepath.Clean(strings.TrimSpace(destRoot))
+ seen := map[string]struct{}{}
+ out := make([]string, 0, len(existing))
+ for _, path := range existing {
+ path = filepath.Clean(strings.TrimSpace(path))
+ if path == "" || path == "." || strings.EqualFold(path, target) {
+ continue
+ }
+ if destRoot != "" && destRoot != "." && !pathWithin(path, destRoot) {
+ continue
+ }
+ if !organizeFileExists(path) {
+ continue
+ }
+ key := strings.ToLower(path)
+ if _, ok := seen[key]; ok {
+ continue
+ }
+ seen[key] = struct{}{}
+ out = append(out, path)
+ }
+ return out
+}
+
+func (o *OrganizerService) cleanupReclassifiedDuplicates(ctx context.Context, req organizeExistingReclassifyRequest, target string, candidates []string) (int, error) {
+ cleaned := 0
+ for _, oldPath := range candidates {
+ if !safeToRemoveReclassifiedDuplicate(oldPath, target) {
+ if o != nil && o.log != nil {
+ o.log.Warn("organize kept duplicate with different size during reclassify",
+ zap.String("path", oldPath),
+ zap.String("target", target))
+ }
+ continue
+ }
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: oldPath, Target: target, Action: "cleanup",
+ MediaType: req.MediaType, Category: req.Category, Title: req.Title,
+ Reason: "duplicate after metadata category changed",
+ })
+ if req.DryRun {
+ cleaned++
+ continue
+ }
+ if err := removeMediaAndNFO(oldPath); err != nil {
+ return cleaned, err
+ }
+ o.deleteMediaRowForPath(ctx, oldPath)
+ cleanupEmptyMediaDirs(filepath.Dir(oldPath), req.DestRoot)
+ if o != nil && o.log != nil {
+ o.log.Info("organize cleaned duplicate after reclassify",
+ zap.String("path", oldPath),
+ zap.String("target", target),
+ zap.String("category", req.Category),
+ zap.String("media_type", req.MediaType))
+ }
+ cleaned++
+ }
+ req.Result.Reclassified += cleaned
+ return cleaned, nil
+}
+
+func safeToRemoveReclassifiedDuplicate(path, target string) bool {
+ info, err := os.Stat(path)
+ if err != nil {
+ return false
+ }
+ targetInfo, err := os.Stat(target)
+ if err != nil {
+ return false
+ }
+ if os.SameFile(info, targetInfo) {
+ return true
+ }
+ return info.Size() > 0 && info.Size() == targetInfo.Size()
+}
+
+func organizeFileExists(path string) bool {
+ _, err := os.Stat(path)
+ return err == nil
+}
+
+func moveSidecarNFO(oldMedia, newMedia string) error {
+ oldNFO := nfoPath(oldMedia)
+ newNFO := nfoPath(newMedia)
+ if oldNFO == newNFO || !organizeFileExists(oldNFO) || organizeFileExists(newNFO) {
+ return nil
+ }
+ if err := os.MkdirAll(filepath.Dir(newNFO), 0o755); err != nil { // #nosec G301 -- sidecar media directories must remain readable by NAS/player users.
+ return err
+ }
+ return moveFile(oldNFO, newNFO)
+}
+
+func removeMediaAndNFO(path string) error {
+ if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
+ return err
+ }
+ if nfo := nfoPath(path); nfo != "" {
+ if err := os.Remove(nfo); err != nil && !os.IsNotExist(err) {
+ return err
+ }
+ }
+ return nil
+}
+
+func (o *OrganizerService) updateReclassifiedMediaRow(ctx context.Context, oldPath, newPath string, req organizeExistingReclassifyRequest) error {
+ if o == nil || o.repo == nil || o.repo.DB == nil {
+ return nil
+ }
+ updates := map[string]any{
+ "path": newPath,
+ }
+ if strings.TrimSpace(req.TargetLibraryID) != "" {
+ updates["library_id"] = strings.TrimSpace(req.TargetLibraryID)
+ }
+ if strings.TrimSpace(req.Title) != "" {
+ updates["title"] = strings.TrimSpace(req.Title)
+ }
+ if req.Year > 0 {
+ updates["year"] = req.Year
+ }
+ if req.Season > 0 {
+ updates["season_num"] = req.Season
+ }
+ if req.Episode > 0 {
+ updates["episode_num"] = req.Episode
+ }
+ return o.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("path = ?", oldPath).Updates(updates).Error
+}
+
+func (o *OrganizerService) deleteMediaRowForPath(ctx context.Context, path string) {
+ if o == nil || o.repo == nil || o.repo.DB == nil {
+ return
+ }
+ _ = o.repo.DB.WithContext(ctx).Where("path = ?", path).Delete(&model.Media{}).Error
+}
+
+func (o *OrganizerService) mediaPathExists(ctx context.Context, path string) bool {
+ if o == nil || o.repo == nil || o.repo.DB == nil {
+ return false
+ }
+ var count int64
+ if err := o.repo.DB.WithContext(ctx).Unscoped().Model(&model.Media{}).Where("path = ?", path).Count(&count).Error; err != nil {
+ return false
+ }
+ return count > 0
+}
+
+func cleanupEmptyMediaDirs(startDir, stopRoot string) {
+ dir := filepath.Clean(strings.TrimSpace(startDir))
+ stopRoot = filepath.Clean(strings.TrimSpace(stopRoot))
+ for dir != "" && dir != "." {
+ if stopRoot != "" && stopRoot != "." && (!pathWithin(dir, stopRoot) || strings.EqualFold(dir, stopRoot)) {
+ return
+ }
+ if err := os.Remove(dir); err != nil {
+ return
+ }
+ parent := filepath.Dir(dir)
+ if parent == dir {
+ return
+ }
+ dir = parent
+ }
+}
diff --git a/internal/service/organizer_directory_source.go b/internal/service/organizer_directory_source.go
new file mode 100644
index 0000000..b293f6d
--- /dev/null
+++ b/internal/service/organizer_directory_source.go
@@ -0,0 +1,338 @@
+package service
+
+import (
+ "context"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+)
+
+type organizeDirectoryLayout struct {
+ MediaType string
+ Category string
+}
+
+type organizeSourceFileRequest struct {
+ Source string
+ SourceRoot string
+ DestRoot string
+ Mode TransferMode
+ MediaTypeOverride string
+ MediaCategoryOverride string
+ DryRun bool
+ AllowReplaceExisting bool
+ MetadataCache map[string]*Match
+ Result *OrganizeResult
+}
+
+// organizeSourceFile organizes a single video file from the source directory
+// into destRoot, applying dedup + 洗版.
+func (o *OrganizerService) organizeSourceFile(ctx context.Context, req organizeSourceFileRequest) error {
+ src := req.Source
+ ext := filepath.Ext(src)
+ season, episode := ParseEpisode(src)
+ title, year := CleanQuery(src)
+ if organizeWeakFileTitle(title) {
+ if folderTitle, folderYear := organizeTitleFromParentFolder(src, req.SourceRoot, season > 0 || episode > 0); folderTitle != "" {
+ title = folderTitle
+ if year <= 0 {
+ year = folderYear
+ }
+ } else {
+ title = strings.TrimSuffix(filepath.Base(src), ext)
+ }
+ }
+ parsedTitle := title
+ sourceMedia := o.lookupOrganizeSourceMedia(ctx, src)
+ title = sanitizeFilename(titleCaseWords(title))
+ if sourceMedia != nil {
+ if mediaTitle := sanitizeFilename(strings.TrimSpace(sourceMedia.Title)); mediaTitle != "" {
+ title = mediaTitle
+ parsedTitle = strings.TrimSpace(sourceMedia.Title)
+ }
+ if sourceMedia.Year > 0 {
+ year = sourceMedia.Year
+ }
+ if sourceMedia.SeasonNum > 0 {
+ season = sourceMedia.SeasonNum
+ }
+ if sourceMedia.EpisodeNum > 0 {
+ episode = sourceMedia.EpisodeNum
+ }
+ }
+ if title == "" {
+ title = "Unknown"
+ }
+ pathLayout := o.inferOrganizeDirectoryLayout(src, req.SourceRoot)
+ layout := pathLayout
+ forcedType := normalizeOrganizeMediaType(req.MediaTypeOverride)
+ inferredType := o.inferMediaTypeForSourceFile(src, title, season, episode)
+ if forcedType != "" {
+ if layout.Category != "" && layout.MediaType != "" && layout.MediaType != forcedType {
+ layout.Category = ""
+ }
+ layout.MediaType = forcedType
+ } else if inferredType != "" {
+ if inferredType == "tv" && layout.MediaType == "movie" {
+ layout = organizeDirectoryLayout{MediaType: inferredType}
+ } else if layout.MediaType == "" {
+ layout.MediaType = inferredType
+ }
+ }
+ var metadataMatch *Match
+ if match := o.lookupOrganizeMetadata(ctx, src, req.SourceRoot, layout.MediaType, title, year, season, episode, req.MetadataCache); match != nil {
+ metadataMatch = match
+ applyOrganizeMetadataMatch(metadataMatch, &title, &parsedTitle, &year)
+ } else if sourceMedia != nil {
+ metadataMatch = organizeMatchFromMedia(sourceMedia)
+ applyOrganizeMetadataMatch(metadataMatch, &title, &parsedTitle, &year)
+ }
+ if category := strings.TrimSpace(req.MediaCategoryOverride); category != "" {
+ layout.Category = sanitizeFilename(category)
+ } else if category := o.smartClassifySourceFile(ctx, src, req.SourceRoot, layout.MediaType, title, parsedTitle, metadataMatch); category != "" {
+ layout.Category = category
+ }
+ if forcedType == "" {
+ if impliedType, normalizedCategory := o.mediaTypeForDirectoryCategory(layout.Category); impliedType != "" {
+ layout.Category = normalizedCategory
+ if layout.MediaType == "" || layout.MediaType == "tv" || layout.MediaType == "anime" || pathLayout.Category != layout.Category {
+ layout.MediaType = impliedType
+ }
+ }
+ }
+ targetLibrary, matchedLibrary := o.organizeLibraryForLayout(ctx, req.DestRoot, layout.MediaType, layout.Category)
+ layoutRoot := targetLibrary.Path
+ targetLibraryID := targetLibrary.ID
+ if !matchedLibrary && layout.MediaType != "" {
+ layoutRoot = o.organizeRoot(req.DestRoot, layout.MediaType, layout.Category)
+ }
+ if !matchedLibrary && layout.Category != "" {
+ layoutRoot = categoryRoot(layoutRoot, sanitizeFilename(layout.Category))
+ }
+ if !matchedLibrary && !req.DryRun {
+ o.ensureOrganizeLibraryForRoot(ctx, layoutRoot, layout.MediaType, layout.Category)
+ }
+
+ isSeries := season > 0 || episode > 0
+ if layout.MediaType != "" {
+ isSeries = isSeriesLibraryType(layout.MediaType) && (season > 0 || episode > 0)
+ }
+ target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
+ Root: layoutRoot,
+ MediaType: layout.MediaType,
+ Category: layout.Category,
+ Title: title,
+ Source: src,
+ Ext: ext,
+ Year: year,
+ Season: season,
+ Episode: episode,
+ Series: isSeries,
+ })
+ if err != nil {
+ return err
+ }
+ destDir, dst, episodeTag := target.Dir, target.Path, target.EpisodeTag
+ if filepath.Clean(src) == filepath.Clean(dst) {
+ req.Result.Skipped++
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: src, Target: dst, Action: "skip", Reason: organizeSkipAlreadyOrganized,
+ MediaType: layout.MediaType, Category: layout.Category, Title: title,
+ })
+ return nil
+ }
+
+ externalExisting := o.existingByExternalIdentity(ctx, req.DestRoot, metadataMatch, season, episode)
+ identityExisting := o.existingByIdentity(ctx, req.DestRoot, parsedTitle, year, season, episode)
+ folderExisting := o.existingByFolder(destDir, episodeTag)
+ existing := mergeExistingVersionPaths(externalExisting, identityExisting, folderExisting)
+ if len(existing) > 0 {
+ reclassified, err := o.reclassifyExistingMedia(ctx, organizeExistingReclassifyRequest{
+ Source: src,
+ Target: dst,
+ DestRoot: req.DestRoot,
+ TargetLibraryID: targetLibraryID,
+ Existing: existing,
+ DryRun: req.DryRun,
+ MediaType: layout.MediaType,
+ Category: layout.Category,
+ Title: title,
+ Year: year,
+ Season: season,
+ Episode: episode,
+ Result: req.Result,
+ })
+ if err != nil {
+ return err
+ }
+ if reclassified {
+ return nil
+ }
+ srcArea := o.resolutionArea(ctx, src)
+ bestArea := 0
+ for _, e := range existing {
+ if a := o.resolutionArea(ctx, e); a > bestArea {
+ bestArea = a
+ }
+ }
+ if req.AllowReplaceExisting && srcArea > 0 && bestArea > 0 && srcArea > bestArea {
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: src, Target: dst, Action: "replace", Reason: "higher resolution",
+ MediaType: layout.MediaType, Category: layout.Category, Title: title,
+ })
+ if req.DryRun {
+ req.Result.Replaced++
+ return nil
+ }
+ if err := o.replaceVersions(ctx, src, existing, dst, req.Mode); err != nil {
+ return err
+ }
+ o.log.Info("organize replaced lower-resolution media",
+ zap.String("from", src),
+ zap.String("to", dst),
+ zap.Int("src_area", srcArea),
+ zap.Int("existing_area", bestArea),
+ )
+ req.Result.Replaced++
+ return nil
+ }
+ reason := organizeSkipTargetExists
+ if len(externalExisting) > 0 || len(identityExisting) > 0 || o.allExistingPathsInDB(ctx, existing) {
+ reason = organizeSkipDuplicateLibrary
+ }
+ o.log.Debug("organize skip duplicate",
+ zap.String("src", src), zap.String("dest_dir", destDir), zap.String("reason", reason))
+ req.Result.Skipped++
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: src, Target: dst, Action: "skip", Reason: reason,
+ MediaType: layout.MediaType, Category: layout.Category, Title: title,
+ })
+ return nil
+ }
+
+ req.Result.Items = append(req.Result.Items, OrganizePreviewItem{
+ Source: src, Target: dst, Action: "organize",
+ MediaType: layout.MediaType, Category: layout.Category, Title: title,
+ })
+ if req.DryRun {
+ req.Result.Organized++
+ return nil
+ }
+ if err := os.MkdirAll(destDir, 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
+ return err
+ }
+ if _, err := os.Stat(dst); err == nil {
+ req.Result.Skipped++
+ if len(req.Result.Items) > 0 {
+ req.Result.Items[len(req.Result.Items)-1].Action = "skip"
+ req.Result.Items[len(req.Result.Items)-1].Reason = organizeSkipTargetExists
+ }
+ return nil
+ }
+ if err := transferFile(src, dst, req.Mode); err != nil {
+ return err
+ }
+ if err := transferSidecarNFO(src, dst, req.Mode); err != nil {
+ o.log.Warn("organize sidecar nfo failed",
+ zap.String("from", src), zap.String("to", dst), zap.Error(err))
+ }
+ req.Result.Organized++
+ return nil
+}
+
+func organizeTitleFromParentFolder(src, sourceRoot string, seriesLike bool) (string, int) {
+ if !seriesLike {
+ return "", 0
+ }
+ raw := seriesFolderTitle(src, sourceRoot)
+ if strings.TrimSpace(raw) == "" {
+ return "", 0
+ }
+ title, year := CleanQuery(raw)
+ if title == "" {
+ title = strings.TrimSpace(raw)
+ }
+ return title, year
+}
+
+func shouldSkipOrganizeSourceVideo(path, sourceRoot string) (bool, string) {
+ cleanPath := filepath.Clean(path)
+ cleanRoot := filepath.Clean(sourceRoot)
+ if rel, err := filepath.Rel(cleanRoot, cleanPath); err == nil && rel != "." && !strings.HasPrefix(rel, "..") {
+ dir := filepath.Dir(rel)
+ if dir != "." {
+ for _, part := range strings.Split(dir, string(os.PathSeparator)) {
+ switch normalizeOrganizeCategoryKey(part) {
+ case "sample", "samples", "trailer", "trailers", "preview", "previews", "teaser", "teasers":
+ return true, organizeSkipSampleClip
+ }
+ }
+ }
+ }
+ base := strings.ToLower(strings.TrimSuffix(filepath.Base(cleanPath), filepath.Ext(cleanPath)))
+ normalized := strings.NewReplacer("_", " ", "-", " ", ".", " ").Replace(base)
+ fields := strings.Fields(normalized)
+ if len(fields) == 0 {
+ return false, ""
+ }
+ if len(fields) == 1 && strings.HasPrefix(fields[0], "sample") {
+ return true, organizeSkipSampleClip
+ }
+ last := fields[len(fields)-1]
+ switch last {
+ case "sample", "trailer", "preview", "teaser":
+ return true, organizeSkipSampleClip
+ }
+ return false, ""
+}
+
+func normalizeOrganizeMediaType(mediaType string) string {
+ switch strings.ToLower(strings.TrimSpace(mediaType)) {
+ case "movie", "film":
+ return "movie"
+ case "tv", "series", "show", "drama":
+ return "tv"
+ case "anime", "animation":
+ return "anime"
+ case "variety":
+ return "variety"
+ case "adult", "nsfw":
+ return "adult"
+ default:
+ return ""
+ }
+}
+
+func organizeWeakFileTitle(title string) bool {
+ title = strings.TrimSpace(title)
+ if title == "" {
+ return true
+ }
+ fields := strings.Fields(strings.ToLower(title))
+ if len(fields) == 0 {
+ return true
+ }
+ meaningful := 0
+ for _, field := range fields {
+ if _, ok := noiseTokenSet[field]; ok {
+ continue
+ }
+ if _, ok := releaseBoundaryTokenSet[field]; ok {
+ continue
+ }
+ if len(field) == 4 && strings.HasPrefix(field, "20") {
+ continue
+ }
+ meaningful++
+ }
+ return meaningful == 0
+}
+
+func (o *OrganizerService) inferMediaTypeForSourceFile(src, title string, season, episode int) string {
+ if season > 0 || episode > 0 {
+ return "tv"
+ }
+ return normalizeMediaType("", title, src)
+}
diff --git a/internal/service/organizer_directory_sources.go b/internal/service/organizer_directory_sources.go
new file mode 100644
index 0000000..78302e3
--- /dev/null
+++ b/internal/service/organizer_directory_sources.go
@@ -0,0 +1,113 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "os"
+ "path/filepath"
+ "strings"
+)
+
+// OrganizeSourceCandidate is a selectable organize source directory surfaced to
+// the UI so operators can organize an arbitrary directory (such as the download
+// directory) and not only registered libraries.
+type OrganizeSourceCandidate struct {
+ Label string `json:"label"`
+ Path string `json:"path"`
+ Kind string `json:"kind"` // "download" | "media"
+}
+
+// OrganizeSourceCandidates returns the configured directories that are valid
+// organize sources (download dir + media dir). It uses the container-visible
+// paths; in NAS direct-read mode those equal the host paths the operator sees.
+func (o *OrganizerService) OrganizeSourceCandidates(ctx context.Context) []OrganizeSourceCandidate {
+ out := []OrganizeSourceCandidate{}
+ seen := map[string]struct{}{}
+ add := func(label, path, kind string) {
+ path = strings.TrimSpace(path)
+ if path == "" || path == "." || strings.HasPrefix(path, ".") {
+ return
+ }
+ clean := filepath.Clean(path)
+ if !isAccessibleDir(clean) {
+ return
+ }
+ if _, ok := seen[clean]; ok {
+ return
+ }
+ seen[clean] = struct{}{}
+ out = append(out, OrganizeSourceCandidate{Label: label, Path: clean, Kind: kind})
+ }
+ add("默认整理源", o.settingValue(ctx, "organize.source_dir"), "source")
+ add("下载器保存目录", o.settingValue(ctx, "qbittorrent.savepath"), "download")
+ add("下载目录", envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads"), "download")
+ add("媒体目录", envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media"), "media")
+ return out
+}
+
+func (o *OrganizerService) settingValue(ctx context.Context, key string) string {
+ if o.repo == nil || o.repo.Setting == nil {
+ return ""
+ }
+ if v, err := o.repo.Setting.Get(ctx, key); err == nil {
+ return strings.TrimSpace(v)
+ }
+ return ""
+}
+
+// defaultSourceRoot resolves the source root for a directory organize:
+// explicit override → organize.source_dir setting → qB default save path →
+// download container dir.
+func (o *OrganizerService) defaultSourceRoot(ctx context.Context, override string) string {
+ if r := strings.TrimSpace(override); r != "" {
+ return r
+ }
+ if v := o.settingValue(ctx, "organize.source_dir"); v != "" {
+ return v
+ }
+ if v := o.settingValue(ctx, "qbittorrent.savepath"); v != "" {
+ return v
+ }
+ return envOrDefault("MEDIASTATION_DOWNLOAD_CONTAINER_DIR", "/downloads")
+}
+
+// defaultDestRoot resolves the destination root for a directory organize:
+// explicit override → organize.target_dir setting → media container dir.
+func (o *OrganizerService) defaultDestRoot(ctx context.Context, override string) string {
+ if r := strings.TrimSpace(override); r != "" {
+ return r
+ }
+ if o.repo != nil && o.repo.Setting != nil {
+ if v, err := o.repo.Setting.Get(ctx, "organize.target_dir"); err == nil && strings.TrimSpace(v) != "" {
+ return strings.TrimSpace(v)
+ }
+ }
+ return envOrDefault("MEDIASTATION_MEDIA_CONTAINER_DIR", "/media")
+}
+
+func ensureOrganizeDestinationWritable(dest string) error {
+ dest = strings.TrimSpace(dest)
+ if dest == "" || dest == "." {
+ return errors.New("destination path required")
+ }
+ if _, ok := ParseCloudLibraryMount(dest); ok {
+ return errors.New("organize destination must be a local writable media directory; enable cloud transfer in external storage when writing to cloud")
+ }
+ if err := os.MkdirAll(dest, 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
+ return fmt.Errorf("destination path is not a writable directory: %s: %w", dest, err)
+ }
+ probe, err := os.CreateTemp(dest, ".mediastation-write-test-*") // #nosec G304 -- dest is operator-configured organize root.
+ if err != nil {
+ return fmt.Errorf("destination path is not writable: %s: %w", dest, err)
+ }
+ name := probe.Name()
+ if closeErr := probe.Close(); closeErr != nil {
+ _ = os.Remove(name)
+ return fmt.Errorf("destination path write probe failed: %s: %w", dest, closeErr)
+ }
+ if err := os.Remove(name); err != nil {
+ return fmt.Errorf("destination path cleanup probe failed: %s: %w", dest, err)
+ }
+ return nil
+}
diff --git a/internal/service/organizer_directory_test.go b/internal/service/organizer_directory_test.go
index 86a28d0..57e1e62 100644
--- a/internal/service/organizer_directory_test.go
+++ b/internal/service/organizer_directory_test.go
@@ -576,263 +576,3 @@ func TestOrganizeDirectoryTVEpisodeDedup(t *testing.T) {
t.Fatalf("expected E02 organized at %q: %v", e02, err)
}
}
-
-func TestOrganizeDirectoryUsesDownloadCategoryLayout(t *testing.T) {
- root := t.TempDir()
- src := filepath.Join(root, "downloads")
- dest := filepath.Join(root, "media")
- writeOrgFile(t, filepath.Join(src, "国产剧", "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "kuangbiao-e01")
- writeOrgFile(t, filepath.Join(src, "华语电影", "流浪地球2.2023.2160p.WEB-DL.H265.mkv"), "wandering-earth-2")
-
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), newOrganizerTestRepo(t))
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: src,
- DestPath: dest,
- TransferMode: TransferCopy,
- })
- if err != nil {
- t.Fatalf("organize directory: %v", err)
- }
- if res.Organized != 2 || res.Replaced != 0 || res.Skipped != 0 {
- t.Fatalf("expected organized=2 replaced=0 skipped=0, got %+v", res)
- }
-
- tv := filepath.Join(dest, "电视剧", "国产剧", "狂飙", "Season 01", "狂飙 - S01E01.mkv")
- if _, err := os.Stat(tv); err != nil {
- t.Fatalf("expected TV episode organized at %q: %v", tv, err)
- }
- movie := filepath.Join(dest, "电影", "华语电影", "流浪地球2 (2023)", "流浪地球2 (2023).mkv")
- if _, err := os.Stat(movie); err != nil {
- t.Fatalf("expected movie organized at %q: %v", movie, err)
- }
-}
-
-func TestOrganizeDirectoryUsesExplicitCategoryLibraryRoot(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)
- libraryRoot := filepath.Join(dest, "电视剧", "国产剧")
- wrongType := model.Library{Name: "国产剧", Path: libraryRoot, Type: "movie", Enabled: true}
- rightType := model.Library{Name: "国产剧", Path: libraryRoot, Type: "tv", Enabled: true}
- if err := repos.Library.Create(t.Context(), &wrongType); err != nil {
- t.Fatal(err)
- }
- if err := repos.Library.Create(t.Context(), &rightType); err != nil {
- t.Fatal(err)
- }
-
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: src,
- DestPath: dest,
- MediaType: "tv",
- MediaCategory: "国产剧",
- TransferMode: TransferCopy,
- })
- if err != nil {
- t.Fatalf("organize explicit category: %v", err)
- }
- if res.Organized != 1 || len(res.Items) != 1 {
- t.Fatalf("result = %+v, want one organized item", res)
- }
- if !pathWithin(res.Items[0].Target, libraryRoot) {
- t.Fatalf("target = %q, want under %q", res.Items[0].Target, libraryRoot)
- }
- if pathWithin(res.Items[0].Target, filepath.Join(dest, "电视剧")) && !pathWithin(res.Items[0].Target, libraryRoot) {
- t.Fatalf("target landed outside category library: %q", res.Items[0].Target)
- }
-}
-
-func TestOrganizeDirectoryCreatesMissingCategoryLibraryForVisibility(t *testing.T) {
- root := t.TempDir()
- srcRoot := filepath.Join(root, "downloads")
- dest := filepath.Join(root, "media")
- source := filepath.Join(srcRoot, "Gourd.Brothers.S01E01.2026.1080p.mkv")
- target := filepath.Join(dest, "电视剧", "未分类", "Gourd Brothers", "Season 01", "Gourd Brothers - S01E01.mkv")
- writeOrgFile(t, source, "source")
- writeOrgFile(t, target, "already-there")
-
- repos := newOrganizerTestRepo(t)
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: srcRoot,
- DestPath: dest,
- MediaType: "tv",
- MediaCategory: "未分类",
- TransferMode: TransferCopy,
- AllowReplaceExisting: false,
- })
- if err != nil {
- t.Fatalf("organize missing category: %v", err)
- }
- if res.Organized != 0 || res.Skipped != 1 || len(res.Items) != 1 || res.Items[0].Reason != organizeSkipTargetExists {
- t.Fatalf("result = %+v, want skipped target exists", res)
- }
-
- var lib model.Library
- if err := repos.DB.Where("path = ?", filepath.Join(dest, "电视剧", "未分类")).First(&lib).Error; err != nil {
- t.Fatalf("missing auto-created category library: %v", err)
- }
- if lib.Name != "未分类" || lib.Type != "tv" || !lib.Enabled {
- t.Fatalf("auto-created library = %+v, want enabled tv 未分类", lib)
- }
-
- scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
- scans := scanner.ScanLibrariesForPath(t.Context(), res.DestPath, "")
- added := 0
- for _, scan := range scans {
- if scan.Error != "" {
- t.Fatalf("scan failed: %#v", scan)
- }
- added += scan.Added
- }
- if added != 1 {
- t.Fatalf("scan added = %d, want 1; scans=%#v", added, scans)
- }
-}
-
-func TestOrganizeDirectorySmartClassifiesUncategorizedSources(t *testing.T) {
- root := t.TempDir()
- src := filepath.Join(root, "downloads")
- dest := filepath.Join(root, "media")
- writeOrgFile(t, filepath.Join(src, "流浪地球2.2023.2160p.WEB-DL.mkv"), "cn-movie")
- writeOrgFile(t, filepath.Join(src, "Dune.2021.2160p.WEB-DL.mkv"), "foreign-movie")
- writeOrgFile(t, filepath.Join(src, "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "cn-tv")
- writeOrgFile(t, filepath.Join(src, "The.Last.of.Us.S01E01.2023.1080p.WEB-DL.mkv"), "western-tv")
-
- repos := newOrganizerTestRepo(t)
- if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
- t.Fatal(err)
- }
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: src,
- DestPath: dest,
- TransferMode: TransferCopy,
- })
- if err != nil {
- t.Fatalf("organize directory: %v", err)
- }
- if res.Organized != 4 {
- t.Fatalf("organized = %d, want 4; result=%+v", res.Organized, res)
- }
-
- for _, want := range []string{
- filepath.Join(dest, "电影", "华语电影", "流浪地球2 (2023)", "流浪地球2 (2023).mkv"),
- filepath.Join(dest, "电影", "外语电影", "Dune (2021)", "Dune (2021).mkv"),
- filepath.Join(dest, "电视剧", "国产剧", "狂飙", "Season 01", "狂飙 - S01E01.mkv"),
- filepath.Join(dest, "电视剧", "未分类", "The Last Of Us", "Season 01", "The Last Of Us - S01E01.mkv"),
- } {
- if _, err := os.Stat(want); err != nil {
- t.Fatalf("expected smart classified file at %q: %v; items=%+v", want, err, res.Items)
- }
- }
-}
-
-func TestOrganizeDirectorySmartClassifiesWithLocalNFO(t *testing.T) {
- root := t.TempDir()
- src := filepath.Join(root, "downloads")
- dest := filepath.Join(root, "media")
- writeOrgFile(t, filepath.Join(src, "Some.Show.S01E01.2024.1080p.mkv"), "jp-anime")
- writeOrgFile(t, filepath.Join(src, "tvshow.nfo"), `
- Some Show
- Animation
- JP
- ja
-`)
-
- repos := newOrganizerTestRepo(t)
- if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
- t.Fatal(err)
- }
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: src,
- DestPath: dest,
- TransferMode: TransferCopy,
- })
- if err != nil {
- t.Fatalf("organize directory: %v", err)
- }
- if res.Organized != 1 {
- t.Fatalf("organized = %d, want 1", res.Organized)
- }
- want := filepath.Join(dest, "动漫", "日番", "Some Show", "Season 01", "Some Show - S01E01.mkv")
- if _, err := os.Stat(want); err != nil {
- t.Fatalf("expected NFO classified episode at %q: %v", want, err)
- }
-}
-
-func TestOrganizeDirectoryScanAfterRecursesNestedDownloadFolders(t *testing.T) {
- root := t.TempDir()
- src := filepath.Join(root, "downloads")
- dest := filepath.Join(root, "media")
- writeOrgFile(t, filepath.Join(src, "国产剧", "子目录", "狂飙.S01E01.2023.1080p.WEB-DL.mkv"), "kuangbiao-e01")
- writeOrgFile(t, filepath.Join(src, "华语电影", "更深", "流浪地球2.2023.2160p.WEB-DL.H265.mkv"), "wandering-earth-2")
-
- repos := newOrganizerTestRepo(t)
- tvLib := model.Library{Name: "国产剧", Path: filepath.Join(dest, "电视剧", "国产剧"), Type: "tv", Enabled: true}
- movieLib := model.Library{Name: "华语电影", Path: filepath.Join(dest, "电影", "华语电影"), Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &tvLib); err != nil {
- t.Fatal(err)
- }
- if err := repos.Library.Create(t.Context(), &movieLib); err != nil {
- t.Fatal(err)
- }
-
- org := NewOrganizerService(&config.Config{}, zap.NewNop(), repos)
- res, err := org.OrganizeDirectory(t.Context(), OrganizeOptions{
- SourcePath: src,
- DestPath: dest,
- TransferMode: TransferCopy,
- })
- if err != nil {
- t.Fatalf("organize directory: %v", err)
- }
- if res.Organized != 2 {
- t.Fatalf("organized = %d, want 2", res.Organized)
- }
-
- scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
- scans := scanner.ScanLibrariesForPath(t.Context(), res.DestPath, "")
- if len(scans) != 2 {
- t.Fatalf("scans = %#v, want two matching libraries", scans)
- }
- added := 0
- for _, scan := range scans {
- if scan.Error != "" {
- t.Fatalf("scan failed: %#v", scan)
- }
- added += scan.Added
- }
- if added != 2 {
- t.Fatalf("scan added = %d, want 2", added)
- }
- var count int64
- if err := repos.DB.Model(&model.Media{}).Count(&count).Error; err != nil {
- t.Fatal(err)
- }
- if count != 2 {
- t.Fatalf("media rows = %d, want 2", count)
- }
-}
-
-func TestSelectOrganizeScanTargetsDedupesSamePathByPathType(t *testing.T) {
- root := t.TempDir()
- path := filepath.Join(root, "media", "电视剧", "国产剧")
- libraries := []model.Library{
- {Name: "国产剧", Path: path, Type: "movie", Enabled: true},
- {Name: "国产剧", Path: path, Type: "tv", Enabled: true},
- }
-
- targets := selectOrganizeScanTargets(libraries, filepath.Join(root, "media"), "")
- if len(targets) != 1 {
- t.Fatalf("targets = %#v, want one deduped target", targets)
- }
- if targets[0].Type != "tv" {
- t.Fatalf("target type = %q, want tv", targets[0].Type)
- }
-}
diff --git a/internal/service/organizer_directory_versions.go b/internal/service/organizer_directory_versions.go
new file mode 100644
index 0000000..a60238a
--- /dev/null
+++ b/internal/service/organizer_directory_versions.go
@@ -0,0 +1,248 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// existingVersionPaths returns existing destination files that represent the
+// same media, combining two strategies and de-duplicating by path:
+//
+// 1. DB identity: media rows already scanned into the destination root whose
+// title (case-insensitive) + year [or + season/episode] match the source.
+// This is robust to directory case/layout differences.
+// 2. Filesystem: video files inside the computed destination folder (matching
+// the SxxExx tag for episodes). Covers destinations that were not scanned.
+func (o *OrganizerService) existingVersionPaths(ctx context.Context, destRoot, destDir, title, episodeTag string, year, season, episode int) []string {
+ return mergeExistingVersionPaths(
+ o.existingByIdentity(ctx, destRoot, title, year, season, episode),
+ o.existingByFolder(destDir, episodeTag),
+ )
+}
+
+func mergeExistingVersionPaths(groups ...[]string) []string {
+ seen := map[string]struct{}{}
+ var out []string
+ add := func(p string) {
+ if p == "" {
+ return
+ }
+ c := filepath.Clean(p)
+ if _, ok := seen[c]; ok {
+ return
+ }
+ if _, err := os.Stat(c); err != nil {
+ return
+ }
+ seen[c] = struct{}{}
+ out = append(out, c)
+ }
+ for _, group := range groups {
+ for _, p := range group {
+ add(p)
+ }
+ }
+ return out
+}
+
+func (o *OrganizerService) allExistingPathsInDB(ctx context.Context, paths []string) bool {
+ if o == nil || o.repo == nil || o.repo.DB == nil || len(paths) == 0 {
+ return false
+ }
+ cleaned := make([]string, 0, len(paths))
+ seen := map[string]struct{}{}
+ for _, path := range paths {
+ path = filepath.Clean(strings.TrimSpace(path))
+ if path == "" || path == "." {
+ continue
+ }
+ if _, ok := seen[path]; ok {
+ continue
+ }
+ seen[path] = struct{}{}
+ cleaned = append(cleaned, path)
+ }
+ if len(cleaned) == 0 {
+ return false
+ }
+ var count int64
+ if err := o.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Where("path IN ?", cleaned).
+ Count(&count).Error; err != nil {
+ return false
+ }
+ return count == int64(len(cleaned))
+}
+
+// existingByIdentity finds scanned destination media matching the parsed
+// identity (case-insensitive title + year for movies; title + season/episode
+// for episodes), located under destRoot.
+func (o *OrganizerService) existingByIdentity(ctx context.Context, destRoot, title string, year, season, episode int) []string {
+ if o.repo == nil || o.repo.DB == nil {
+ return nil
+ }
+ title = strings.TrimSpace(title)
+ if title == "" {
+ return nil
+ }
+ q := o.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("deleted_at IS NULL").
+ Where("LOWER(title) = ?", strings.ToLower(title))
+ if season > 0 || episode > 0 {
+ q = q.Where("season_num = ? AND episode_num = ?", season, episode)
+ } else if year > 0 {
+ q = q.Where("year = ?", year)
+ }
+ var rows []model.Media
+ if err := q.Find(&rows).Error; err != nil {
+ return nil
+ }
+ var out []string
+ for _, r := range rows {
+ if r.Path != "" && pathWithin(r.Path, destRoot) {
+ out = append(out, r.Path)
+ }
+ }
+ return out
+}
+
+func (o *OrganizerService) existingByExternalIdentity(ctx context.Context, destRoot string, match *Match, season, episode int) []string {
+ if o.repo == nil || o.repo.DB == nil || match == nil {
+ return nil
+ }
+ var conds []string
+ var args []any
+ if match.TMDbID > 0 {
+ conds = append(conds, "tm_db_id = ?")
+ args = append(args, match.TMDbID)
+ }
+ if match.BangumiID > 0 {
+ conds = append(conds, "bangumi_id = ?")
+ args = append(args, match.BangumiID)
+ }
+ if strings.TrimSpace(match.DoubanID) != "" {
+ conds = append(conds, "douban_id = ?")
+ args = append(args, strings.TrimSpace(match.DoubanID))
+ }
+ if strings.TrimSpace(match.TheTVDBID) != "" {
+ conds = append(conds, "thetvdb_id = ?")
+ args = append(args, strings.TrimSpace(match.TheTVDBID))
+ }
+ if len(conds) == 0 {
+ return nil
+ }
+ q := o.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Where("deleted_at IS NULL").
+ Where("("+strings.Join(conds, " OR ")+")", args...)
+ if season > 0 || episode > 0 {
+ q = q.Where("season_num = ? AND episode_num = ?", season, episode)
+ }
+ var rows []model.Media
+ if err := q.Find(&rows).Error; err != nil {
+ return nil
+ }
+ var out []string
+ for _, row := range rows {
+ if row.Path != "" && pathWithin(row.Path, destRoot) {
+ out = append(out, row.Path)
+ }
+ }
+ return out
+}
+
+// existingByFolder returns video files already present in destDir that
+// represent the same media. For an episode (episodeTag != "") it matches files
+// carrying the same SxxExx tag; for a movie it matches every video file in the
+// movie folder.
+func (o *OrganizerService) existingByFolder(destDir, episodeTag string) []string {
+ entries, err := os.ReadDir(destDir)
+ if err != nil {
+ return nil
+ }
+ tag := strings.ToLower(episodeTag)
+ var out []string
+ for _, e := range entries {
+ if e.IsDir() {
+ continue
+ }
+ name := e.Name()
+ if _, ok := videoExtensions[strings.ToLower(filepath.Ext(name))]; !ok {
+ continue
+ }
+ if tag != "" && !strings.Contains(strings.ToLower(name), tag) {
+ continue
+ }
+ out = append(out, filepath.Join(destDir, name))
+ }
+ return out
+}
+
+// replaceVersions removes the existing lower-resolution files (and their NFO
+// sidecars + DB rows) and transfers src into dst.
+func (o *OrganizerService) replaceVersions(ctx context.Context, src string, existing []string, dst string, mode TransferMode) error {
+ for _, e := range existing {
+ if err := os.Remove(e); err != nil && !os.IsNotExist(err) {
+ return fmt.Errorf("remove existing %s: %w", e, err)
+ }
+ if nfo := nfoPath(e); nfo != "" {
+ _ = os.Remove(nfo)
+ }
+ if o.repo != nil && o.repo.DB != nil {
+ _ = o.repo.DB.WithContext(ctx).Where("path = ?", e).Delete(&model.Media{}).Error
+ }
+ }
+ if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
+ return err
+ }
+ if err := transferFile(src, dst, mode); err != nil {
+ return err
+ }
+ if err := transferSidecarNFO(src, dst, mode); err != nil {
+ o.log.Warn("organize sidecar nfo failed",
+ zap.String("from", src), zap.String("to", dst), zap.Error(err))
+ }
+ return nil
+}
+
+// resolutionArea returns the pixel area (width*height) of a video file for 洗版
+// comparison. It prefers ffprobe; when unavailable it falls back to a
+// resolution token in the filename (2160p/1080p/720p). Returns 0 when the
+// resolution cannot be determined, in which case the caller treats the file as
+// "unknown" and never performs a destructive replace.
+func (o *OrganizerService) resolutionArea(ctx context.Context, path string) int {
+ // Prefer a scanned media row's stored dimensions. The destination library
+ // is normally scanned with ffprobe, so its files have accurate Width/Height
+ // even after organize stripped the resolution token from the filename.
+ if o.repo != nil && o.repo.DB != nil {
+ var m model.Media
+ if err := o.repo.DB.WithContext(ctx).
+ Select("width", "height").
+ Where("path = ?", path).
+ Limit(1).Take(&m).Error; err == nil && m.Width > 0 && m.Height > 0 {
+ return m.Width * m.Height
+ }
+ }
+ if o.probe != nil {
+ if pr, err := o.probe.Probe(ctx, path); err == nil && pr != nil && pr.Width > 0 && pr.Height > 0 {
+ return pr.Width * pr.Height
+ }
+ }
+ switch detectResolutionScore(strings.ToLower(filepath.Base(path))) {
+ case 4:
+ return 3840 * 2160
+ case 3:
+ return 1920 * 1080
+ case 2:
+ return 1280 * 720
+ default:
+ return 0
+ }
+}
diff --git a/internal/service/organizer_media.go b/internal/service/organizer_media.go
new file mode 100644
index 0000000..d3bfd91
--- /dev/null
+++ b/internal/service/organizer_media.go
@@ -0,0 +1,175 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type organizeMediaRequest struct {
+ media *model.Media
+ library *model.Library
+ baseRoot string
+ mediaType string
+ mediaCategory string
+ dryRun bool
+ transferMode TransferMode
+}
+
+type organizeMediaDestination struct {
+ path string
+ libraryID string
+ mediaType string
+ category string
+}
+
+func (o *OrganizerService) resolveOrganizeMediaRequest(ctx context.Context, mediaID string, opts OrganizeOptions) (organizeMediaRequest, error) {
+ m, err := o.repo.Media.FindByID(ctx, mediaID)
+ if err != nil || m == nil {
+ return organizeMediaRequest{}, errors.New("media not found")
+ }
+ lib, err := o.repo.Library.FindByID(ctx, m.LibraryID)
+ if err != nil || lib == nil {
+ return organizeMediaRequest{}, errors.New("library not found")
+ }
+ if _, ok := ParseCloudLibraryMount(lib.Path); ok {
+ return organizeMediaRequest{}, errors.New("local organize cannot use cloud libraries directly; use external storage scan/mount for cloud media or enable cloud transfer to write to cloud")
+ }
+ baseRoot := redirectOrganizeStagingRoot(o.resolveBaseRoot(ctx, lib, opts.DestPath))
+ if _, ok := ParseCloudLibraryMount(baseRoot); ok {
+ return organizeMediaRequest{}, errors.New("organize destination must be a local writable media directory; enable cloud transfer in external storage when writing to cloud")
+ }
+ if !opts.DryRun {
+ if err := ensureOrganizeDestinationWritable(baseRoot); err != nil {
+ return organizeMediaRequest{}, err
+ }
+ }
+ if isSeriesLibraryType(lib.Type) {
+ if err := o.refreshEpisodeIdentity(m, lib); err != nil {
+ return organizeMediaRequest{}, err
+ }
+ }
+ return organizeMediaRequest{
+ media: m,
+ library: lib,
+ baseRoot: baseRoot,
+ mediaType: strings.TrimSpace(opts.MediaType),
+ mediaCategory: strings.TrimSpace(opts.MediaCategory),
+ dryRun: opts.DryRun,
+ transferMode: o.resolveTransferMode(ctx, opts.TransferMode),
+ }, nil
+}
+
+func (o *OrganizerService) buildOrganizeMediaDestination(ctx context.Context, req organizeMediaRequest) (organizeMediaDestination, error) {
+ m := req.media
+ lib := req.library
+ title := sanitizeFilename(m.Title)
+ if title == "" {
+ title = "Unknown"
+ }
+
+ mediaType := normalizeOrganizeMediaType(req.mediaType)
+ if mediaType == "" {
+ mediaType = normalizeOrganizeMediaType(lib.Type)
+ }
+ category := sanitizeFilename(strings.TrimSpace(req.mediaCategory))
+ if category == "" && o.isSmartClassifyEnabled(ctx) {
+ category = o.classifyMedia(ctx, m, mediaType)
+ }
+ if impliedType, normalizedCategory := o.mediaTypeForDirectoryCategory(category); impliedType != "" {
+ mediaType = impliedType
+ category = normalizedCategory
+ }
+ root := o.organizeRoot(req.baseRoot, mediaType, category)
+ targetLibraryID := ""
+ matchedLibrary := false
+ if category != "" {
+ targetLibrary, ok := o.organizeLibraryForLayout(ctx, req.baseRoot, mediaType, category)
+ if ok {
+ root = targetLibrary.Path
+ targetLibraryID = targetLibrary.ID
+ matchedLibrary = true
+ }
+ }
+ if !matchedLibrary && category != "" {
+ root = categoryRoot(root, category)
+ }
+ if !matchedLibrary && !req.dryRun {
+ o.ensureOrganizeLibraryForRoot(ctx, root, mediaType, category)
+ }
+ target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
+ Root: root,
+ MediaType: mediaType,
+ Category: category,
+ Title: title,
+ Source: m.Path,
+ Ext: filepath.Ext(m.Path),
+ Year: m.Year,
+ Season: m.SeasonNum,
+ Episode: m.EpisodeNum,
+ Series: isSeriesLibraryType(mediaType),
+ })
+ if err != nil {
+ return organizeMediaDestination{}, err
+ }
+ return organizeMediaDestination{path: target.Path, libraryID: targetLibraryID, mediaType: mediaType, category: category}, nil
+}
+
+func (o *OrganizerService) applyOrganizeMedia(ctx context.Context, req organizeMediaRequest, dst organizeMediaDestination) (string, error) {
+ m := req.media
+
+ // Refuse to overwrite an existing different file. 当多个 release(如
+ // 不同字幕组、不同源)刮削后被统一改名,原本不重复的文件会被映射到
+ // 同一个目标路径,盲目 move 会导致后者覆盖前者,造成数据丢失。
+ if _, err := os.Stat(dst.path); err == nil {
+ o.log.Warn("organize skipped: destination already exists",
+ zap.String("media", m.ID),
+ zap.String("from", m.Path),
+ zap.String("to", dst.path))
+ return dst.path, nil
+ }
+
+ if err := os.MkdirAll(filepath.Dir(dst.path), 0o755); err != nil { // #nosec G301 -- organized media directories must remain readable by NAS/player users.
+ return "", err
+ }
+ if err := transferFile(m.Path, dst.path, req.transferMode); err != nil {
+ return "", err
+ }
+
+ updates := map[string]any{
+ "path": dst.path,
+ "season_num": m.SeasonNum,
+ "episode_num": m.EpisodeNum,
+ }
+ if strings.TrimSpace(dst.libraryID) != "" {
+ updates["library_id"] = strings.TrimSpace(dst.libraryID)
+ }
+ if err := o.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Where("id = ?", m.ID).
+ Updates(updates).Error; err != nil {
+ return dst.path, err
+ }
+ if err := transferSidecarNFO(m.Path, dst.path, req.transferMode); err != nil {
+ o.log.Warn("organize sidecar nfo failed",
+ zap.String("media", m.ID),
+ zap.String("from", nfoPath(m.Path)),
+ zap.String("to", nfoPath(dst.path)),
+ zap.Error(err))
+ }
+ o.log.Info("organized",
+ zap.String("media", m.ID),
+ zap.String("from", m.Path),
+ zap.String("to", dst.path),
+ zap.String("category", dst.category),
+ zap.String("media_type", dst.mediaType),
+ zap.String("mode", string(req.transferMode)),
+ )
+ return dst.path, nil
+}
diff --git a/internal/service/organizer_paths.go b/internal/service/organizer_paths.go
new file mode 100644
index 0000000..229aed0
--- /dev/null
+++ b/internal/service/organizer_paths.go
@@ -0,0 +1,174 @@
+package service
+
+import (
+ "path/filepath"
+ "strings"
+)
+
+// sanitizeFilename removes characters not safe for filesystem names.
+func sanitizeFilename(s string) string {
+ r := strings.NewReplacer(
+ "/", " ", "\\", " ", ":", " ", "*", "", "?", "",
+ "\"", "", "<", "", ">", "", "|", "",
+ )
+ return strings.TrimSpace(r.Replace(s))
+}
+
+func (o *OrganizerService) organizeRoot(libraryPath, mediaType, category string) string {
+ typeDir := o.mediaTypeRootDirForCategory(mediaType, category)
+ if typeDir == "" || pathAlreadyEndsWith(libraryPath, typeDir) {
+ return libraryPath
+ }
+ if isGenericMediaRoot(libraryPath) {
+ return filepath.Join(libraryPath, typeDir)
+ }
+ return libraryPath
+}
+
+func (o *OrganizerService) mediaTypeRootDirForCategory(mediaType, category string) string {
+ if root := o.categoryPhysicalRootDir(category); root != "" {
+ return root
+ }
+ return mediaTypeRootDir(mediaType)
+}
+
+func (o *OrganizerService) categoryPhysicalRootDir(category string) string {
+ key := normalizeOrganizeCategoryKey(category)
+ if key == "" {
+ return ""
+ }
+ categories := o.categoryMap()
+ match := func(values ...string) bool {
+ for _, value := range values {
+ if key == normalizeOrganizeCategoryKey(value) {
+ return true
+ }
+ }
+ return false
+ }
+ switch {
+ case match(
+ categoryName(categories, "cn_anime", "国漫"),
+ categoryName(categories, "jp_anime", "日番"),
+ categoryName(categories, "children", "儿童"),
+ "国漫", "国产动漫", "日番", "番剧", "日漫", "日本动漫", "日本动画", "儿童", "少儿",
+ ):
+ return "动漫"
+ case match(
+ categoryName(categories, "domestic_tv", "国产剧"),
+ categoryName(categories, "euus_tv", "欧美剧"),
+ categoryName(categories, "jk_tv", "日韩剧"),
+ categoryName(categories, "variety", "综艺"),
+ categoryName(categories, "documentary", "纪录片"),
+ categoryName(categories, "uncategorized_tv", "未分类"),
+ "国产剧", "欧美剧", "日韩剧", "日剧", "韩剧", "综艺", "真人秀", "纪录片", "纪录", "未分类",
+ ):
+ return "电视剧"
+ case match(
+ categoryName(categories, "animation_movie", "动画电影"),
+ categoryName(categories, "chinese_movie", "华语电影"),
+ categoryName(categories, "foreign_movie", "外语电影"),
+ categoryName(categories, "euus_movie", "欧美电影"),
+ categoryName(categories, "jk_movie", "日韩电影"),
+ "动画电影", "动漫电影", "华语电影", "国产电影", "外语电影", "欧美电影", "日韩电影",
+ ):
+ return "电影"
+ case match(categoryName(categories, "adult", "成人"), categoryName(categories, "adult_9kg", "9KG"), categoryName(categories, "adult_jav", "番号"), "成人", "9kg", "番号", "jav"):
+ return "成人"
+ default:
+ return ""
+ }
+}
+
+func categoryRoot(root, category string) string {
+ if strings.TrimSpace(category) == "" || pathAlreadyEndsWith(root, category) {
+ return root
+ }
+ return filepath.Join(root, category)
+}
+
+func pathWithin(path, root string) bool {
+ cleanPath := filepath.Clean(path)
+ cleanRoot := filepath.Clean(root)
+ if strings.EqualFold(cleanPath, cleanRoot) {
+ return true
+ }
+ rel, err := filepath.Rel(cleanRoot, cleanPath)
+ if err != nil {
+ return false
+ }
+ return rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator))
+}
+
+func mediaTypeRootDir(mediaType string) string {
+ switch normalizeMediaType(mediaType, "", "") {
+ case "movie":
+ return "电影"
+ case "anime":
+ return "动漫"
+ case "tv", "variety":
+ return "电视剧"
+ case "adult":
+ return "成人"
+ default:
+ return ""
+ }
+}
+
+func isGenericMediaRoot(path string) bool {
+ base := strings.ToLower(strings.TrimSpace(filepath.Base(filepath.Clean(path))))
+ switch base {
+ case "media", "medias", "library", "libraries", "organized", "整理":
+ return true
+ default:
+ return false
+ }
+}
+
+// organizeStagingFolderNames lists manual-organize style staging folder names.
+// These workspaces should not remain as first-level category folders after an
+// organize operation; redirectOrganizeStagingRoot lifts them to the media root.
+func organizeStagingFolderNames() map[string]struct{} {
+ return map[string]struct{}{
+ "手动整理": {}, "手动整理入库": {}, "待整理": {}, "待分类": {},
+ "manual": {}, "manual_organize": {}, "manualorganize": {}, "staging": {}, "inbox": {},
+ }
+}
+
+func isOrganizeStagingDir(path string) bool {
+ base := strings.ToLower(strings.TrimSpace(filepath.Base(filepath.Clean(path))))
+ if base == "" {
+ return false
+ }
+ _, ok := organizeStagingFolderNames()[base]
+ return ok
+}
+
+func redirectOrganizeStagingRoot(root string) string {
+ cleaned := filepath.Clean(strings.TrimSpace(root))
+ if cleaned == "" || cleaned == "." {
+ return root
+ }
+ for isOrganizeStagingDir(cleaned) {
+ parent := filepath.Dir(cleaned)
+ if parent == cleaned || parent == "." || parent == string(filepath.Separator) {
+ break
+ }
+ cleaned = parent
+ }
+ return cleaned
+}
+
+func pathAlreadyEndsWith(path, suffix string) bool {
+ base := strings.TrimSpace(filepath.Base(filepath.Clean(path)))
+ return strings.EqualFold(base, suffix)
+}
+
+func isSeriesLibraryType(mediaType string) bool {
+ switch normalizeMediaType(mediaType, "", "") {
+ case "tv", "anime", "variety":
+ return true
+ default:
+ return false
+ }
+}
diff --git a/internal/service/organizer_reclassify_scanned.go b/internal/service/organizer_reclassify_scanned.go
new file mode 100644
index 0000000..38ac643
--- /dev/null
+++ b/internal/service/organizer_reclassify_scanned.go
@@ -0,0 +1,200 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// MediaCategoryReclassifyOptions controls the metadata-based category audit.
+// Empty LibraryIDs means all enabled local libraries.
+type MediaCategoryReclassifyOptions struct {
+ LibraryIDs []string
+ DryRun bool
+}
+
+// ReclassifyMisclassifiedMedia corrects already-scanned local media whose
+// stored metadata clearly disagrees with the current library/category.
+func (o *OrganizerService) ReclassifyMisclassifiedMedia(ctx context.Context, opts MediaCategoryReclassifyOptions) (*OrganizeResult, error) {
+ res := &OrganizeResult{DryRun: opts.DryRun}
+ if o == nil || o.repo == nil || o.repo.DB == nil || o.repo.Library == nil {
+ return res, nil
+ }
+ libraries, err := o.repo.Library.List(ctx)
+ if err != nil {
+ return res, err
+ }
+ filterIDs := compactLibraryIDs(opts.LibraryIDs...)
+ filter := map[string]struct{}{}
+ for _, id := range filterIDs {
+ filter[id] = struct{}{}
+ }
+ libByID := make(map[string]model.Library, len(libraries))
+ for _, lib := range libraries {
+ if !lib.Enabled || strings.TrimSpace(lib.ID) == "" {
+ continue
+ }
+ if len(filter) > 0 {
+ if _, ok := filter[lib.ID]; !ok {
+ continue
+ }
+ }
+ libByID[lib.ID] = lib
+ }
+ if len(libByID) == 0 {
+ return res, nil
+ }
+
+ query := o.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("deleted_at IS NULL")
+ if len(filter) > 0 {
+ query = query.Where("library_id IN ?", filterIDs)
+ }
+ var rows []model.Media
+ err = query.FindInBatches(&rows, 500, func(_ *gorm.DB, _ int) error {
+ for i := range rows {
+ lib, ok := libByID[rows[i].LibraryID]
+ if !ok {
+ continue
+ }
+ changed, err := o.reclassifyScannedMedia(ctx, rows[i], lib, opts.DryRun, res)
+ if err != nil {
+ res.Errors = append(res.Errors, fmt.Sprintf("%s: %s", rows[i].Title, err.Error()))
+ if o.log != nil {
+ o.log.Warn("metadata category reclassify failed",
+ zap.String("media", rows[i].ID),
+ zap.String("path", rows[i].Path),
+ zap.Error(err))
+ }
+ continue
+ }
+ if changed && o.log != nil {
+ o.log.Debug("metadata category reclassify applied",
+ zap.String("media", rows[i].ID),
+ zap.String("title", rows[i].Title))
+ }
+ }
+ return nil
+ }).Error
+ return res, err
+}
+
+func (o *OrganizerService) reclassifyScannedMedia(ctx context.Context, media model.Media, lib model.Library, dryRun bool, res *OrganizeResult) (bool, error) {
+ if res == nil || !lib.Enabled || strings.TrimSpace(media.Path) == "" {
+ return false, nil
+ }
+ if _, ok := ParseCloudLibraryMount(lib.Path); ok {
+ return false, nil
+ }
+ if !organizeFileExists(media.Path) || !mediaHasReliableCategoryMetadata(media) {
+ return false, nil
+ }
+
+ mediaType := normalizeOrganizeMediaType(lib.Type)
+ category := o.classifyMedia(ctx, &media, mediaType)
+ if category == "" {
+ return false, nil
+ }
+ if impliedType, normalizedCategory := o.mediaTypeForDirectoryCategory(category); impliedType != "" {
+ mediaType = impliedType
+ category = normalizedCategory
+ }
+ if mediaType == "" {
+ mediaType = normalizeOrganizeMediaType(lib.Type)
+ }
+ if mediaType == "" {
+ return false, nil
+ }
+
+ baseRoot := redirectOrganizeStagingRoot(o.resolveBaseRoot(ctx, &lib, ""))
+ targetLibrary, matched := o.organizeLibraryForLayout(ctx, baseRoot, mediaType, category)
+ if !matched || strings.TrimSpace(targetLibrary.ID) == "" || strings.TrimSpace(targetLibrary.Path) == "" {
+ return false, nil
+ }
+ if strings.EqualFold(targetLibrary.ID, lib.ID) && pathWithin(media.Path, targetLibrary.Path) {
+ return false, nil
+ }
+ if pathWithin(media.Path, targetLibrary.Path) {
+ return o.reclassifyScannedMediaLibraryOnly(ctx, media, lib, targetLibrary, category, mediaType, dryRun, res)
+ }
+
+ title := sanitizeFilename(strings.TrimSpace(media.Title))
+ if title == "" {
+ title = "Unknown"
+ }
+ target, err := o.buildOrganizeTargetPath(ctx, organizeTargetInput{
+ Root: targetLibrary.Path,
+ MediaType: mediaType,
+ Category: category,
+ Title: title,
+ Source: media.Path,
+ Ext: filepath.Ext(media.Path),
+ Year: media.Year,
+ Season: media.SeasonNum,
+ Episode: media.EpisodeNum,
+ Series: isSeriesLibraryType(mediaType),
+ })
+ if err != nil {
+ return false, err
+ }
+ return o.reclassifyExistingMedia(ctx, organizeExistingReclassifyRequest{
+ Source: media.Path,
+ Target: target.Path,
+ DestRoot: firstNonEmpty(baseRoot, lib.Path),
+ TargetLibraryID: targetLibrary.ID,
+ Existing: []string{media.Path},
+ DryRun: dryRun,
+ MediaType: mediaType,
+ Category: category,
+ Title: title,
+ Year: media.Year,
+ Season: media.SeasonNum,
+ Episode: media.EpisodeNum,
+ Result: res,
+ })
+}
+
+func (o *OrganizerService) reclassifyScannedMediaLibraryOnly(ctx context.Context, media model.Media, oldLib, targetLib model.Library, category, mediaType string, dryRun bool, res *OrganizeResult) (bool, error) {
+ res.Items = append(res.Items, OrganizePreviewItem{
+ Source: media.Path,
+ Target: media.Path,
+ Action: "reclassify",
+ Reason: "metadata category library changed",
+ MediaType: mediaType,
+ Category: category,
+ Title: media.Title,
+ })
+ if dryRun {
+ res.Reclassified++
+ return true, nil
+ }
+ if err := o.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Where("id = ?", media.ID).
+ Update("library_id", targetLib.ID).Error; err != nil {
+ return false, err
+ }
+ if o.log != nil {
+ o.log.Info("media library reclassified by metadata",
+ zap.String("media", media.ID),
+ zap.String("path", media.Path),
+ zap.String("from_library", oldLib.ID),
+ zap.String("to_library", targetLib.ID),
+ zap.String("category", category),
+ zap.String("media_type", mediaType))
+ }
+ res.Reclassified++
+ return true, nil
+}
+
+func mediaHasReliableCategoryMetadata(media model.Media) bool {
+ return media.NSFW ||
+ strings.TrimSpace(media.Languages) != "" ||
+ strings.TrimSpace(media.Countries) != "" ||
+ strings.TrimSpace(media.Genres) != ""
+}
diff --git a/internal/service/organizer_scan.go b/internal/service/organizer_scan.go
index 5e85a69..a4aa722 100644
--- a/internal/service/organizer_scan.go
+++ b/internal/service/organizer_scan.go
@@ -28,6 +28,7 @@ type OrganizeScrapeSummary struct {
Name string `json:"name"`
Path string `json:"path"`
Matched int `json:"matched"`
+ Processed int `json:"processed"`
Skipped bool `json:"skipped,omitempty"`
Reason string `json:"reason,omitempty"`
Error string `json:"error,omitempty"`
@@ -53,7 +54,7 @@ func OrganizeScrapeAfterEnabled(ctx context.Context, repo *repository.Container)
// change: scanning after a no-op organize can turn a harmless restart into a
// full library ffprobe sweep.
func OrganizeResultHasChanges(res *OrganizeResult) bool {
- return res != nil && (res.Organized > 0 || res.Replaced > 0)
+ return res != nil && (res.Organized > 0 || res.Replaced > 0 || res.Reclassified > 0)
}
// OrganizeResultNeedsVisibilitySync reports whether a just-finished organize
@@ -163,11 +164,12 @@ func (s *ScannerService) scrapeOrganizeTargets(ctx context.Context, targets []mo
// Organize is an explicit ingest workflow: after rename/classification,
// previously failed no_match rows should be retried so the operator does
// not need to run a separate manual scrape.
- matched, err := s.scraper.EnrichLibrary(ctx, lib.ID, true)
+ result, err := s.scraper.EnrichLibraryDetailedWithOptions(ctx, lib.ID, skipEpisodeArtworkOptions(true))
if err != nil {
summary.Error = err.Error()
} else {
- summary.Matched = matched
+ summary.Matched = result.Matched
+ summary.Processed = result.Processed
}
out = append(out, summary)
}
diff --git a/internal/service/organizer_scrape_test.go b/internal/service/organizer_scrape_test.go
index db4c91c..4ea64fe 100644
--- a/internal/service/organizer_scrape_test.go
+++ b/internal/service/organizer_scrape_test.go
@@ -279,6 +279,318 @@ func TestOrganizeDirectoryMetadataCategoryOverridesDownloadFolder(t *testing.T)
}
}
+func TestOrganizeDirectoryReclassifiesExistingWrongCategoryMedia(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ if r.URL.Path != "/search/tv" {
+ http.NotFound(w, r)
+ return
+ }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 292696,
+ "name": "莫离",
+ "original_name": "The First Jasmine",
+ "original_language": "zh",
+ "origin_country": []string{"CN"},
+ "genre_ids": []int{18},
+ "first_air_date": "2026-06-23",
+ }},
+ })
+ }))
+ defer upstream.Close()
+
+ repos := newOrganizerTestRepo(t)
+ cfg := &config.Config{}
+ cfg.Organizer.SmartClassify = true
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ scraper := NewScraperService(cfg, zap.NewNop(), repos, NewTMDbProvider(cfg, zap.NewNop(), nil), nil, nil, nil, NewHub(zap.NewNop()))
+
+ root := t.TempDir()
+ srcRoot := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ sourceFile := filepath.Join(srcRoot, "欧美剧", "The.First.Jasmine.S01.1080p.TX.WEB-DL.AAC2.0.H.264-MWeb", "The.First.Jasmine.S01E01.1080p.TX.WEB-DL.AAC2.0.H.264-MWeb.mkv")
+ writeOrgFile(t, sourceFile, "episode")
+
+ euusLib := model.Library{Name: "欧美剧", Path: filepath.Join(dest, "电视剧", "欧美剧"), Type: "tv", Enabled: true}
+ domesticLib := model.Library{Name: "国产剧", Path: filepath.Join(dest, "电视剧", "国产剧"), Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &euusLib); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &domesticLib); err != nil {
+ t.Fatal(err)
+ }
+
+ wrongPath := filepath.Join(euusLib.Path, "The First Jasmine", "Season 01", "The First Jasmine - S01E01.mkv")
+ writeOrgFile(t, wrongPath, "existing")
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: euusLib.ID,
+ Title: "莫离",
+ OriginalName: "The First Jasmine",
+ Path: wrongPath,
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ TMDbID: 292696,
+ Languages: "zh",
+ Countries: "CN",
+ Genres: "剧情",
+ ScrapeStatus: "matched",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ organizer := NewOrganizerService(cfg, zap.NewNop(), repos)
+ organizer.SetScraper(scraper)
+ res, err := organizer.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: srcRoot,
+ DestPath: euusLib.Path,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ want := filepath.Join(domesticLib.Path, "莫离", "Season 01", "莫离 - S01E01.mkv")
+ if res.Reclassified != 1 || res.Organized != 0 || res.Skipped != 0 {
+ t.Fatalf("result = %+v, want reclassified=1 only", res)
+ }
+ if _, err := os.Stat(wrongPath); !os.IsNotExist(err) {
+ t.Fatalf("wrong category path should be moved away, stat err=%v", err)
+ }
+ if _, err := os.Stat(want); err != nil {
+ t.Fatalf("reclassified media missing at %q: %v; items=%#v", want, err, res.Items)
+ }
+ if _, err := os.Stat(sourceFile); err != nil {
+ t.Fatalf("source download should remain untouched: %v", err)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "path = ?", want).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.LibraryID != domesticLib.ID {
+ t.Fatalf("library_id = %q, want domestic library %q", got.LibraryID, domesticLib.ID)
+ }
+}
+
+func TestOrganizeDirectoryReclassifiesScannedAnimeUsingDBMetadata(t *testing.T) {
+ repos := newOrganizerTestRepo(t)
+ cfg := &config.Config{}
+ cfg.Organizer.SmartClassify = true
+
+ root := t.TempDir()
+ dest := filepath.Join(root, "media")
+ euusLib := model.Library{Name: "欧美剧", Path: filepath.Join(dest, "电视剧", "欧美剧"), Type: "tv", Enabled: true}
+ tvAnimeLib := model.Library{Name: "国漫", Path: filepath.Join(dest, "电视剧", "国漫"), Type: "tv", Enabled: true}
+ animeLib := model.Library{Name: "国漫", Path: filepath.Join(dest, "动漫", "国漫"), Type: "anime", Enabled: true}
+ for _, lib := range []*model.Library{&euusLib, &tvAnimeLib, &animeLib} {
+ if err := repos.Library.Create(t.Context(), lib); err != nil {
+ t.Fatal(err)
+ }
+ }
+
+ wrongPath := filepath.Join(euusLib.Path, "Blades Of The Guardians", "Season 2", "Blades Of The Guardians - S02E01-1080p.TX.WEB-DL.mkv")
+ writeOrgFile(t, wrongPath, "episode")
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: euusLib.ID,
+ Title: "镖人",
+ OriginalName: "Blades Of The Guardians",
+ Path: wrongPath,
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ TMDbID: 107463,
+ Languages: "zh",
+ Countries: "CN",
+ Genres: "动画,动作冒险",
+ ScrapeStatus: "matched",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ organizer := NewOrganizerService(cfg, zap.NewNop(), repos)
+ res, err := organizer.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: filepath.Join(euusLib.Path, "Blades Of The Guardians"),
+ DestPath: euusLib.Path,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ want := filepath.Join(animeLib.Path, "镖人", "Season 02", "镖人 - S02E01.mkv")
+ if res.Reclassified != 1 || res.Organized != 0 {
+ t.Fatalf("result = %+v, want scanned DB metadata reclassified only", res)
+ }
+ if _, err := os.Stat(wrongPath); !os.IsNotExist(err) {
+ t.Fatalf("wrong anime path should be moved away, stat err=%v", err)
+ }
+ if _, err := os.Stat(want); err != nil {
+ t.Fatalf("anime should move to physical anime library at %q: %v; items=%#v", want, err, res.Items)
+ }
+ var got model.Media
+ if err := repos.DB.First(&got, "path = ?", want).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.LibraryID != animeLib.ID {
+ t.Fatalf("library_id = %q, want anime library %q", got.LibraryID, animeLib.ID)
+ }
+}
+
+func TestReclassifyMisclassifiedMediaMovesScannedAnimeToPhysicalAnimeLibrary(t *testing.T) {
+ repos := newOrganizerTestRepo(t)
+ cfg := &config.Config{}
+ cfg.Organizer.SmartClassify = true
+
+ root := t.TempDir()
+ dest := filepath.Join(root, "media")
+ euusLib := model.Library{Name: "欧美剧", Path: filepath.Join(dest, "电视剧", "欧美剧"), Type: "tv", Enabled: true}
+ tvAnimeLib := model.Library{Name: "国漫", Path: filepath.Join(dest, "电视剧", "国漫"), Type: "tv", Enabled: true}
+ animeLib := model.Library{Name: "国漫", Path: filepath.Join(dest, "动漫", "国漫"), Type: "anime", Enabled: true}
+ for _, lib := range []*model.Library{&euusLib, &tvAnimeLib, &animeLib} {
+ if err := repos.Library.Create(t.Context(), lib); err != nil {
+ t.Fatal(err)
+ }
+ }
+
+ wrongPath := filepath.Join(euusLib.Path, "Blades Of The Guardians", "Season 2", "Blades Of The Guardians - S02E01.mkv")
+ writeOrgFile(t, wrongPath, "episode")
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: euusLib.ID,
+ Title: "镖人",
+ OriginalName: "Blades Of The Guardians",
+ Path: wrongPath,
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ TMDbID: 107463,
+ Languages: "zh",
+ Countries: "CN",
+ Genres: "动画,动作冒险",
+ ScrapeStatus: "matched",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ organizer := NewOrganizerService(cfg, zap.NewNop(), repos)
+ res, err := organizer.ReclassifyMisclassifiedMedia(t.Context(), MediaCategoryReclassifyOptions{})
+ if err != nil {
+ t.Fatalf("reclassify media: %v", err)
+ }
+ want := filepath.Join(animeLib.Path, "镖人", "Season 02", "镖人 - S02E01.mkv")
+ if res.Reclassified != 1 {
+ t.Fatalf("reclassified = %d, want 1; items=%#v errors=%#v", res.Reclassified, res.Items, res.Errors)
+ }
+ if _, err := os.Stat(want); err != nil {
+ t.Fatalf("bulk reclassify target missing at %q: %v", want, err)
+ }
+ var got model.Media
+ if err := repos.DB.First(&got, "path = ?", want).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.LibraryID != animeLib.ID {
+ t.Fatalf("library_id = %q, want anime library %q", got.LibraryID, animeLib.ID)
+ }
+}
+
+func TestOrganizeDirectoryCleansWrongCategoryDuplicateWhenTargetExists(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ if r.URL.Path != "/search/tv" {
+ http.NotFound(w, r)
+ return
+ }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 292696,
+ "name": "莫离",
+ "original_name": "The First Jasmine",
+ "original_language": "zh",
+ "origin_country": []string{"CN"},
+ "genre_ids": []int{18},
+ "first_air_date": "2026-06-23",
+ }},
+ })
+ }))
+ defer upstream.Close()
+
+ repos := newOrganizerTestRepo(t)
+ cfg := &config.Config{}
+ cfg.Organizer.SmartClassify = true
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ scraper := NewScraperService(cfg, zap.NewNop(), repos, NewTMDbProvider(cfg, zap.NewNop(), nil), nil, nil, nil, NewHub(zap.NewNop()))
+
+ root := t.TempDir()
+ srcRoot := filepath.Join(root, "downloads")
+ dest := filepath.Join(root, "media")
+ sourceFile := filepath.Join(srcRoot, "欧美剧", "The.First.Jasmine.S01.1080p.TX.WEB-DL.AAC2.0.H.264-MWeb", "The.First.Jasmine.S01E01.1080p.TX.WEB-DL.AAC2.0.H.264-MWeb.mkv")
+ writeOrgFile(t, sourceFile, "episode")
+
+ euusLib := model.Library{Name: "欧美剧", Path: filepath.Join(dest, "电视剧", "欧美剧"), Type: "tv", Enabled: true}
+ domesticLib := model.Library{Name: "国产剧", Path: filepath.Join(dest, "电视剧", "国产剧"), Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &euusLib); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &domesticLib); err != nil {
+ t.Fatal(err)
+ }
+
+ targetPath := filepath.Join(domesticLib.Path, "莫离", "Season 01", "莫离 - S01E01.mkv")
+ wrongPath := filepath.Join(euusLib.Path, "The First Jasmine", "Season 01", "The First Jasmine - S01E01.mkv")
+ writeOrgFile(t, targetPath, "same-bytes")
+ writeOrgFile(t, wrongPath, "same-bytes")
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: domesticLib.ID,
+ Title: "莫离",
+ Path: targetPath,
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ TMDbID: 292696,
+ ScrapeStatus: "matched",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: euusLib.ID,
+ Title: "莫离",
+ Path: wrongPath,
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ TMDbID: 292696,
+ ScrapeStatus: "matched",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ organizer := NewOrganizerService(cfg, zap.NewNop(), repos)
+ organizer.SetScraper(scraper)
+ res, err := organizer.OrganizeDirectory(t.Context(), OrganizeOptions{
+ SourcePath: srcRoot,
+ DestPath: euusLib.Path,
+ TransferMode: TransferCopy,
+ })
+ if err != nil {
+ t.Fatalf("organize directory: %v", err)
+ }
+ if res.Reclassified != 1 || res.Organized != 0 || res.Skipped != 0 {
+ t.Fatalf("result = %+v, want reclassified=1 only", res)
+ }
+ if _, err := os.Stat(targetPath); err != nil {
+ t.Fatalf("canonical target should remain: %v", err)
+ }
+ if _, err := os.Stat(wrongPath); !os.IsNotExist(err) {
+ t.Fatalf("wrong category duplicate should be removed, stat err=%v", err)
+ }
+ var wrongRows int64
+ if err := repos.DB.Model(&model.Media{}).Where("path = ?", wrongPath).Count(&wrongRows).Error; err != nil {
+ t.Fatal(err)
+ }
+ if wrongRows != 0 {
+ t.Fatalf("wrong DB rows = %d, want 0", wrongRows)
+ }
+ if _, err := os.Stat(sourceFile); err != nil {
+ t.Fatalf("source download should remain untouched: %v", err)
+ }
+}
+
func TestOrganizeDirectoryDoesNotScrapeByDownloadCategoryFolder(t *testing.T) {
var queries []string
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -650,6 +962,9 @@ func TestOrganizeResultNeedsVisibilitySyncIgnoresScannedDuplicates(t *testing.T)
if !OrganizeResultNeedsVisibilitySync(&OrganizeResult{Organized: 1}) {
t.Fatal("organized files must trigger visibility scan")
}
+ if !OrganizeResultNeedsVisibilitySync(&OrganizeResult{Reclassified: 1}) {
+ t.Fatal("reclassified files must trigger visibility scan")
+ }
}
func TestOrganizeScrapeAfterEnabledDefaultsOn(t *testing.T) {
diff --git a/internal/service/organizer_settings.go b/internal/service/organizer_settings.go
new file mode 100644
index 0000000..338af5e
--- /dev/null
+++ b/internal/service/organizer_settings.go
@@ -0,0 +1,75 @@
+package service
+
+import (
+ "context"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// resolveBaseRoot picks the organize destination root (目的地目录): a
+// per-request override wins, then the organize.target_dir setting, then the
+// library's own path.
+func (o *OrganizerService) resolveBaseRoot(ctx context.Context, lib *model.Library, override string) string {
+ if r := strings.TrimSpace(override); r != "" {
+ return r
+ }
+ if o.repo != nil && o.repo.Setting != nil {
+ if v, err := o.repo.Setting.Get(ctx, "organize.target_dir"); err == nil && strings.TrimSpace(v) != "" {
+ return strings.TrimSpace(v)
+ }
+ }
+ return lib.Path
+}
+
+// resolveSourceRoot picks the organize source root (源目录,待整理文件所在目录):
+// a per-request override wins, then the organize.source_dir setting, then the
+// library's own path. Library organize only touches media located under this
+// root, so operators can point at a specific download/staging folder.
+func (o *OrganizerService) resolveSourceRoot(ctx context.Context, lib *model.Library, override string) string {
+ if r := strings.TrimSpace(override); r != "" {
+ return r
+ }
+ if o.repo != nil && o.repo.Setting != nil {
+ if v, err := o.repo.Setting.Get(ctx, "organize.source_dir"); err == nil && strings.TrimSpace(v) != "" {
+ return strings.TrimSpace(v)
+ }
+ }
+ return lib.Path
+}
+
+// resolveTransferMode picks the transfer mode: a per-request override wins,
+// otherwise the organize.transfer_mode setting (default move). When the
+// effective mode is move and 做种保种 (organize.keep_seeding) is enabled, it is
+// upgraded to hardlink so the source stays in place for the torrent client.
+func (o *OrganizerService) resolveTransferMode(ctx context.Context, override TransferMode) TransferMode {
+ mode := override
+ if mode == "" {
+ mode = TransferMove
+ if o.repo != nil && o.repo.Setting != nil {
+ if v, err := o.repo.Setting.Get(ctx, "organize.transfer_mode"); err == nil && strings.TrimSpace(v) != "" {
+ mode = parseTransferMode(v)
+ }
+ }
+ }
+ if mode == TransferMove && o.keepSeedingEnabled(ctx) {
+ // 移动会删除源文件导致 qBittorrent 停止做种;保种开启时改用硬链接
+ // 既规范命名又保留源文件继续做种上传。硬链接失败时会报错,避免静默
+ // 退化复制后占用双份磁盘空间。
+ return TransferHardlink
+ }
+ return mode
+}
+
+// keepSeedingEnabled reports whether 做种保种 is on. Defaults to true so an
+// unconfigured instance never silently breaks seeding on organize.
+func (o *OrganizerService) keepSeedingEnabled(ctx context.Context) bool {
+ if o.repo == nil || o.repo.Setting == nil {
+ return true
+ }
+ v, err := o.repo.Setting.Get(ctx, "organize.keep_seeding")
+ if err != nil || strings.TrimSpace(v) == "" {
+ return true
+ }
+ return v == "true" || v == "1" || v == "on"
+}
diff --git a/internal/service/organizer_sidecar.go b/internal/service/organizer_sidecar.go
new file mode 100644
index 0000000..7b8495b
--- /dev/null
+++ b/internal/service/organizer_sidecar.go
@@ -0,0 +1,29 @@
+package service
+
+import (
+ "os"
+ "path/filepath"
+)
+
+// transferSidecarNFO moves/copies/links the .nfo sidecar alongside its media
+// using the same transfer mode, so metadata follows the organized file.
+func transferSidecarNFO(srcMedia, dstMedia string, mode TransferMode) error {
+ src := nfoPath(srcMedia)
+ dst := nfoPath(dstMedia)
+ if src == dst {
+ return nil
+ }
+ if _, err := os.Stat(src); err != nil {
+ if os.IsNotExist(err) {
+ return nil
+ }
+ return err
+ }
+ if _, err := os.Stat(dst); err == nil {
+ return nil
+ }
+ if err := os.MkdirAll(filepath.Dir(dst), 0o755); err != nil { // #nosec G301 -- sidecar media directories must remain readable by NAS/player users.
+ return err
+ }
+ return transferFile(src, dst, mode)
+}
diff --git a/internal/service/organizer_test.go b/internal/service/organizer_test.go
index 0b1f172..641b176 100644
--- a/internal/service/organizer_test.go
+++ b/internal/service/organizer_test.go
@@ -5,9 +5,7 @@ import (
"path/filepath"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -25,13 +23,7 @@ func TestOrganizeMediaReDetectsSeasonFromPath(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "TV", Path: root, Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -86,13 +78,7 @@ func TestOrganizeMediaUsesEpisodeNFOSeason(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "TV", Path: root, Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -134,13 +120,7 @@ func TestOrganizeMediaAddsTypeRootForGenericMediaRoot(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
t.Fatal(err)
@@ -189,13 +169,7 @@ func TestOrganizeMediaDoesNotRepeatCategoryWhenLibraryIsCategoryRoot(t *testing.
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
t.Fatal(err)
@@ -245,13 +219,7 @@ func TestOrganizeLibrarySkipsFilesAlreadyInsideLibrary(t *testing.T) {
t.Fatal(err)
}
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil {
t.Fatal(err)
diff --git a/internal/service/organizer_transfer_test.go b/internal/service/organizer_transfer_test.go
index 82c0857..2327915 100644
--- a/internal/service/organizer_transfer_test.go
+++ b/internal/service/organizer_transfer_test.go
@@ -6,9 +6,7 @@ import (
"strings"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -17,14 +15,7 @@ import (
func newOrganizerTestRepo(t *testing.T) *repository.Container {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- return repository.New(db)
+ return repository.New(newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.AccessLog{}))
}
func TestOrganizeMediaHonorsTargetDirAndCopyMode(t *testing.T) {
diff --git a/internal/service/organizer_types.go b/internal/service/organizer_types.go
new file mode 100644
index 0000000..9bf518f
--- /dev/null
+++ b/internal/service/organizer_types.go
@@ -0,0 +1,51 @@
+package service
+
+// OrganizeResult reports what happened.
+type OrganizeResult struct {
+ Organized int `json:"organized"`
+ Skipped int `json:"skipped"`
+ Replaced int `json:"replaced,omitempty"`
+ Reclassified int `json:"reclassified,omitempty"`
+ Errors []string `json:"errors,omitempty"`
+ SourcePath string `json:"source_path,omitempty"`
+ DestPath string `json:"dest_path,omitempty"`
+ DryRun bool `json:"dry_run,omitempty"`
+ Items []OrganizePreviewItem `json:"items,omitempty"`
+ Scans []OrganizeScanSummary `json:"scans,omitempty"`
+ Scrapes []OrganizeScrapeSummary `json:"scrapes,omitempty"`
+}
+
+type OrganizePreviewItem struct {
+ Source string `json:"source"`
+ Target string `json:"target,omitempty"`
+ Action string `json:"action"` // organize / skip / replace / reclassify / cleanup / error
+ Reason string `json:"reason,omitempty"`
+ MediaType string `json:"media_type,omitempty"`
+ Category string `json:"category,omitempty"`
+ Title string `json:"title,omitempty"`
+}
+
+// OrganizeOptions carries per-request overrides for an organize operation.
+// Empty values use system defaults.
+//
+// 整理是「从源目录整理到目的地目录」:SourcePath 指定待整理文件所在的源目录,
+// DestPath 指定整理输出的目的地目录。两者相互独立,不再混用同一个目录。
+type OrganizeOptions struct {
+ // SourcePath 本次整理的源目录(待整理文件所在目录),覆盖 organize.source_dir
+ // 设置与媒体库路径。仅整理位于该目录下的媒体;留空表示整个媒体库。
+ SourcePath string
+ // DestPath 本次整理的目的地根路径(整理输出到哪里),覆盖 organize.target_dir 设置。
+ // 留空则使用设置中的默认目的地目录,再退回媒体库路径。
+ DestPath string
+ // TransferMode 本次整理的转移方式,覆盖 organize.transfer_mode 设置。
+ TransferMode TransferMode
+ // MediaType 手动整理时由 UI 指定的媒体类型。空值时按文件名/目录推断。
+ MediaType string
+ // MediaCategory 由订阅/下载任务或 UI 指定的分类。空值时按目录/NFO/规则推断。
+ MediaCategory string
+ // DryRun 仅生成整理预览,不实际移动/复制/硬链接文件。
+ DryRun bool
+ // AllowReplaceExisting 允许用本次来源替换目标库中已存在的同一媒体。
+ // 默认 false:只去重不洗版,避免未开启洗版的订阅/手动整理留下或替换出多份版本。
+ AllowReplaceExisting bool
+}
diff --git a/internal/service/play_profile_test.go b/internal/service/play_profile_test.go
index c6d438a..ed9dd2a 100644
--- a/internal/service/play_profile_test.go
+++ b/internal/service/play_profile_test.go
@@ -6,20 +6,12 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
)
func newPlayProfileTestService(t *testing.T) *PlayProfileService {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.PlayProfile{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.PlayProfile{})
return NewPlayProfileService(zap.NewNop(), repository.New(db))
}
diff --git a/internal/service/scanner.go b/internal/service/scanner.go
index 12b609f..88ff870 100644
--- a/internal/service/scanner.go
+++ b/internal/service/scanner.go
@@ -13,11 +13,6 @@ package service
import (
"context"
"errors"
- "fmt"
- "net/url"
- "os"
- "path/filepath"
- "sort"
"strings"
"sync"
"time"
@@ -25,7 +20,6 @@ import (
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
- "github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
@@ -148,134 +142,6 @@ func (s *ScannerService) SetImageProxy(imageProxy *ImageProxy) {
}
}
-func (s *ScannerService) cloudImagePrefetchWorker() {
- for task := range s.cloudImagePrefetchQueue {
- s.prefetchCloudImage(task)
- }
-}
-
-func (s *ScannerService) queueCloudArtworkPrefetch(raw string) {
- if s == nil || s.storage == nil || s.imageProxy == nil {
- return
- }
- typ, ref, ok := parseCloudImagePlaybackURL(raw)
- if !ok {
- return
- }
- stableKey := typ + ":" + ref
- if s.imageProxy.CloudImageCached(stableKey) {
- return
- }
- s.cloudImagePrefetchMu.Lock()
- if _, ok := s.cloudImagePrefetching[stableKey]; ok {
- s.cloudImagePrefetchMu.Unlock()
- return
- }
- s.cloudImagePrefetching[stableKey] = struct{}{}
- s.cloudImagePrefetchMu.Unlock()
-
- task := cloudImagePrefetchTask{typ: typ, ref: ref, stableKey: stableKey}
- select {
- case s.cloudImagePrefetchQueue <- task:
- default:
- s.cloudImagePrefetchMu.Lock()
- delete(s.cloudImagePrefetching, stableKey)
- s.cloudImagePrefetchMu.Unlock()
- if s.log != nil {
- s.log.Debug("cloud artwork prefetch queue full", zap.String("provider", typ), zap.String("ref", ref))
- }
- }
-}
-
-func (s *ScannerService) prefetchCloudImage(task cloudImagePrefetchTask) {
- defer func() {
- s.cloudImagePrefetchMu.Lock()
- delete(s.cloudImagePrefetching, task.stableKey)
- s.cloudImagePrefetchMu.Unlock()
- }()
- if s == nil || s.storage == nil || s.imageProxy == nil || s.imageProxy.CloudImageCached(task.stableKey) {
- return
- }
- ctx, cancel := context.WithTimeout(context.Background(), 45*time.Second)
- defer cancel()
- link, err := s.storage.CloudResolve(ctx, task.typ, task.ref, "")
- if err != nil {
- if s.log != nil {
- s.log.Debug("resolve cloud artwork for prefetch failed", zap.String("provider", task.typ), zap.String("ref", task.ref), zap.Error(err))
- }
- return
- }
- if err := s.imageProxy.PrefetchCloudResolved(ctx, task.stableKey, link); err != nil && s.log != nil {
- s.log.Debug("prefetch cloud artwork failed", zap.String("provider", task.typ), zap.String("ref", task.ref), zap.Error(err))
- }
-}
-
-func (s *ScannerService) cacheCloudArtworkNow(ctx context.Context, raw string) {
- if s == nil || s.storage == nil || s.imageProxy == nil {
- return
- }
- typ, ref, ok := parseCloudImagePlaybackURL(raw)
- if !ok {
- return
- }
- stableKey := typ + ":" + ref
- if s.imageProxy.CloudImageCached(stableKey) {
- return
- }
- cacheCtx, cancel := context.WithTimeout(ctx, 20*time.Second)
- defer cancel()
- link, err := s.storage.CloudResolve(cacheCtx, typ, ref, "")
- if err != nil {
- if s.log != nil {
- s.log.Debug("resolve cloud artwork for priority cache failed", zap.String("provider", typ), zap.String("ref", ref), zap.Error(err))
- }
- s.queueCloudArtworkPrefetch(raw)
- return
- }
- if err := s.imageProxy.PrefetchCloudResolved(cacheCtx, stableKey, link); err != nil {
- if s.log != nil {
- s.log.Debug("priority cache cloud artwork failed", zap.String("provider", typ), zap.String("ref", ref), zap.Error(err))
- }
- s.queueCloudArtworkPrefetch(raw)
- }
-}
-
-func (s *ScannerService) cacheCloudMetadataArtworkNow(ctx context.Context, meta *LocalMetadata) {
- if meta == nil {
- return
- }
- s.cacheCloudArtworkNow(ctx, meta.PosterURL)
- s.cacheCloudArtworkNow(ctx, meta.BackdropURL)
-}
-
-func parseCloudImagePlaybackURL(raw string) (string, string, bool) {
- u, err := url.Parse(strings.TrimSpace(raw))
- if err != nil {
- return "", "", false
- }
- path := strings.Trim(u.Path, "/")
- const prefix = "api/cloud/play/"
- if !strings.HasPrefix(strings.ToLower(path), prefix) {
- return "", "", false
- }
- typ := strings.TrimSpace(path[len(prefix):])
- ref := strings.TrimSpace(u.Query().Get("ref"))
- if typ == "" || ref == "" || !isCloudArtworkRef(ref) {
- return "", "", false
- }
- return typ, ref, true
-}
-
-func isCloudArtworkRef(ref string) bool {
- ref = strings.ToLower(strings.TrimSpace(ref))
- for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp"} {
- if strings.HasSuffix(ref, suffix) {
- return true
- }
- }
- return false
-}
-
// ScanResult summarises a scan run.
type ScanResult struct {
LibraryID string `json:"library_id"`
@@ -311,7 +177,7 @@ func addScanError(res *ScanResult, path string, err error) {
res.Errors = append(res.Errors, msg)
}
-const maxCloudMediaProbeQueuePerScan = 32
+const maxCloudMediaProbeQueuePerScan = 256
const cloudMediaProbeFailureBackoff = 6 * time.Hour
@@ -348,12 +214,6 @@ type cloudScanEntry struct {
cancel context.CancelFunc
}
-type cloudImagePrefetchTask struct {
- typ string
- ref string
- stableKey string
-}
-
type cloudMediaProbeTask struct {
typ string
ref string
@@ -365,2107 +225,63 @@ type localMediaProbeTask struct {
}
type existingCloudMedia struct {
- SizeBytes int64
- DurationSec int
- Width int
- Height int
- VideoCodec string
- AudioCodec string
- Container string
- PosterURL string
- BackdropURL string
- STRMURL string
- Year int
- TMDbID int
- BangumiID int
- DoubanID string
- TheTVDBID string
+ LibraryID string
+ Title string
+ OriginalName string
+ EpisodeTitle string
+ SizeBytes int64
+ DurationSec int
+ Width int
+ Height int
+ VideoCodec string
+ AudioCodec string
+ Container string
+ PosterURL string
+ BackdropURL string
+ STRMURL string
+ Overview string
+ Year int
+ Rating float32
+ TMDbID int
+ BangumiID int
+ DoubanID string
+ TheTVDBID string
+ SeasonNum int
+ EpisodeNum int
+ Genres string
+ Countries string
+ Languages string
+ NSFW bool
+ ScrapeStatus string
}
type existingLocalMedia struct {
- SizeBytes int64
- DurationSec int
- Width int
- Height int
- VideoCodec string
- AudioCodec string
- Container string
- STRMURL string
- FileID string
-}
-
-func (s *ScannerService) cloudMediaProbeWorker() {
- for task := range s.cloudMediaProbeQueue {
- s.probeCloudMediaAsync(task)
- }
-}
-
-func (s *ScannerService) queueCloudMediaProbe(typ, ref, path string) bool {
- if s == nil || s.storage == nil || s.probe == nil {
- return false
- }
- typ = strings.TrimSpace(typ)
- ref = strings.TrimSpace(ref)
- path = strings.TrimSpace(path)
- if typ == "" || ref == "" || path == "" {
- return false
- }
- s.cloudMediaProbeMu.Lock()
- if until, ok := s.cloudMediaProbeBackoff[path]; ok {
- if time.Now().Before(until) {
- s.cloudMediaProbeMu.Unlock()
- return false
- }
- delete(s.cloudMediaProbeBackoff, path)
- }
- if _, ok := s.cloudMediaProbing[path]; ok {
- s.cloudMediaProbeMu.Unlock()
- return false
- }
- s.cloudMediaProbing[path] = struct{}{}
- s.cloudMediaProbeMu.Unlock()
-
- task := cloudMediaProbeTask{typ: typ, ref: ref, path: path}
- select {
- case s.cloudMediaProbeQueue <- task:
- return true
- default:
- s.cloudMediaProbeMu.Lock()
- delete(s.cloudMediaProbing, path)
- // 队列满说明探测工人已饱和;给该文件挂一个短退避,避免下一轮
- // 扫描立刻重复尝试同一批文件。
- if s.cloudMediaProbeBackoff == nil {
- s.cloudMediaProbeBackoff = make(map[string]time.Time)
- }
- s.cloudMediaProbeBackoff[path] = time.Now().Add(cloudMediaProbeQueueFullBackoff)
- s.cloudMediaProbeMu.Unlock()
- if s.log != nil {
- // 限速告警:队列满在大库扫描中是常态而非异常,逐条 WARN 会
- // 在几小时内产生数万行日志(真实环境出现过 41165 条)。
- now := time.Now()
- s.cloudMediaProbeWarnMu.Lock()
- shouldWarn := now.Sub(s.cloudMediaProbeLastWarn) >= time.Minute
- if shouldWarn {
- s.cloudMediaProbeLastWarn = now
- }
- s.cloudMediaProbeWarnMu.Unlock()
- if shouldWarn {
- s.log.Warn("cloud media probe queue full; deferring remaining probes (logged at most once per minute)",
- zap.String("provider", typ), zap.String("path", path))
- } else {
- s.log.Debug("cloud media probe queue full", zap.String("provider", typ), zap.String("path", path))
- }
- }
- return false
- }
-}
-
-func (s *ScannerService) queueCloudMediaProbeWithBudget(typ, ref, path string, budget *int) bool {
- if budget != nil {
- if *budget <= 0 {
- return false
- }
- // 预算按「尝试」扣减而不是按「成功入队」扣减。否则当探测队列被
- // 其他扫描填满时,本次扫描会对剩下的每一个文件都尝试入队并各打
- // 一条日志——真实环境里曾因此产生过 4 万多条 "queue full" WARN,
- // 这本身就是一笔可观的 CPU/磁盘开销。
- *budget--
- }
- return s.queueCloudMediaProbe(typ, ref, path)
-}
-
-func (s *ScannerService) localMediaProbeWorker() {
- for task := range s.localMediaProbeQueue {
- s.probeLocalMediaAsync(task)
- }
-}
-
-func (s *ScannerService) queueLocalMediaProbe(path string) bool {
- if s == nil || s.probe == nil {
- return false
- }
- path = strings.TrimSpace(path)
- if path == "" {
- return false
- }
- s.localMediaProbeOnce.Do(func() {
- workers := s.ffprobeWorkerCount()
- for i := 0; i < workers; i++ {
- go s.localMediaProbeWorker()
- }
- })
- s.localMediaProbeMu.Lock()
- if s.localMediaProbing == nil {
- s.localMediaProbing = make(map[string]struct{})
- }
- if _, ok := s.localMediaProbing[path]; ok {
- s.localMediaProbeMu.Unlock()
- return false
- }
- s.localMediaProbing[path] = struct{}{}
- s.localMediaProbeMu.Unlock()
-
- task := localMediaProbeTask{path: path}
- select {
- case s.localMediaProbeQueue <- task:
- return true
- default:
- s.localMediaProbeMu.Lock()
- delete(s.localMediaProbing, path)
- s.localMediaProbeMu.Unlock()
- if s.log != nil {
- s.log.Debug("local media probe queue full", zap.String("path", path))
- }
- return false
- }
-}
-
-func (s *ScannerService) beginCloudScan(ctx context.Context, lib *model.Library, mount CloudMountInfo) (context.Context, func(*ScanResult, error), error) {
- if s == nil || lib == nil {
- return ctx, func(*ScanResult, error) {}, nil
- }
- s.cloudScanMu.Lock()
- if s.cloudScans == nil {
- s.cloudScans = make(map[string]*cloudScanEntry)
- }
- if entry := s.cloudScans[lib.ID]; entry != nil && (entry.status.State == "running" || entry.status.State == "canceling") {
- s.cloudScanMu.Unlock()
- return ctx, nil, ErrCloudScanAlreadyRunning
- }
- runCtx, cancel := context.WithCancel(ctx)
- now := time.Now()
- entry := &cloudScanEntry{
- status: CloudScanStatus{
- LibraryID: lib.ID,
- Provider: mount.Provider,
- Stage: "listing",
- State: "running",
- StartedAt: now,
- UpdatedAt: now,
- ResumeHint: "中断后再次点击扫描会从头遍历,但已入库媒体会去重更新,只补齐缺失项。",
- Estimate: "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度。",
- },
- cancel: cancel,
- }
- s.cloudScans[lib.ID] = entry
- s.cloudScanMu.Unlock()
-
- finish := func(res *ScanResult, err error) {
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- current := s.cloudScans[lib.ID]
- if current == nil {
- return
- }
- now := time.Now()
- if res != nil {
- current.status.Visited = res.Visited
- current.status.Added = res.Added
- current.status.Updated = res.Updated
- current.status.Skipped = res.Skipped
- current.status.Removed = res.Removed
- current.status.ErrorCount = res.ErrorCount
- current.status.Errors = append([]string(nil), res.Errors...)
- }
- current.status.UpdatedAt = now
- current.status.FinishedAt = now
- current.cancel = nil
- switch {
- case errors.Is(err, context.Canceled):
- current.status.State = "canceled"
- current.status.Stage = "canceled"
- current.status.Error = ""
- case errors.Is(err, context.DeadlineExceeded):
- current.status.State = "error"
- current.status.Stage = "error"
- current.status.Error = "扫描超时:" + err.Error()
- case err != nil:
- current.status.State = "error"
- current.status.Stage = "error"
- current.status.Error = err.Error()
- default:
- current.status.State = "finished"
- current.status.Stage = "finished"
- if current.status.ErrorCount > 0 {
- current.status.Error = fmt.Sprintf("部分文件入库失败:%d 个,详情见 errors", current.status.ErrorCount)
- } else {
- current.status.Error = ""
- }
- }
- if s.hub != nil {
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "provider": mount.Provider,
- "cloud": true,
- "finished": true,
- "state": current.status.State,
- "stage": current.status.Stage,
- "error": current.status.Error,
- "visited": current.status.Visited,
- "added": current.status.Added,
- "updated": current.status.Updated,
- "skipped": current.status.Skipped,
- "removed": current.status.Removed,
- "error_count": current.status.ErrorCount,
- "errors": current.status.Errors,
- })
- }
- s.notifyScanFinished(lib, res, err, true)
- }
- return runCtx, finish, nil
-}
-
-func (s *ScannerService) updateCloudScanProgress(libraryID, stage string, dirs, discovered, visited, added, updated, skipped int, removed int64, filesPerSecond float64) {
- if s == nil {
- return
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- entry := s.cloudScans[libraryID]
- if entry == nil {
- return
- }
- entry.status.Stage = stage
- entry.status.UpdatedAt = time.Now()
- entry.status.Dirs = dirs
- entry.status.Discovered = discovered
- entry.status.Visited = visited
- entry.status.Added = added
- entry.status.Updated = updated
- entry.status.Skipped = skipped
- entry.status.Removed = removed
- entry.status.FilesPerSecond = filesPerSecond
-}
-
-func (s *ScannerService) acquireCloudScanSlot(ctx context.Context, libraryID string) (func(), error) {
- if s == nil {
- return func() {}, nil
- }
- s.cloudScanMu.Lock()
- if s.cloudSlots == nil {
- s.cloudSlots = make(chan struct{}, 1)
- }
- slots := s.cloudSlots
- if entry := s.cloudScans[libraryID]; entry != nil {
- entry.status.Stage = "queued"
- entry.status.UpdatedAt = time.Now()
- }
- s.cloudScanMu.Unlock()
-
- select {
- case slots <- struct{}{}:
- s.cloudScanMu.Lock()
- if entry := s.cloudScans[libraryID]; entry != nil && entry.status.State == "running" {
- entry.status.Stage = "listing"
- entry.status.UpdatedAt = time.Now()
- }
- s.cloudScanMu.Unlock()
- return func() { <-slots }, nil
- case <-ctx.Done():
- return nil, ctx.Err()
- }
-}
-
-// CloudScanStatuses returns the current or most recent status per cloud library.
-func (s *ScannerService) CloudScanStatuses() []CloudScanStatus {
- if s == nil {
- return nil
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- out := make([]CloudScanStatus, 0, len(s.cloudScans))
- for _, entry := range s.cloudScans {
- out = append(out, entry.status)
- }
- return out
-}
-
-func (s *ScannerService) CancelCloudScan(libraryID string) bool {
- if s == nil || strings.TrimSpace(libraryID) == "" {
- return false
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- entry := s.cloudScans[libraryID]
- if entry == nil || (entry.status.State != "running" && entry.status.State != "queued" && entry.status.State != "canceling") {
- return false
- }
- entry.status.State = "canceling"
- entry.status.Stage = "canceling"
- entry.status.UpdatedAt = time.Now()
- if entry.cancel != nil {
- entry.cancel()
- } else {
- entry.status.State = "canceled"
- entry.status.Stage = "canceled"
- entry.status.FinishedAt = time.Now()
- }
- return true
-}
-
-func (s *ScannerService) CancelAllCloudScans() int {
- if s == nil {
- return 0
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- cancelled := 0
- for _, entry := range s.cloudScans {
- if entry == nil || (entry.status.State != "running" && entry.status.State != "queued" && entry.status.State != "canceling") {
- continue
- }
- entry.status.State = "canceling"
- entry.status.Stage = "canceling"
- entry.status.UpdatedAt = time.Now()
- if entry.cancel != nil {
- entry.cancel()
- } else {
- entry.status.State = "canceled"
- entry.status.Stage = "canceled"
- entry.status.FinishedAt = time.Now()
- }
- cancelled++
- }
- return cancelled
-}
-
-func (s *ScannerService) CancelCloudScansForProvider(provider string) int {
- if s == nil {
- return 0
- }
- provider = strings.TrimSpace(provider)
- if provider == "" {
- return 0
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- cancelled := 0
- for _, entry := range s.cloudScans {
- if entry == nil || entry.status.Provider != provider || (entry.status.State != "running" && entry.status.State != "queued" && entry.status.State != "canceling") {
- continue
- }
- entry.status.State = "canceling"
- entry.status.Stage = "canceling"
- entry.status.UpdatedAt = time.Now()
- if entry.cancel != nil {
- entry.cancel()
- } else {
- entry.status.State = "canceled"
- entry.status.Stage = "canceled"
- entry.status.FinishedAt = time.Now()
- }
- cancelled++
- }
- return cancelled
-}
-
-func (s *ScannerService) StartCloudLibraryScan(libraryID string, autoScrape bool) (CloudScanStatus, bool, error) {
- if s == nil {
- return CloudScanStatus{}, false, errors.New("scanner unavailable")
- }
- lib, err := s.repo.Library.FindByID(context.Background(), libraryID)
- if err != nil {
- return CloudScanStatus{}, false, err
- }
- if lib == nil {
- return CloudScanStatus{}, false, errors.New("library not found")
- }
- mount, ok := ParseCloudLibraryMount(lib.Path)
- if !ok {
- return CloudScanStatus{}, false, errors.New("library is not a cloud mount")
- }
- s.cloudScanMu.Lock()
- if entry := s.cloudScans[libraryID]; entry != nil && (entry.status.State == "running" || entry.status.State == "canceling") {
- status := entry.status
- s.cloudScanMu.Unlock()
- return status, false, nil
- }
- s.cloudScanMu.Unlock()
-
- go func() {
- ctx, cancel := cloudScanContext(context.Background(), cloudScanTimeout(context.Background(), s.repo, 24*time.Hour))
- defer cancel()
- if autoScrape {
- _, err = s.ScanLibrary(ctx, libraryID)
- } else {
- _, err = s.ScanLibraryWithoutAutoScrape(ctx, libraryID)
- }
- if err != nil && !errors.Is(err, ErrCloudScanAlreadyRunning) && s.log != nil {
- s.log.Warn("cloud library background scan failed", zap.String("library_id", libraryID), zap.Error(err))
- }
- }()
- status := CloudScanStatus{
- LibraryID: libraryID,
- Provider: mount.Provider,
- Stage: "queued",
- State: "queued",
- StartedAt: time.Now(),
- UpdatedAt: time.Now(),
- ResumeHint: "中断后再次点击扫描会从头遍历,但已入库媒体会去重更新,只补齐缺失项。",
- Estimate: "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度。",
- }
- return status, true, nil
-}
-
-func cloudScanContext(parent context.Context, timeout time.Duration) (context.Context, context.CancelFunc) {
- if timeout <= 0 {
- return context.WithCancel(parent)
- }
- return context.WithTimeout(parent, timeout)
-}
-
-func cloudScanTimeout(ctx context.Context, repo *repository.Container, fallback time.Duration) time.Duration {
- if repo == nil || repo.Setting == nil {
- return fallback
- }
- value, err := repo.Setting.Get(ctx, "cloud.scan_timeout_hours")
- if err != nil || strings.TrimSpace(value) == "" {
- return fallback
- }
- hours := parseIntSettingDefault(strings.TrimSpace(value), int(fallback/time.Hour))
- if hours <= 0 {
- return 0
- }
- return time.Duration(hours) * time.Hour
-}
-
-func (s *ScannerService) StartAllCloudLibraryScans() ([]CloudScanStatus, error) {
- if s == nil {
- return nil, errors.New("scanner unavailable")
- }
- libs, err := s.repo.Library.List(context.Background())
- if err != nil {
- return nil, err
- }
- libs = FilterScannableCloudLibraries(context.Background(), s.repo, libs)
- statuses := make([]CloudScanStatus, 0, len(libs))
- queue := make([]string, 0, len(libs))
- for _, lib := range libs {
- if !lib.Enabled {
- continue
- }
- mount, ok := ParseCloudLibraryMount(lib.Path)
- if !ok {
- continue
- }
- status, queued := s.queueCloudLibraryScan(lib, mount)
- if queued {
- queue = append(queue, lib.ID)
- }
- statuses = append(statuses, status)
- }
- if len(queue) > 0 {
- go s.runQueuedCloudLibraryScans(queue)
- }
- return statuses, nil
-}
-
-func (s *ScannerService) queueCloudLibraryScan(lib model.Library, mount CloudMountInfo) (CloudScanStatus, bool) {
- now := time.Now()
- status := CloudScanStatus{
- LibraryID: lib.ID,
- Provider: mount.Provider,
- Stage: "queued",
- State: "queued",
- StartedAt: now,
- UpdatedAt: now,
- ResumeHint: "中断后再次点击扫描会从头遍历,但已入库媒体会去重更新,只补齐缺失项。",
- Estimate: "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度。",
- }
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- if s.cloudScans == nil {
- s.cloudScans = make(map[string]*cloudScanEntry)
- }
- if entry := s.cloudScans[lib.ID]; entry != nil {
- switch entry.status.State {
- case "running", "queued", "canceling":
- return entry.status, false
- }
- }
- s.cloudScans[lib.ID] = &cloudScanEntry{status: status}
- return status, true
-}
-
-func (s *ScannerService) runQueuedCloudLibraryScans(libraryIDs []string) {
- ctx, cancel := cloudScanContext(context.Background(), cloudScanTimeout(context.Background(), s.repo, 24*time.Hour))
- defer cancel()
- for _, libraryID := range libraryIDs {
- if ctx.Err() != nil {
- return
- }
- if s.cloudScanWasCanceled(libraryID) {
- continue
- }
- if _, err := s.ScanLibraryWithoutAutoScrape(ctx, libraryID); err != nil && !errors.Is(err, ErrCloudScanAlreadyRunning) && !errors.Is(err, context.Canceled) && s.log != nil {
- s.log.Warn("cloud library queued scan failed", zap.String("library_id", libraryID), zap.Error(err))
- }
- }
-}
-
-func (s *ScannerService) cloudScanWasCanceled(libraryID string) bool {
- s.cloudScanMu.Lock()
- defer s.cloudScanMu.Unlock()
- entry := s.cloudScans[libraryID]
- return entry != nil && entry.status.State == "canceled"
-}
-
-// ScanLibrary walks the library root and persists discovered media files.
-func (s *ScannerService) ScanLibrary(ctx context.Context, libraryID string) (*ScanResult, error) {
- return s.scanLibrary(ctx, libraryID, true)
-}
-
-// ScanLibraryWithoutAutoScrape walks a library without kicking off online
-// metadata enrichment. Cloud mounts can contain very large trees; keeping mount
-// scans import-only prevents scraper bursts from overwhelming small NAS boxes.
-func (s *ScannerService) ScanLibraryWithoutAutoScrape(ctx context.Context, libraryID string) (*ScanResult, error) {
- return s.scanLibrary(ctx, libraryID, false)
-}
-
-func (s *ScannerService) TryBeginLocalScan(libraryID string) (func(), bool) {
- if s == nil || strings.TrimSpace(libraryID) == "" {
- return func() {}, true
- }
- s.localScanMu.Lock()
- if s.localScans == nil {
- s.localScans = make(map[string]struct{})
- }
- if _, ok := s.localScans[libraryID]; ok {
- s.localScanMu.Unlock()
- return nil, false
- }
- s.localScans[libraryID] = struct{}{}
- s.localScanMu.Unlock()
- return func() {
- s.localScanMu.Lock()
- delete(s.localScans, libraryID)
- s.localScanMu.Unlock()
- }, true
-}
-
-func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, autoScrape bool) (*ScanResult, error) {
- lib, err := s.repo.Library.FindByID(ctx, libraryID)
- if err != nil || lib == nil {
- return nil, err
- }
- if mount, ok := ParseCloudLibraryMount(lib.Path); ok {
- if shadow := s.shadowedCloudLibrary(ctx, lib); shadow != nil {
- res := &ScanResult{LibraryID: lib.ID, Skipped: 1}
- s.log.Warn("skip shadowed cloud library scan",
- zap.String("library_id", lib.ID),
- zap.String("shadowed_by", shadow.Library.ID),
- zap.String("provider", mount.Provider))
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "finished": true,
- "skipped": res.Skipped,
- "cloud": true,
- "shadowed": true,
- })
- return res, nil
- }
- scanCtx, finish, err := s.beginCloudScan(ctx, lib, mount)
- if err != nil {
- if errors.Is(err, ErrCloudScanAlreadyRunning) {
- return &ScanResult{LibraryID: lib.ID, Skipped: 1}, nil
- }
- return nil, err
- }
- release, err := s.acquireCloudScanSlot(scanCtx, lib.ID)
- if err != nil {
- res := &ScanResult{LibraryID: lib.ID}
- if finish != nil {
- finish(res, err)
- }
- return res, err
- }
- defer release()
- res, err := s.scanCloudLibrary(scanCtx, lib, mount, autoScrape)
- if finish != nil {
- finish(res, err)
- }
- return res, err
- }
- if err := s.resolveLocalLibraryPath(ctx, lib); err != nil {
- return &ScanResult{LibraryID: lib.ID}, err
- }
- res := &ScanResult{LibraryID: lib.ID}
- seen := make(map[string]struct{})
- seenInodes := make(map[string]string)
- writeBatch := newLocalMediaWriteBatch(s, ctx, res, 100)
- existingMedia, err := s.existingLocalMediaSnapshot(ctx, lib.ID)
- if err != nil {
- s.log.Warn("load existing local media snapshot failed", zap.String("library_id", lib.ID), zap.Error(err))
- existingMedia = nil
- } else {
- for path, existing := range existingMedia {
- if existing.FileID != "" {
- seenInodes[existing.FileID] = path
- }
- }
- }
-
- walkFn := func(path string, info walkInfo) error {
- select {
- case <-ctx.Done():
- return ctx.Err()
- default:
- }
- if info.isDir {
- return nil
- }
- ext := strings.ToLower(filepath.Ext(path))
- if _, ok := videoExtensions[ext]; !ok {
- return nil
- }
- seen[filepath.Clean(path)] = struct{}{}
- s.ingestFile(ctx, lib, path, info.size, seenInodes, existingMedia, writeBatch, res)
- return nil
- }
-
- walkErr := walk(lib.Path, walkFn)
- writeBatch.Flush()
- if walkErr != nil {
- addScanError(res, lib.Path, walkErr)
- if res.Added+res.Updated > 0 {
- s.invalidateMediaCache(ctx)
- }
- return res, walkErr
- }
- removed, err := s.pruneMissingMedia(ctx, lib.ID, seen)
- if err != nil {
- s.log.Warn("prune missing media failed", zap.String("library_id", lib.ID), zap.Error(err))
- } else {
- res.Removed = removed
- }
-
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "finished": true,
- "visited": res.Visited,
- "added": res.Added,
- "updated": res.Updated,
- "probed": res.Probed,
- "local_meta": res.LocalMetadata,
- "removed": res.Removed,
- "error_count": res.ErrorCount,
- "errors": res.Errors,
- })
- s.notifyScanFinished(lib, res, nil, false)
- s.invalidateMediaCache(ctx)
- s.maybeGenerateSTRMAfterScan(lib.ID)
-
- // Online enrichment is opt-in. Local NFO is always consumed first during
- // the scan, and matched rows are excluded from EnrichLibrary's pending set.
- if autoScrape && s.scraper != nil && s.scraper.AnyEnabled() && s.autoScrapeEnabled(ctx) {
- s.startAutoScrape(ctx, lib.ID)
- }
- return res, nil
-}
-
-func (s *ScannerService) notifyScanFinished(lib *model.Library, res *ScanResult, err error, cloud bool) {
- if s == nil || s.notify == nil || lib == nil || res == nil {
- return
- }
- if err != nil {
- go func() {
- ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
- defer cancel()
- s.notify.Broadcast(ctx, "MediaStationGo 扫描异常", fmt.Sprintf("媒体库:%s\n错误:%s", lib.Name, err.Error()), EventSystemAlert)
- }()
- return
- }
- if res.Added+res.Updated <= 0 {
- return
- }
- source := "本地媒体库"
- if cloud {
- source = "网盘媒体库"
- }
- body := fmt.Sprintf("%s:%s\n新增:%d\n更新:%d\n跳过:%d\n移除:%d", source, lib.Name, res.Added, res.Updated, res.Skipped, res.Removed)
- go func() {
- ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
- defer cancel()
- s.notify.Broadcast(ctx, "MediaStationGo 入库完成", body, EventLibraryIngest)
- }()
-}
-
-// IngestPath ingests a single file into the given library without walking the
-// whole tree. Used by the watcher for incremental, event-driven additions so
-// adding one new file no longer triggers a full library re-scan (减少硬盘损耗).
-// Non-video files and directories are ignored. Returns true if a media row was
-// added or updated.
-func (s *ScannerService) IngestPath(ctx context.Context, libraryID, path string) (bool, error) {
- lib, err := s.repo.Library.FindByID(ctx, libraryID)
- if err != nil || lib == nil {
- return false, err
- }
- if err := s.resolveLocalLibraryPath(ctx, lib); err != nil {
- return false, err
- }
- fi, err := os.Stat(path)
- if err != nil || fi.IsDir() {
- return false, err
- }
- ext := strings.ToLower(filepath.Ext(path))
- if _, ok := videoExtensions[ext]; !ok {
- return false, nil
- }
- res := &ScanResult{LibraryID: lib.ID}
- s.ingestFile(ctx, lib, path, fi.Size(), make(map[string]string), nil, nil, res)
- if res.Added+res.Updated > 0 {
- s.invalidateMediaCache(ctx)
- }
- return res.Added+res.Updated > 0, nil
-}
-
-func (s *ScannerService) resolveLocalLibraryPath(ctx context.Context, lib *model.Library) error {
- if lib == nil || strings.TrimSpace(lib.Path) == "" {
- return nil
- }
- resolved, err := resolveAccessibleLibraryPath(lib.Path)
- if err != nil {
- return err
- }
- if sameLibraryPath(resolved, lib.Path) {
- lib.Path = filepath.Clean(lib.Path)
- return nil
- }
- if s.repo != nil && s.repo.DB != nil {
- if updateErr := s.repo.DB.WithContext(ctx).Model(&model.Library{}).Where("id = ?", lib.ID).Update("path", resolved).Error; updateErr != nil && s.log != nil {
- s.log.Warn("update mapped library path failed",
- zap.String("library_id", lib.ID),
- zap.String("from", lib.Path),
- zap.String("to", resolved),
- zap.Error(updateErr))
- }
- }
- if s.log != nil {
- s.log.Info("mapped library path for scan",
- zap.String("library_id", lib.ID),
- zap.String("from", lib.Path),
- zap.String("to", resolved))
- }
- lib.Path = resolved
- return nil
-}
-
-func (s *ScannerService) scanCloudLibrary(ctx context.Context, lib *model.Library, mount CloudMountInfo, autoScrape bool) (*ScanResult, error) {
- res := &ScanResult{LibraryID: lib.ID}
- if s.storage == nil {
- return res, fmt.Errorf("cloud storage service unavailable")
- }
-
- // 验证存储配置是否存在且已启用
- cfg, err := s.repo.StorageConfig.Get(ctx, mount.Provider)
- if err != nil || cfg == nil {
- return res, fmt.Errorf("storage config not found: %s", mount.Provider)
- }
- if !cfg.Enabled {
- return res, fmt.Errorf("storage %s is disabled", mount.Provider)
- }
- typ := mount.Provider
- rootDir := mount.ScanDir
- rootDisplayDir := mount.DisplayDir
- type cloudCandidate struct {
- ref string
- name string
- size int64
- path string
- localMeta *LocalMetadata
- }
- seen := make(map[string]struct{})
- seenRefs := make(map[string]struct{})
- candidates := make([]cloudCandidate, 0, 256)
- candidateByKey := make(map[string]int)
- visitedDirs := map[string]struct{}{}
- startedAt := time.Now()
- lastProgress := time.Time{}
- dirsVisited := 0
- filesDiscovered := 0
- var stateMu sync.Mutex
- publishProgress := func(stage string, force bool) {
- if s.hub == nil {
- return
- }
- stateMu.Lock()
- if !force && time.Since(lastProgress) < 2*time.Second {
- stateMu.Unlock()
- return
- }
- lastProgress = time.Now()
- dirs := dirsVisited
- discovered := filesDiscovered
- visited := res.Visited
- added := res.Added
- updated := res.Updated
- skipped := res.Skipped
- removed := res.Removed
- stateMu.Unlock()
- elapsed := time.Since(startedAt)
- filesPerSecond := 0.0
- processed := discovered
- if visited > processed {
- processed = visited
- }
- if elapsed.Seconds() > 0 {
- filesPerSecond = float64(processed) / elapsed.Seconds()
- }
- s.updateCloudScanProgress(lib.ID, stage, dirs, discovered, visited, added, updated, skipped, removed, filesPerSecond)
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "cloud": true,
- "stage": stage,
- "dirs": dirs,
- "discovered": discovered,
- "visited": visited,
- "added": added,
- "updated": updated,
- "skipped": skipped,
- "elapsed_seconds": int(elapsed.Seconds()),
- "files_per_second": filesPerSecond,
- "estimate_message": "云盘接口不提供总文件数,剩余时间会随目录大小和网盘响应速度变化",
- })
- }
- publishProgress("listing", true)
- var walkWG sync.WaitGroup
- var walkErr error
- var walkErrOnce sync.Once
- setWalkErr := func(err error) {
- if err != nil {
- walkErrOnce.Do(func() {
- walkErr = err
- })
- }
- }
- listSlots := make(chan struct{}, s.cloudScanWorkerCount())
- var walkCloud func(dirID, displayDir string, inheritedMeta *LocalMetadata) error
- walkCloud = func(dirID, displayDir string, inheritedMeta *LocalMetadata) error {
- defer walkWG.Done()
- if err := ctx.Err(); err != nil {
- setWalkErr(err)
- return err
- }
- stateMu.Lock()
- if _, ok := visitedDirs[dirID]; ok {
- stateMu.Unlock()
- return nil
- }
- visitedDirs[dirID] = struct{}{}
- stateMu.Unlock()
-
- select {
- case listSlots <- struct{}{}:
- defer func() { <-listSlots }()
- case <-ctx.Done():
- setWalkErr(ctx.Err())
- return ctx.Err()
- }
- entries, err := s.storage.CloudList(ctx, typ, dirID)
- if err != nil {
- if dirID != rootDir {
- stateMu.Lock()
- res.Skipped++
- stateMu.Unlock()
- s.log.Warn("skip inaccessible cloud directory",
- zap.String("library_id", lib.ID),
- zap.String("provider", typ),
- zap.String("dir", dirID),
- zap.Error(err))
- return nil
- }
- setWalkErr(err)
- return err
- }
- stateMu.Lock()
- dirsVisited++
- dirProgress := dirsVisited == 1 || dirsVisited%20 == 0
- stateMu.Unlock()
- publishProgress("listing", dirProgress)
- sidecars := newCloudSidecarSet(typ, entries)
- dirMeta := s.cloudDirectoryMetadata(ctx, typ, displayDir, sidecars, inheritedMeta)
- s.cacheCloudMetadataArtworkNow(ctx, dirMeta)
- for _, entry := range entries {
- select {
- case <-ctx.Done():
- setWalkErr(ctx.Err())
- return ctx.Err()
- default:
- }
- if entry.IsDir {
- if strings.TrimSpace(entry.ID) != "" {
- walkWG.Add(1)
- go func(childID, childDisplay string, childMeta *LocalMetadata) {
- _ = walkCloud(childID, childDisplay, childMeta)
- }(entry.ID, joinCloudDisplayPath(displayDir, entry.Name), dirMeta)
- }
- continue
- }
- ext := strings.ToLower(filepath.Ext(entry.Name))
- if _, ok := videoExtensions[ext]; !ok {
- continue
- }
- ref := cloudEntryRef(typ, entry.ID, entry.PickCode)
- if ref == "" {
- stateMu.Lock()
- res.Skipped++
- stateMu.Unlock()
- continue
- }
- stateMu.Lock()
- if _, ok := seenRefs[ref]; ok {
- res.Skipped++
- stateMu.Unlock()
- continue
- }
- seenRefs[ref] = struct{}{}
- filesDiscovered++
- fileProgress := filesDiscovered%100 == 0
- stateMu.Unlock()
- publishProgress("listing", fileProgress)
- displayPath := joinCloudDisplayPath(displayDir, entry.Name)
- path := cloudMediaPath(typ, displayPath)
- localMeta := s.cloudFileMetadata(ctx, typ, displayPath, entry.Name, sidecars, dirMeta, librarySupportsSeasons(lib))
- // 每个文件的海报/背景图改走后台预取队列。此前这里是同步
- // CloudResolve+下载(每张最多 20s 超时),几千个文件的云盘库
- // 扫描会变成持续数小时的串行下载,把 CPU/带宽长期吃满。
- if localMeta != nil {
- s.queueCloudArtworkPrefetch(localMeta.PosterURL)
- s.queueCloudArtworkPrefetch(localMeta.BackdropURL)
- }
- candidate := cloudCandidate{
- ref: ref,
- name: entry.Name,
- size: entry.Size,
- path: path,
- localMeta: localMeta,
- }
- key := cloudMediaDedupeKey(lib, displayDir, entry.Name, entry.Size)
- stateMu.Lock()
- if key != "" {
- if prevIndex, ok := candidateByKey[key]; ok {
- res.Skipped++
- if candidate.size > candidates[prevIndex].size {
- candidates[prevIndex] = candidate
- }
- stateMu.Unlock()
- continue
- }
- candidateByKey[key] = len(candidates)
- }
- candidates = append(candidates, candidate)
- stateMu.Unlock()
- }
- return nil
- }
- walkWG.Add(1)
- go func() {
- _ = walkCloud(rootDir, rootDisplayDir, nil)
- }()
- walkWG.Wait()
- if walkErr != nil {
- return res, walkErr
- }
- if err := ctx.Err(); err != nil {
- return res, err
- }
- existingMedia, err := s.existingCloudMediaSnapshot(ctx, lib.ID)
- if err != nil {
- s.log.Warn("load existing cloud media snapshot failed", zap.String("library_id", lib.ID), zap.Error(err))
- existingMedia = nil
- }
- if existingMedia != nil {
- priority := func(candidate cloudCandidate) int {
- existing, ok := existingMedia[candidate.path]
- if !ok {
- return 2
- }
- if cloudTrackMetadataMissing(existing) || cloudMetadataNeedsRefresh(existing, candidate.localMeta) {
- return 0
- }
- return 1
- }
- sort.SliceStable(candidates, func(i, j int) bool {
- return priority(candidates[i]) < priority(candidates[j])
- })
- }
- writeBatch := newLocalMediaWriteBatch(s, ctx, res, 100)
- probeBudget := maxCloudMediaProbeQueuePerScan
- for _, candidate := range candidates {
- select {
- case <-ctx.Done():
- return res, ctx.Err()
- default:
- }
- seen[candidate.path] = struct{}{}
- s.ingestCloudFile(ctx, lib, typ, candidate.ref, candidate.path, candidate.name, candidate.size, candidate.localMeta, existingMedia, writeBatch, &probeBudget, res)
- publishProgress("importing", res.Visited == 1 || res.Visited%100 == 0)
- }
- writeBatch.Flush()
- removed, err := s.pruneMissingCloudMedia(ctx, lib.ID, seen)
- if err != nil {
- s.log.Warn("prune missing cloud media failed", zap.String("library_id", lib.ID), zap.Error(err))
- } else {
- res.Removed = removed
- }
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "finished": true,
- "visited": res.Visited,
- "added": res.Added,
- "updated": res.Updated,
- "skipped": res.Skipped,
- "removed": res.Removed,
- "error_count": res.ErrorCount,
- "errors": res.Errors,
- "discovered": filesDiscovered,
- "dirs": dirsVisited,
- "elapsed_seconds": int(time.Since(startedAt).Seconds()),
- "cloud": true,
- })
- s.invalidateMediaCache(ctx)
- s.maybeGenerateSTRMAfterScan(lib.ID)
- if autoScrape && s.scraper != nil && s.scraper.AnyEnabled() && s.autoScrapeEnabled(ctx) {
- s.startAutoScrape(ctx, lib.ID)
- }
- return res, nil
-}
-
-func (s *ScannerService) invalidateMediaCache(ctx context.Context) {
- if s != nil && s.cache != nil {
- s.cache.DeletePrefix(ctx, "media:")
- s.cache.DeletePrefix(ctx, "stats:")
- }
-}
-
-func (s *ScannerService) startAutoScrape(ctx context.Context, libraryID string) {
- scrapeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Minute)
- go func() {
- defer cancel()
- if _, err := s.scraper.EnrichLibrary(scrapeCtx, libraryID); err != nil {
- s.log.Warn("scraper enrich failed", zap.Error(err))
- }
- }()
-}
-
-func (s *ScannerService) existingCloudMediaSnapshot(ctx context.Context, libraryID string) (map[string]existingCloudMedia, error) {
- var rows []struct {
- Path string
- SizeBytes int64
- DurationSec int
- Width int
- Height int
- VideoCodec string
- AudioCodec string
- Container string
- PosterURL string
- BackdropURL string
- STRMURL string
- Year int
- TMDbID int
- BangumiID int
- DoubanID string
- TheTVDBID string
- }
- if err := s.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Select("path, size_bytes, duration_sec, width, height, video_codec, audio_codec, container, poster_url, backdrop_url, strm_url, year, tm_db_id, bangumi_id, douban_id, thetvdb_id").
- Where("library_id = ? AND path LIKE ?", libraryID, "cloud://%").
- Find(&rows).Error; err != nil {
- return nil, err
- }
- out := make(map[string]existingCloudMedia, len(rows))
- for _, row := range rows {
- if row.Path != "" {
- out[row.Path] = existingCloudMedia{
- SizeBytes: row.SizeBytes,
- DurationSec: row.DurationSec,
- Width: row.Width,
- Height: row.Height,
- VideoCodec: row.VideoCodec,
- AudioCodec: row.AudioCodec,
- Container: row.Container,
- PosterURL: row.PosterURL,
- BackdropURL: row.BackdropURL,
- STRMURL: row.STRMURL,
- Year: row.Year,
- TMDbID: row.TMDbID,
- BangumiID: row.BangumiID,
- DoubanID: row.DoubanID,
- TheTVDBID: row.TheTVDBID,
- }
- }
- }
- return out, nil
-}
-
-func (s *ScannerService) existingLocalMediaSnapshot(ctx context.Context, libraryID string) (map[string]existingLocalMedia, error) {
- var rows []struct {
- Path string
- SizeBytes int64
- DurationSec int
- Width int
- Height int
- VideoCodec string
- AudioCodec string
- Container string
- STRMURL string
- FileID string
- }
- if err := s.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Select("path, size_bytes, duration_sec, width, height, video_codec, audio_codec, container, strm_url, file_id").
- Where("library_id = ? AND path NOT LIKE ?", libraryID, "cloud://%").
- Find(&rows).Error; err != nil {
- return nil, err
- }
- out := make(map[string]existingLocalMedia, len(rows))
- for _, row := range rows {
- if row.Path != "" {
- out[filepath.Clean(row.Path)] = existingLocalMedia{
- SizeBytes: row.SizeBytes,
- DurationSec: row.DurationSec,
- Width: row.Width,
- Height: row.Height,
- VideoCodec: row.VideoCodec,
- AudioCodec: row.AudioCodec,
- Container: row.Container,
- STRMURL: row.STRMURL,
- FileID: row.FileID,
- }
- }
- }
- return out, nil
-}
-
-func (s *ScannerService) shadowedCloudLibrary(ctx context.Context, lib *model.Library) *CloudMountConflict {
- libs, err := s.repo.Library.List(ctx)
- if err != nil {
- s.log.Warn("list libraries for cloud shadow check failed", zap.String("library_id", lib.ID), zap.Error(err))
- return nil
- }
- visible := FilterScannableCloudLibraries(ctx, s.repo, libs)
- for _, kept := range visible {
- if kept.ID == lib.ID {
- return nil
- }
- }
- current, ok := ParseCloudLibraryMount(lib.Path)
- if ok {
- currentKey, _ := cloudLibraryDisplayKey(*lib)
- for _, kept := range visible {
- info, ok := ParseCloudLibraryMount(kept.Path)
- if !ok || info.Provider != current.Provider {
- continue
- }
- keptKey, _ := cloudLibraryDisplayKey(kept)
- exact := currentKey != "" && currentKey == keptKey
- return &CloudMountConflict{
- Library: kept,
- Exact: exact,
- Nested: !exact,
- ExistingIsAncestor: cloudMountAncestor(info.DisplayDir, current.DisplayDir),
- }
- }
- }
- return CloudLibraryShadowed(libs, *lib)
-}
-
-func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library, typ, ref, path, name string, size int64, localMeta *LocalMetadata, existingMedia map[string]existingCloudMedia, writeBatch *localMediaWriteBatch, probeBudget *int, res *ScanResult) {
- res.Visited++
- ext := strings.ToLower(filepath.Ext(name))
- title, year := CleanQuery(name)
- if title == "" {
- title = strings.TrimSuffix(filepath.Base(name), ext)
- }
- if title == "" {
- title = ref
- }
- parsedSeason, parsedEpisode := ParseEpisode(path)
- if librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0 {
- if seriesTitle, seriesYear := cloudSeriesTitleFromMediaPath(path); seriesTitle != "" {
- title = seriesTitle
- if seriesYear > 0 {
- year = seriesYear
- }
- }
- }
- expectedSTRMURL := BuildPublicAPIURL(ctx, s.repo, s.cfg, "/api/cloud/play/"+typ, url.Values{"ref": []string{ref}})
- isNewMedia := false
- needsTrackProbe := true
- if existingMedia != nil {
- existing, exists := existingMedia[path]
- isNewMedia = !exists
- needsTrackProbe = !exists || cloudTrackMetadataMissing(existing)
- if exists && existing.SizeBytes == size && existing.STRMURL == expectedSTRMURL && !cloudMetadataNeedsRefresh(existing, localMeta) {
- if needsTrackProbe && ext != ".strm" {
- s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
- }
- res.Skipped++
- return
- }
- } else {
- isNewMedia = !s.mediaPathExists(ctx, path)
- }
- m := &model.Media{
- LibraryID: lib.ID,
- Title: title,
- Year: year,
- Path: path,
- SizeBytes: size,
- Container: strings.TrimPrefix(ext, "."),
- STRMURL: expectedSTRMURL,
- ScrapeStatus: "pending",
- }
- if ext == ".strm" {
- if targetURL, err := s.resolveCloudSTRMTarget(ctx, typ, ref); err == nil && targetURL != "" {
- m.STRMURL = targetURL
- } else if err != nil {
- s.log.Debug("read cloud strm failed", zap.String("ref", ref), zap.Error(err))
- }
- }
- m.SeasonNum = parsedSeason
- m.EpisodeNum = parsedEpisode
- if localMeta != nil {
- applyLocalMetadata(m, localMeta)
- res.LocalMetadata++
- s.queueCloudArtworkPrefetch(localMeta.PosterURL)
- s.queueCloudArtworkPrefetch(localMeta.BackdropURL)
- }
- if _, hints := pathHintMetadata(path, librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0); hints.useful() {
- if hints.TMDbID > 0 && m.TMDbID <= 0 {
- m.TMDbID = hints.TMDbID
- }
- if hints.BangumiID > 0 && m.BangumiID <= 0 {
- m.BangumiID = hints.BangumiID
- }
- if strings.TrimSpace(hints.DoubanID) != "" && strings.TrimSpace(m.DoubanID) == "" {
- m.DoubanID = strings.TrimSpace(hints.DoubanID)
- }
- if strings.TrimSpace(hints.TheTVDBID) != "" && strings.TrimSpace(m.TheTVDBID) == "" {
- m.TheTVDBID = strings.TrimSpace(hints.TheTVDBID)
- }
- }
- if isNewMedia && writeBatch != nil {
- var after func()
- if needsTrackProbe && ext != ".strm" {
- after = func() {
- s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
- }
- }
- writeBatch.AddWithAfter(path, m, after)
- return
- }
- if err := s.repo.Media.Upsert(ctx, m); err != nil {
- addScanError(res, path, err)
- s.log.Warn("upsert cloud media failed", zap.String("path", path), zap.Error(err))
- return
- }
- if needsTrackProbe && ext != ".strm" {
- s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
- }
- if isNewMedia {
- res.Added++
- } else {
- res.Updated++
- }
- if s.hub != nil && (res.Visited == 1 || res.Visited%100 == 0) {
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "path": path,
- "visited": res.Visited,
- "added": res.Added,
- "updated": res.Updated,
- "cloud": true,
- })
- }
-}
-
-func (s *ScannerService) probeCloudMediaAsync(task cloudMediaProbeTask) {
- defer func() {
- s.cloudMediaProbeMu.Lock()
- delete(s.cloudMediaProbing, task.path)
- s.cloudMediaProbeMu.Unlock()
- }()
- ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
- defer cancel()
- probe, err := s.probeCloudFileMetadata(ctx, task.typ, task.ref)
- if err != nil {
- if s.log != nil {
- s.log.Debug("cloud media async probe failed", zap.String("provider", task.typ), zap.String("path", task.path), zap.Error(err))
- }
- s.cloudMediaProbeMu.Lock()
- if s.cloudMediaProbeBackoff == nil {
- s.cloudMediaProbeBackoff = make(map[string]time.Time)
- }
- s.cloudMediaProbeBackoff[task.path] = time.Now().Add(cloudMediaProbeFailureBackoff)
- s.cloudMediaProbeMu.Unlock()
- return
- }
- updates := probeResultUpdates(probe)
- if len(updates) == 0 {
- return
- }
- if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("path = ?", task.path).Updates(updates).Error; err != nil {
- if s.log != nil {
- s.log.Debug("update cloud media track metadata failed", zap.String("path", task.path), zap.Error(err))
- }
- return
- }
- s.cloudMediaProbeMu.Lock()
- delete(s.cloudMediaProbeBackoff, task.path)
- s.cloudMediaProbeMu.Unlock()
- if s.hub != nil {
- s.hub.Publish("scan", map[string]any{
- "path": task.path,
- "cloud": true,
- "track_probed": true,
- "duration_sec": probe.DurationSec,
- "video_codec": probe.VideoCodec,
- "audio_codec": probe.AudioCodec,
- "width": probe.Width,
- "height": probe.Height,
- "probe_message": "云盘媒体轨道元数据已后台补齐",
- })
- }
-}
-
-func (s *ScannerService) probeLocalMediaAsync(task localMediaProbeTask) {
- defer func() {
- s.localMediaProbeMu.Lock()
- delete(s.localMediaProbing, task.path)
- s.localMediaProbeMu.Unlock()
- }()
- if s == nil || s.probe == nil || strings.TrimSpace(task.path) == "" {
- return
- }
- ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
- defer cancel()
- probe, err := s.probe.Probe(ctx, task.path)
- if err != nil {
- if s.log != nil {
- s.log.Debug("local media async probe failed", zap.String("path", task.path), zap.Error(err))
- }
- return
- }
- updates := probeResultUpdates(probe)
- if len(updates) == 0 {
- return
- }
- if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("path = ?", task.path).Updates(updates).Error; err != nil {
- if s.log != nil {
- s.log.Debug("update local media track metadata failed", zap.String("path", task.path), zap.Error(err))
- }
- return
- }
- if s.hub != nil {
- s.hub.Publish("scan", map[string]any{
- "path": task.path,
- "track_probed": true,
- "duration_sec": probe.DurationSec,
- "video_codec": probe.VideoCodec,
- "audio_codec": probe.AudioCodec,
- "width": probe.Width,
- "height": probe.Height,
- "probe_message": "本地媒体轨道元数据已后台补齐",
- })
- }
-}
-
-func (s *ScannerService) ffprobeWorkerCount() int {
- if s == nil || s.cfg == nil {
- return 1
- }
- return normalizeFFprobeMaxConcurrent(s.cfg.App.FFprobeMaxConcurrent)
-}
-
-func (s *ScannerService) cloudScanWorkerCount() int {
- if s == nil || s.cfg == nil {
- return 4
- }
- return normalizeCloudScanMaxConcurrent(s.cfg.App.CloudScanMaxConcurrent)
-}
-
-func normalizeCloudScanMaxConcurrent(n int) int {
- if n <= 0 {
- return 1
- }
- if n > 16 {
- return 16
- }
- return n
-}
-
-func (s *ScannerService) probeCloudFileMetadata(ctx context.Context, typ, ref string) (*ProbeResult, error) {
- if s == nil || s.probe == nil || s.storage == nil {
- return nil, errors.New("cloud probe unavailable")
- }
- link, err := s.storage.CloudResolve(ctx, typ, ref, "")
- if err != nil {
- return nil, err
- }
- return s.probe.ProbeHTTP(ctx, link.URL, link.Headers)
-}
-
-func probeResultUpdates(probe *ProbeResult) map[string]any {
- updates := map[string]any{}
- if probe == nil {
- return updates
- }
- if probe.DurationSec > 0 {
- updates["duration_sec"] = probe.DurationSec
- }
- if probe.Width > 0 {
- updates["width"] = probe.Width
- }
- if probe.Height > 0 {
- updates["height"] = probe.Height
- }
- if strings.TrimSpace(probe.VideoCodec) != "" {
- updates["video_codec"] = probe.VideoCodec
- }
- if strings.TrimSpace(probe.AudioCodec) != "" {
- updates["audio_codec"] = probe.AudioCodec
- }
- if probe.Container != "" {
- updates["container"] = probe.Container
- }
- return updates
-}
-
-func cloudMetadataNeedsRefresh(existing existingCloudMedia, localMeta *LocalMetadata) bool {
- if localMeta == nil {
- return false
- }
- if localMeta.Year > 0 && existing.Year <= 0 {
- return true
- }
- if localMeta.TMDbID > 0 && existing.TMDbID != localMeta.TMDbID {
- return true
- }
- if localMeta.BangumiID > 0 && existing.BangumiID != localMeta.BangumiID {
- return true
- }
- if strings.TrimSpace(localMeta.DoubanID) != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(localMeta.DoubanID) {
- return true
- }
- if strings.TrimSpace(localMeta.TheTVDBID) != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(localMeta.TheTVDBID) {
- return true
- }
- if strings.TrimSpace(localMeta.PosterURL) != "" && strings.TrimSpace(existing.PosterURL) == "" {
- return true
- }
- if strings.TrimSpace(localMeta.BackdropURL) != "" && strings.TrimSpace(existing.BackdropURL) == "" {
- return true
- }
- return false
-}
-
-func cloudTrackMetadataMissing(existing existingCloudMedia) bool {
- return existing.DurationSec <= 0 ||
- existing.Width <= 0 ||
- existing.Height <= 0 ||
- strings.TrimSpace(existing.VideoCodec) == "" ||
- strings.TrimSpace(existing.AudioCodec) == ""
-}
-
-func localMetadataNeedsRefresh(local *LocalMetadata) bool {
- return local != nil && (local.HasNFO || local.HasArtwork || localHasDescriptiveMetadata(local))
-}
-
-func cloudSeriesTitleFromMediaPath(mediaPath string) (string, int) {
- displayPath := strings.TrimSpace(mediaPath)
- if strings.HasPrefix(strings.ToLower(displayPath), "cloud://") {
- rest := strings.TrimPrefix(displayPath, "cloud://")
- if idx := strings.Index(rest, "/"); idx >= 0 {
- displayPath = rest[idx+1:]
- } else {
- return "", 0
- }
- }
- displayPath = strings.Trim(strings.ReplaceAll(displayPath, "\\", "/"), "/")
- if displayPath == "" {
- return "", 0
- }
- parts := strings.Split(displayPath, "/")
- if len(parts) < 2 {
- return "", 0
- }
- dirs := parts[:len(parts)-1]
- if len(dirs) == 0 {
- return "", 0
- }
- base := strings.TrimSpace(dirs[len(dirs)-1])
- usedSeasonFolder := false
- if _, ok := seasonFromDir(base); ok {
- usedSeasonFolder = true
- dirs = dirs[:len(dirs)-1]
- if len(dirs) == 0 {
- return "", 0
- }
- base = strings.TrimSpace(dirs[len(dirs)-1])
- }
- if base == "" || (!usedSeasonFolder && len(dirs) < 2) {
- return "", 0
- }
- title, year := CleanQuery(base)
- if title == "" {
- title = base
- }
- return strings.TrimSpace(title), year
-}
-
-// RemovePath deletes the media row for a path that has disappeared from disk
-// (incremental delete used by the watcher on Remove/Rename events).
-func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, error) {
- if _, err := os.Stat(path); err == nil {
- return 0, nil // still exists; nothing to remove
- }
- res := s.repo.DB.WithContext(ctx).
- Where("path = ?", path).
- Delete(&model.Media{})
- if res.Error == nil && res.RowsAffected > 0 {
- s.invalidateMediaCache(ctx)
- }
- return res.RowsAffected, res.Error
-}
-
-// ingestFile upserts a single media file. seenInodes dedups hardlinks within a
-// single scan; pass a fresh map for one-off ingests. It mutates res counters.
-func (s *ScannerService) ingestFile(ctx context.Context, lib *model.Library, path string, size int64, seenInodes map[string]string, existingMedia map[string]existingLocalMedia, writeBatch *localMediaWriteBatch, res *ScanResult) {
- res.Visited++
- ext := strings.ToLower(filepath.Ext(path))
- cleanPath := filepath.Clean(path)
-
- // Hardlink dedup: a seeding source kept by keep_seeding shares its inode
- // with the organized hardlink. Importing both would create duplicate rows
- // and double-count storage, so skip any file whose identity we've already
- // taken (within this scan or via an existing DB row pointing elsewhere).
- fileID, hasID := fileIdentity(path)
- if hasID {
- if first, ok := seenInodes[fileID]; ok && first != path {
- res.Skipped++
- s.log.Debug("scan skip hardlink duplicate",
- zap.String("path", path), zap.String("primary", first))
- return
- }
- if existingMedia == nil {
- if other, ok := s.duplicateByFileID(ctx, fileID, path); ok {
- res.Skipped++
- s.log.Debug("scan skip hardlink duplicate (existing)",
- zap.String("path", path), zap.String("primary", other))
- return
- }
- }
- seenInodes[fileID] = path
- }
-
- parsedSeason, parsedEpisode := ParseEpisode(path)
- localMeta, localMetaErr := ReadLocalMetadata(path, lib.Path, librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0)
- if localMetaErr != nil {
- s.log.Warn("read local metadata failed", zap.String("path", path), zap.Error(localMetaErr))
- }
- isNewMedia := false
- if existingMedia != nil {
- existing, exists := existingMedia[cleanPath]
- isNewMedia = !exists
- if exists &&
- ext != ".strm" &&
- existing.SizeBytes == size &&
- !localMetadataNeedsRefresh(localMeta) {
- res.Skipped++
- return
- }
- } else {
- isNewMedia = !s.mediaPathExists(ctx, path)
- }
-
- title, year := CleanQuery(path)
- if title == "" {
- title = strings.TrimSuffix(filepath.Base(path), ext)
- }
-
- m := &model.Media{
- LibraryID: lib.ID,
- Title: title,
- Year: year,
- Path: path,
- SizeBytes: size,
- Container: strings.TrimPrefix(ext, "."),
- FileID: fileID,
- }
- if ext == ".strm" {
- m.Container = "strm"
- if targetURL, err := readLocalSTRMTarget(path); err == nil && targetURL != "" {
- m.STRMURL = targetURL
- } else if err != nil {
- s.log.Debug("read local strm failed", zap.String("path", path), zap.Error(err))
- }
- }
-
- m.SeasonNum = parsedSeason
- m.EpisodeNum = parsedEpisode
-
- if localMeta != nil {
- applyLocalMetadata(m, localMeta)
- res.LocalMetadata++
- }
-
- var after func()
- if ext != ".strm" && s.probe != nil {
- after = func() {
- s.queueLocalMediaProbe(path)
- }
- }
-
- if isNewMedia && writeBatch != nil {
- writeBatch.AddWithAfter(path, m, after)
- return
- }
- if err := s.repo.Media.Upsert(ctx, m); err != nil {
- addScanError(res, path, err)
- s.log.Warn("upsert media failed", zap.String("path", path), zap.Error(err))
- return
- }
- if after != nil {
- after()
- }
- if isNewMedia {
- res.Added++
- } else {
- res.Updated++
- }
- s.hub.Publish("scan", map[string]any{
- "library_id": lib.ID,
- "path": path,
- "visited": res.Visited,
- "added": res.Added,
- "updated": res.Updated,
- "probed": res.Probed,
- "local_meta": res.LocalMetadata,
- })
-}
-
-type localMediaWriteBatch struct {
- scanner *ScannerService
- ctx context.Context
- res *ScanResult
- limit int
- items []localMediaWriteItem
-}
-
-type localMediaWriteItem struct {
- path string
- media *model.Media
- after func()
-}
-
-func newLocalMediaWriteBatch(scanner *ScannerService, ctx context.Context, res *ScanResult, limit int) *localMediaWriteBatch {
- if limit <= 0 {
- limit = 100
- }
- return &localMediaWriteBatch{scanner: scanner, ctx: ctx, res: res, limit: limit}
-}
-
-func (b *localMediaWriteBatch) Add(path string, media *model.Media) {
- b.AddWithAfter(path, media, nil)
-}
-
-func (b *localMediaWriteBatch) AddWithAfter(path string, media *model.Media, after func()) {
- if b == nil || b.scanner == nil || media == nil {
- return
- }
- if media.ScrapeStatus == "" {
- media.ScrapeStatus = "pending"
- }
- b.items = append(b.items, localMediaWriteItem{path: path, media: media, after: after})
- if len(b.items) >= b.limit {
- b.Flush()
- }
-}
-
-func (b *localMediaWriteBatch) Flush() {
- if b == nil || len(b.items) == 0 || b.scanner == nil || b.scanner.repo == nil || b.scanner.repo.DB == nil {
- return
- }
- items := b.items
- b.items = nil
- media := make([]model.Media, 0, len(items))
- for _, item := range items {
- if item.media != nil {
- media = append(media, *item.media)
- }
- }
- if len(media) == 0 {
- return
- }
- if err := b.scanner.repo.DB.WithContext(b.ctx).CreateInBatches(&media, b.limit).Error; err == nil {
- b.res.Added += len(media)
- for _, item := range items {
- if item.after != nil {
- item.after()
- }
- }
- b.publish()
- return
- }
- for _, item := range items {
- if item.media == nil {
- continue
- }
- if err := b.scanner.repo.Media.Upsert(b.ctx, item.media); err != nil {
- addScanError(b.res, item.path, err)
- b.scanner.log.Warn("upsert media failed", zap.String("path", item.path), zap.Error(err))
- continue
- }
- b.res.Added++
- if item.after != nil {
- item.after()
- }
- }
- b.publish()
-}
-
-func (b *localMediaWriteBatch) publish() {
- if b == nil || b.scanner == nil || b.scanner.hub == nil || b.res == nil {
- return
- }
- b.scanner.hub.Publish("scan", map[string]any{
- "library_id": b.res.LibraryID,
- "visited": b.res.Visited,
- "added": b.res.Added,
- "updated": b.res.Updated,
- "probed": b.res.Probed,
- "local_meta": b.res.LocalMetadata,
- "batched": true,
- })
-}
-
-// duplicateByFileID reports an existing media path that shares the given inode
-// identity but lives at a different path and still exists on disk.
-func (s *ScannerService) duplicateByFileID(ctx context.Context, fileID, path string) (string, bool) {
- if fileID == "" {
- return "", false
- }
- var rows []model.Media
- if err := s.repo.DB.WithContext(ctx).
- Where("file_id = ? AND path <> ?", fileID, path).
- Limit(8).Find(&rows).Error; err != nil {
- return "", false
- }
- for _, r := range rows {
- if r.Path == "" {
- continue
- }
- if _, err := os.Stat(r.Path); err == nil {
- return r.Path, true
- }
- }
- return "", false
-}
-
-func (s *ScannerService) mediaPathExists(ctx context.Context, path string) bool {
- var count int64
- err := s.repo.DB.WithContext(ctx).Unscoped().Model(&model.Media{}).
- Where("path = ?", path).Count(&count).Error
- return err == nil && count > 0
-}
-
-func (s *ScannerService) pruneMissingMedia(ctx context.Context, libraryID string, seen map[string]struct{}) (int64, error) {
- // 只取 id/path,并把删除按批提交:此前整表载入完整 Media 结构体、
- // 每行一条 DELETE,大库 prune 既费内存又长期占用写锁。
- var rows []struct {
- ID string
- Path string
- }
- if err := s.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Select("id, path").
- Where("library_id = ?", libraryID).
- Find(&rows).Error; err != nil {
- return 0, err
- }
- stale := make([]string, 0)
- for _, row := range rows {
- if row.Path == "" {
- continue
- }
- if _, ok := seen[filepath.Clean(row.Path)]; ok {
- continue
- }
- if _, err := os.Stat(row.Path); err == nil {
- continue
- } else if !os.IsNotExist(err) {
- continue
- }
- stale = append(stale, row.ID)
- }
- return s.deleteMediaByIDs(ctx, stale, false)
-}
-
-// deleteMediaByIDs removes media rows in fixed-size batches so each write
-// transaction stays short and the global write gate is released frequently.
-func (s *ScannerService) deleteMediaByIDs(ctx context.Context, ids []string, hard bool) (int64, error) {
- const batch = 500
- var removed int64
- for i := 0; i < len(ids); i += batch {
- end := i + batch
- if end > len(ids) {
- end = len(ids)
- }
- q := s.repo.DB.WithContext(ctx)
- if hard {
- q = q.Unscoped()
- }
- res := q.Where("id IN ?", ids[i:end]).Delete(&model.Media{})
- if res.Error != nil {
- return removed, res.Error
- }
- removed += res.RowsAffected
- }
- return removed, nil
-}
-
-func (s *ScannerService) pruneMissingCloudMedia(ctx context.Context, libraryID string, seen map[string]struct{}) (int64, error) {
- var rows []struct {
- ID string
- Path string
- }
- if err := s.repo.DB.WithContext(ctx).
- Model(&model.Media{}).
- Select("id, path").
- Where("library_id = ? AND path LIKE ?", libraryID, "cloud://%").
- Find(&rows).Error; err != nil {
- return 0, err
- }
- stale := make([]string, 0)
- for _, row := range rows {
- if _, ok := seen[row.Path]; ok {
- continue
- }
- stale = append(stale, row.ID)
- }
- return s.deleteMediaByIDs(ctx, stale, true)
-}
-
-func parseCloudLibraryPath(raw string) (typ, dirID string, ok bool) {
- info, ok := ParseCloudLibraryMount(raw)
- if !ok {
- return "", "", false
- }
- return info.Provider, info.ScanDir, true
-}
-
-func cloudEntryRef(typ, id, pickCode string) string {
- if typ == "cloud115" && strings.TrimSpace(pickCode) != "" {
- return strings.TrimSpace(pickCode)
- }
- return strings.TrimSpace(id)
-}
-
-func cloudMediaPath(typ, ref string) string {
- return "cloud://" + strings.TrimSpace(typ) + "/" + strings.TrimLeft(strings.TrimSpace(ref), "/")
-}
-
-func cloudMediaDedupeKey(lib *model.Library, dirID, name string, size int64) string {
- base := strings.TrimSpace(strings.TrimSuffix(filepath.Base(name), filepath.Ext(name)))
- if base == "" {
- return ""
- }
- season, episode := ParseEpisode(name)
- title, year := CleanQuery(name)
- title = normalizeCloudDedupeText(title)
- if (season > 0 || episode > 0) && title != "" {
- return fmt.Sprintf("episode:%s:%s:%d:%d:%d", strings.ToLower(strings.TrimSpace(lib.Type)), title, year, season, episode)
- }
- if (season > 0 || episode > 0) && title == "" {
- return fmt.Sprintf("episode-dir:%s:%s:%d:%d:%d", strings.ToLower(strings.TrimSpace(lib.Type)), normalizeCloudDedupeText(dirID), season, episode, size)
- }
- return fmt.Sprintf("file:%s:%d", normalizeCloudDedupeText(base), size)
-}
-
-func normalizeCloudDedupeText(value string) string {
- value = strings.ToLower(strings.TrimSpace(value))
- if value == "" {
- return ""
- }
- fields := strings.FieldsFunc(value, func(r rune) bool {
- switch r {
- case '.', '_', '-', ' ', '\t', '/', '\\', '[', ']', '(', ')':
- return true
- default:
- return false
- }
- })
- return strings.Join(fields, " ")
-}
-
-func (s *ScannerService) resolveCloudSTRMTarget(ctx context.Context, typ, ref string) (string, error) {
- if s.storage == nil {
- return "", nil
- }
- content, err := s.storage.CloudReadText(ctx, typ, ref, 64<<10)
- if err != nil {
- return "", err
- }
- for _, line := range strings.Split(content, "\n") {
- candidate := strings.TrimSpace(strings.TrimPrefix(line, "\ufeff"))
- if candidate == "" || strings.HasPrefix(candidate, "#") {
- continue
- }
- u, err := url.Parse(candidate)
- if err != nil {
- continue
- }
- switch strings.ToLower(u.Scheme) {
- case "http", "https", "webdav", "davs", "alist", "alists", "openlist", "openlists":
- return candidate, nil
- }
- }
- return "", nil
-}
-
-func readLocalSTRMTarget(path string) (string, error) {
- data, err := os.ReadFile(path) // #nosec G304 -- path is a discovered .strm file under the configured library root.
- if err != nil {
- return "", err
- }
- for _, line := range strings.Split(string(data), "\n") {
- candidate := strings.TrimSpace(strings.TrimPrefix(line, "\ufeff"))
- if candidate == "" || strings.HasPrefix(candidate, "#") {
- continue
- }
- if strings.HasPrefix(candidate, "/api/") || strings.HasPrefix(candidate, "/Videos/") || strings.HasPrefix(candidate, "/videos/") {
- return candidate, nil
- }
- u, err := url.Parse(candidate)
- if err != nil {
- continue
- }
- switch strings.ToLower(u.Scheme) {
- case "http", "https", "webdav", "davs", "alist", "alists", "openlist", "openlists":
- return candidate, nil
- }
- }
- return "", nil
-}
-
-func applyLocalMetadata(m *model.Media, local *LocalMetadata) {
- if local.Title != "" {
- m.Title = local.Title
- }
- if local.OriginalName != "" {
- m.OriginalName = local.OriginalName
- }
- if local.AdultCode != "" {
- m.OriginalName = local.AdultCode
- }
- if local.Year > 0 {
- m.Year = local.Year
- }
- if local.Overview != "" {
- m.Overview = local.Overview
- }
- if local.Rating > 0 {
- m.Rating = local.Rating
- }
- if local.PosterURL != "" {
- m.PosterURL = local.PosterURL
- }
- if local.BackdropURL != "" {
- m.BackdropURL = local.BackdropURL
- }
- if local.TMDbID > 0 {
- m.TMDbID = local.TMDbID
- }
- if local.BangumiID > 0 {
- m.BangumiID = local.BangumiID
- }
- if local.DoubanID != "" {
- m.DoubanID = local.DoubanID
- }
- if local.TheTVDBID != "" {
- m.TheTVDBID = local.TheTVDBID
- }
- if local.SeasonNum > 0 || local.EpisodeNum > 0 {
- m.SeasonNum = local.SeasonNum
- }
- if local.EpisodeNum > 0 {
- m.EpisodeNum = local.EpisodeNum
- }
- if local.Genres != "" {
- m.Genres = local.Genres
- }
- if local.Countries != "" {
- m.Countries = local.Countries
- }
- if local.Languages != "" {
- m.Languages = local.Languages
- }
- if local.NSFW {
- m.NSFW = true
- }
- if localMetadataMarksMatched(local) {
- m.ScrapeStatus = "matched"
- }
-}
-
-func localMetadataMarksMatched(local *LocalMetadata) bool {
- return local != nil && (local.HasNFO || (!local.PathHint && localHasDescriptiveMetadata(local)))
-}
-
-func localHasDescriptiveMetadata(local *LocalMetadata) bool {
- if local == nil {
- return false
- }
- return local.Title != "" ||
- local.OriginalName != "" ||
- local.AdultCode != "" ||
- local.Year > 0 ||
- local.Overview != "" ||
- local.Rating > 0 ||
- local.TMDbID > 0 ||
- local.BangumiID > 0 ||
- local.DoubanID != "" ||
- local.TheTVDBID != "" ||
- local.Genres != "" ||
- local.Countries != "" ||
- local.Languages != ""
-}
-
-func (s *ScannerService) autoScrapeEnabled(ctx context.Context) bool {
- if s.repo == nil || s.repo.Setting == nil {
- return false
- }
- value, err := s.repo.Setting.Get(ctx, "scrape.auto_on_scan")
- if err != nil {
- s.log.Warn("read scrape.auto_on_scan failed", zap.Error(err))
- return false
- }
- switch strings.ToLower(strings.TrimSpace(value)) {
- case "1", "true", "yes", "on", "enabled":
- return true
- default:
- return false
- }
-}
-
-func (s *ScannerService) maybeGenerateSTRMAfterScan(libraryID string) {
- if s == nil || s.repo == nil || s.repo.Setting == nil {
- return
- }
- value, err := s.repo.Setting.Get(context.Background(), "strm.auto_generate_enabled")
- if err != nil || !parseBoolSetting(value, false) {
- return
- }
- go func() {
- strmSvc := NewSTRMService(s.log, s.repo, s.cfg)
- if _, err := strmSvc.GenerateForLibrary(context.Background(), GenerateSTRMOptions{
- LibraryID: libraryID,
- Enabled: true,
- IncludeLocal: true,
- }); err != nil && s.log != nil {
- s.log.Warn("auto generate strm failed", zap.String("library_id", libraryID), zap.Error(err))
- }
- }()
+ Title string
+ OriginalName string
+ EpisodeTitle string
+ SizeBytes int64
+ DurationSec int
+ Width int
+ Height int
+ VideoCodec string
+ AudioCodec string
+ Container string
+ STRMURL string
+ FileID string
+ PosterURL string
+ BackdropURL string
+ Overview string
+ Year int
+ Rating float32
+ TMDbID int
+ BangumiID int
+ DoubanID string
+ TheTVDBID string
+ SeasonNum int
+ EpisodeNum int
+ Genres string
+ Countries string
+ Languages string
+ NSFW bool
+ ScrapeStatus string
}
diff --git a/internal/service/scanner_cloud_artwork.go b/internal/service/scanner_cloud_artwork.go
new file mode 100644
index 0000000..92360f7
--- /dev/null
+++ b/internal/service/scanner_cloud_artwork.go
@@ -0,0 +1,158 @@
+package service
+
+import (
+ "context"
+ "net/url"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+type cloudImagePrefetchTask struct {
+ typ string
+ ref string
+ stableKey string
+}
+
+func (s *ScannerService) cloudImagePrefetchWorker() {
+ for task := range s.cloudImagePrefetchQueue {
+ s.prefetchCloudImage(task)
+ }
+}
+
+func (s *ScannerService) queueCloudArtworkPrefetch(raw string) {
+ if s == nil || s.storage == nil || s.imageProxy == nil {
+ return
+ }
+ typ, ref, ok := ParseCloudArtworkURL(raw)
+ if !ok {
+ return
+ }
+ stableKey := typ + ":" + ref
+ if s.imageProxy.CloudImageCached(stableKey) {
+ return
+ }
+ s.cloudImagePrefetchMu.Lock()
+ if _, ok := s.cloudImagePrefetching[stableKey]; ok {
+ s.cloudImagePrefetchMu.Unlock()
+ return
+ }
+ s.cloudImagePrefetching[stableKey] = struct{}{}
+ s.cloudImagePrefetchMu.Unlock()
+
+ task := cloudImagePrefetchTask{typ: typ, ref: ref, stableKey: stableKey}
+ select {
+ case s.cloudImagePrefetchQueue <- task:
+ default:
+ s.cloudImagePrefetchMu.Lock()
+ delete(s.cloudImagePrefetching, stableKey)
+ s.cloudImagePrefetchMu.Unlock()
+ if s.log != nil {
+ s.log.Debug("cloud artwork prefetch queue full", zap.String("provider", typ), zap.String("ref", ref))
+ }
+ }
+}
+
+func (s *ScannerService) prefetchCloudImage(task cloudImagePrefetchTask) {
+ defer func() {
+ s.cloudImagePrefetchMu.Lock()
+ delete(s.cloudImagePrefetching, task.stableKey)
+ s.cloudImagePrefetchMu.Unlock()
+ }()
+ if s == nil || s.storage == nil || s.imageProxy == nil || s.imageProxy.CloudImageCached(task.stableKey) {
+ return
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 45*time.Second)
+ defer cancel()
+ link, err := s.storage.CloudResolve(ctx, task.typ, task.ref, "")
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("resolve cloud artwork for prefetch failed", zap.String("provider", task.typ), zap.String("ref", task.ref), zap.Error(err))
+ }
+ return
+ }
+ if err := s.imageProxy.PrefetchCloudResolved(ctx, task.stableKey, link); err != nil && s.log != nil {
+ s.log.Debug("prefetch cloud artwork failed", zap.String("provider", task.typ), zap.String("ref", task.ref), zap.Error(err))
+ }
+}
+
+func (s *ScannerService) cacheCloudArtworkNow(ctx context.Context, raw string) {
+ if s == nil || s.storage == nil || s.imageProxy == nil {
+ return
+ }
+ typ, ref, ok := ParseCloudArtworkURL(raw)
+ if !ok {
+ return
+ }
+ stableKey := typ + ":" + ref
+ if s.imageProxy.CloudImageCached(stableKey) {
+ return
+ }
+ cacheCtx, cancel := context.WithTimeout(ctx, 20*time.Second)
+ defer cancel()
+ link, err := s.storage.CloudResolve(cacheCtx, typ, ref, "")
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("resolve cloud artwork for priority cache failed", zap.String("provider", typ), zap.String("ref", ref), zap.Error(err))
+ }
+ s.queueCloudArtworkPrefetch(raw)
+ return
+ }
+ if err := s.imageProxy.PrefetchCloudResolved(cacheCtx, stableKey, link); err != nil {
+ if s.log != nil {
+ s.log.Debug("priority cache cloud artwork failed", zap.String("provider", typ), zap.String("ref", ref), zap.Error(err))
+ }
+ s.queueCloudArtworkPrefetch(raw)
+ }
+}
+
+func (s *ScannerService) cacheCloudMetadataArtworkNow(ctx context.Context, meta *LocalMetadata) {
+ if meta == nil {
+ return
+ }
+ s.cacheCloudArtworkNow(ctx, meta.PosterURL)
+ s.cacheCloudArtworkNow(ctx, meta.BackdropURL)
+}
+
+func ParseCloudArtworkURL(raw string) (string, string, bool) {
+ u, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil {
+ return "", "", false
+ }
+ path := strings.Trim(u.Path, "/")
+ typ := ""
+ for _, prefix := range []string{"api/img/cloud/", "api/cloud/play/"} {
+ if strings.HasPrefix(strings.ToLower(path), prefix) {
+ typ = strings.TrimSpace(path[len(prefix):])
+ break
+ }
+ }
+ if typ == "" {
+ return "", "", false
+ }
+ ref := strings.TrimSpace(u.Query().Get("ref"))
+ if typ == "" || ref == "" || !isCloudArtworkRef(ref) {
+ return "", "", false
+ }
+ return typ, ref, true
+}
+
+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)
+}
+
+func isCloudArtworkRef(ref string) bool {
+ ref = strings.ToLower(strings.TrimSpace(ref))
+ for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
+ if strings.HasSuffix(ref, suffix) {
+ return true
+ }
+ }
+ return false
+}
diff --git a/internal/service/scanner_cloud_conflict.go b/internal/service/scanner_cloud_conflict.go
new file mode 100644
index 0000000..f712b5d
--- /dev/null
+++ b/internal/service/scanner_cloud_conflict.go
@@ -0,0 +1,42 @@
+package service
+
+import (
+ "context"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) shadowedCloudLibrary(ctx context.Context, lib *model.Library) *CloudMountConflict {
+ libs, err := s.repo.Library.List(ctx)
+ if err != nil {
+ s.log.Warn("list libraries for cloud shadow check failed", zap.String("library_id", lib.ID), zap.Error(err))
+ return nil
+ }
+ visible := FilterScannableCloudLibraries(ctx, s.repo, libs)
+ for _, kept := range visible {
+ if kept.ID == lib.ID {
+ return nil
+ }
+ }
+ current, ok := ParseCloudLibraryMount(lib.Path)
+ if ok {
+ currentKey, _ := cloudLibraryDisplayKey(*lib)
+ for _, kept := range visible {
+ info, ok := ParseCloudLibraryMount(kept.Path)
+ if !ok || info.Provider != current.Provider {
+ continue
+ }
+ keptKey, _ := cloudLibraryDisplayKey(kept)
+ exact := currentKey != "" && currentKey == keptKey
+ return &CloudMountConflict{
+ Library: kept,
+ Exact: exact,
+ Nested: !exact,
+ ExistingIsAncestor: cloudMountAncestor(info.DisplayDir, current.DisplayDir),
+ }
+ }
+ }
+ return CloudLibraryShadowed(libs, *lib)
+}
diff --git a/internal/service/scanner_cloud_enrich.go b/internal/service/scanner_cloud_enrich.go
new file mode 100644
index 0000000..f1abc0f
--- /dev/null
+++ b/internal/service/scanner_cloud_enrich.go
@@ -0,0 +1,152 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) enrichCloudMetadataFromExternalIDs(ctx context.Context, lib *model.Library, path string, meta *LocalMetadata) *LocalMetadata {
+ if s == nil || s.scraper == nil || meta == nil || !cloudMetadataNeedsExternalEnrich(meta) {
+ return meta
+ }
+ localPoster, localBackdrop := cloudLocalArtworkURLs(meta)
+ media := &model.Media{
+ LibraryID: "",
+ Title: firstNonEmpty(meta.Title, pathBaseSlash(path)),
+ Path: path,
+ Year: meta.Year,
+ TMDbID: meta.TMDbID,
+ BangumiID: meta.BangumiID,
+ DoubanID: meta.DoubanID,
+ TheTVDBID: meta.TheTVDBID,
+ SeasonNum: meta.SeasonNum,
+ EpisodeNum: meta.EpisodeNum,
+ PosterURL: meta.PosterURL,
+ BackdropURL: meta.BackdropURL,
+ }
+ if lib != nil {
+ media.LibraryID = lib.ID
+ }
+ enrichCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
+ defer cancel()
+ match := s.scraper.matchFromMediaExternalIDs(enrichCtx, media, lib)
+ if match == nil {
+ return meta
+ }
+ s.scraper.applyFanartArtwork(enrichCtx, match)
+ mergeLocalMetadataIntoMatch(match, meta)
+
+ enriched := cloneLocalMetadata(meta)
+ if enriched == nil {
+ enriched = &LocalMetadata{}
+ }
+ mergeMatchIntoLocalMetadata(enriched, match)
+ if localPoster != "" {
+ enriched.PosterURL = localPoster
+ enriched.HasArtwork = true
+ }
+ if localBackdrop != "" {
+ enriched.BackdropURL = localBackdrop
+ enriched.HasArtwork = true
+ }
+ enriched.PathHint = false
+ enriched.HasNFO = true
+ if enriched.PosterURL != "" || enriched.BackdropURL != "" {
+ enriched.HasArtwork = true
+ }
+ s.prefetchRemoteArtworkFromScan(ctx, enriched.PosterURL)
+ s.prefetchRemoteArtworkFromScan(ctx, enriched.BackdropURL)
+ return enriched
+}
+
+func cloudMetadataNeedsExternalEnrich(meta *LocalMetadata) bool {
+ if meta == nil {
+ return false
+ }
+ hasExternalID := meta.TMDbID > 0 || meta.BangumiID > 0 || strings.TrimSpace(meta.DoubanID) != "" || strings.TrimSpace(meta.TheTVDBID) != ""
+ if !hasExternalID {
+ return false
+ }
+ return meta.PosterURL == "" || meta.BackdropURL == "" || meta.Overview == "" || meta.Title == ""
+}
+
+func cloudLocalArtworkURLs(meta *LocalMetadata) (poster, backdrop string) {
+ if meta == nil || !meta.HasArtwork {
+ return "", ""
+ }
+ if _, _, ok := ParseCloudArtworkURL(meta.PosterURL); ok {
+ poster = meta.PosterURL
+ }
+ if _, _, ok := ParseCloudArtworkURL(meta.BackdropURL); ok {
+ backdrop = meta.BackdropURL
+ }
+ return poster, backdrop
+}
+
+func mergeMatchIntoLocalMetadata(meta *LocalMetadata, match *Match) {
+ if meta == nil || match == nil {
+ return
+ }
+ if match.Title != "" {
+ meta.Title = match.Title
+ }
+ if match.OriginalName != "" {
+ meta.OriginalName = match.OriginalName
+ }
+ if match.Year > 0 {
+ meta.Year = match.Year
+ }
+ if match.Overview != "" {
+ meta.Overview = match.Overview
+ }
+ if match.Rating > 0 {
+ meta.Rating = match.Rating
+ }
+ if match.PosterURL != "" {
+ meta.PosterURL = match.PosterURL
+ }
+ if match.BackdropURL != "" {
+ meta.BackdropURL = match.BackdropURL
+ }
+ if match.TMDbID > 0 {
+ meta.TMDbID = match.TMDbID
+ }
+ if match.BangumiID > 0 {
+ meta.BangumiID = match.BangumiID
+ }
+ if match.DoubanID != "" {
+ meta.DoubanID = match.DoubanID
+ }
+ if match.TheTVDBID != "" {
+ meta.TheTVDBID = match.TheTVDBID
+ }
+ if len(match.Genres) > 0 {
+ meta.Genres = strings.Join(match.Genres, ",")
+ }
+ if len(match.Countries) > 0 {
+ meta.Countries = strings.Join(match.Countries, ",")
+ }
+ if len(match.Languages) > 0 {
+ meta.Languages = strings.Join(match.Languages, ",")
+ }
+ if match.NSFW {
+ meta.NSFW = true
+ }
+}
+
+func (s *ScannerService) prefetchRemoteArtworkFromScan(ctx context.Context, raw string) {
+ if s == nil || s.imageProxy == nil || !isHTTPish(raw) {
+ return
+ }
+ fetchCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second)
+ err := s.imageProxy.PrefetchRemote(fetchCtx, raw)
+ cancel()
+ if err != nil && s.log != nil {
+ s.log.Debug("scan remote artwork prefetch failed", zap.String("url", raw), zap.Error(err))
+ }
+}
diff --git a/internal/service/scanner_cloud_ingest.go b/internal/service/scanner_cloud_ingest.go
new file mode 100644
index 0000000..eb7ac22
--- /dev/null
+++ b/internal/service/scanner_cloud_ingest.go
@@ -0,0 +1,121 @@
+package service
+
+import (
+ "context"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) ingestCloudFile(ctx context.Context, lib *model.Library, typ, ref, path, name string, size int64, localMeta *LocalMetadata, existingMedia map[string]existingCloudMedia, writeBatch *localMediaWriteBatch, probeBudget *int, res *ScanResult) {
+ res.Visited++
+ ext := strings.ToLower(filepath.Ext(name))
+ title, year := CleanQuery(name)
+ if title == "" {
+ title = strings.TrimSuffix(filepath.Base(name), ext)
+ }
+ if title == "" {
+ title = ref
+ }
+ parsedSeason, parsedEpisode := ParseEpisode(path)
+ if librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0 {
+ if seriesTitle, seriesYear := cloudSeriesTitleFromMediaPath(path); seriesTitle != "" {
+ title = seriesTitle
+ if seriesYear > 0 {
+ year = seriesYear
+ }
+ }
+ }
+ expectedSTRMURL := BuildRelativeCloudPlayURL(typ, ref)
+ isNewMedia := false
+ needsTrackProbe := true
+ if existingMedia != nil {
+ existing, exists := existingMedia[path]
+ isNewMedia = !exists
+ needsTrackProbe = !exists || cloudTrackMetadataMissing(existing)
+ if exists && existing.LibraryID == lib.ID && existing.SizeBytes == size && existing.STRMURL == expectedSTRMURL && !cloudMetadataNeedsRefresh(existing, localMeta) {
+ if needsTrackProbe && ext != ".strm" {
+ s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
+ }
+ res.Skipped++
+ return
+ }
+ } else {
+ isNewMedia = !s.mediaPathExists(ctx, path)
+ }
+ m := &model.Media{
+ LibraryID: lib.ID,
+ Title: title,
+ Year: year,
+ Path: path,
+ SizeBytes: size,
+ Container: strings.TrimPrefix(ext, "."),
+ STRMURL: expectedSTRMURL,
+ ScrapeStatus: "pending",
+ }
+ if ext == ".strm" {
+ if targetURL, err := s.resolveCloudSTRMTarget(ctx, typ, ref); err == nil && targetURL != "" {
+ m.STRMURL = targetURL
+ } else if err != nil {
+ s.log.Debug("read cloud strm failed", zap.String("ref", ref), zap.Error(err))
+ }
+ }
+ m.SeasonNum = parsedSeason
+ m.EpisodeNum = parsedEpisode
+ if localMeta != nil {
+ applyLocalMetadata(m, localMeta)
+ res.LocalMetadata++
+ s.queueCloudArtworkPrefetch(localMeta.PosterURL)
+ s.queueCloudArtworkPrefetch(localMeta.BackdropURL)
+ }
+ if _, hints := pathHintMetadata(path, librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0); hints.useful() {
+ if hints.TMDbID > 0 && m.TMDbID <= 0 {
+ m.TMDbID = hints.TMDbID
+ }
+ if hints.BangumiID > 0 && m.BangumiID <= 0 {
+ m.BangumiID = hints.BangumiID
+ }
+ if strings.TrimSpace(hints.DoubanID) != "" && strings.TrimSpace(m.DoubanID) == "" {
+ m.DoubanID = strings.TrimSpace(hints.DoubanID)
+ }
+ if strings.TrimSpace(hints.TheTVDBID) != "" && strings.TrimSpace(m.TheTVDBID) == "" {
+ m.TheTVDBID = strings.TrimSpace(hints.TheTVDBID)
+ }
+ }
+ if isNewMedia && writeBatch != nil {
+ var after func()
+ if needsTrackProbe && ext != ".strm" {
+ after = func() {
+ s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
+ }
+ }
+ writeBatch.AddWithAfter(path, m, after)
+ return
+ }
+ if err := s.repo.Media.Upsert(ctx, m); err != nil {
+ addScanError(res, path, err)
+ s.log.Warn("upsert cloud media failed", zap.String("path", path), zap.Error(err))
+ return
+ }
+ if needsTrackProbe && ext != ".strm" {
+ s.queueCloudMediaProbeWithBudget(typ, ref, path, probeBudget)
+ }
+ if isNewMedia {
+ res.Added++
+ } else {
+ res.Updated++
+ }
+ if s.hub != nil && (res.Visited == 1 || res.Visited%100 == 0) {
+ s.hub.Publish("scan", map[string]any{
+ "library_id": lib.ID,
+ "path": path,
+ "visited": res.Visited,
+ "added": res.Added,
+ "updated": res.Updated,
+ "cloud": true,
+ })
+ }
+}
diff --git a/internal/service/scanner_cloud_jobs.go b/internal/service/scanner_cloud_jobs.go
new file mode 100644
index 0000000..8e99ce4
--- /dev/null
+++ b/internal/service/scanner_cloud_jobs.go
@@ -0,0 +1,151 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func (s *ScannerService) StartCloudLibraryScan(libraryID string, autoScrape bool) (CloudScanStatus, bool, error) {
+ if s == nil {
+ return CloudScanStatus{}, false, errors.New("scanner unavailable")
+ }
+ lib, err := s.repo.Library.FindByID(context.Background(), libraryID)
+ if err != nil {
+ return CloudScanStatus{}, false, err
+ }
+ if lib == nil {
+ return CloudScanStatus{}, false, errors.New("library not found")
+ }
+ mount, ok := ParseCloudLibraryMount(lib.Path)
+ if !ok {
+ return CloudScanStatus{}, false, errors.New("library is not a cloud mount")
+ }
+ if IsDeprecatedNativeCloudProvider(mount.Provider) {
+ return CloudScanStatus{}, false, fmt.Errorf("cloud provider %q is deprecated; use OpenList or CloudDrive2 bridge", mount.Provider)
+ }
+ s.cloudScanMu.Lock()
+ if entry := s.cloudScans[libraryID]; cloudScanBlocksBegin(entry) {
+ status := entry.status
+ s.cloudScanMu.Unlock()
+ return status, false, nil
+ }
+ s.cloudScanMu.Unlock()
+
+ go func() {
+ ctx, cancel := cloudScanContext(context.Background(), cloudScanTimeout(context.Background(), s.repo, 24*time.Hour))
+ defer cancel()
+ if autoScrape {
+ _, err = s.ScanLibrary(ctx, libraryID)
+ } else {
+ _, err = s.ScanLibraryWithoutAutoScrape(ctx, libraryID)
+ }
+ if err != nil && !errors.Is(err, ErrCloudScanAlreadyRunning) && s.log != nil {
+ s.log.Warn("cloud library background scan failed", zap.String("library_id", libraryID), zap.Error(err))
+ }
+ }()
+ return newCloudScanEntry(libraryID, mount.Provider, nil).status.withQueuedState(), true, nil
+}
+
+func (status CloudScanStatus) withQueuedState() CloudScanStatus {
+ status.Stage = "queued"
+ status.State = "queued"
+ return status
+}
+
+func cloudScanContext(parent context.Context, timeout time.Duration) (context.Context, context.CancelFunc) {
+ if timeout <= 0 {
+ return context.WithCancel(parent)
+ }
+ return context.WithTimeout(parent, timeout)
+}
+
+func cloudScanTimeout(ctx context.Context, repo *repository.Container, fallback time.Duration) time.Duration {
+ if repo == nil || repo.Setting == nil {
+ return fallback
+ }
+ value, err := repo.Setting.Get(ctx, "cloud.scan_timeout_hours")
+ if err != nil || strings.TrimSpace(value) == "" {
+ return fallback
+ }
+ hours := parseIntSettingDefault(strings.TrimSpace(value), int(fallback/time.Hour))
+ if hours <= 0 {
+ return 0
+ }
+ return time.Duration(hours) * time.Hour
+}
+
+func (s *ScannerService) StartAllCloudLibraryScans() ([]CloudScanStatus, error) {
+ if s == nil {
+ return nil, errors.New("scanner unavailable")
+ }
+ libs, err := s.repo.Library.List(context.Background())
+ if err != nil {
+ return nil, err
+ }
+ libs = FilterScannableCloudLibraries(context.Background(), s.repo, libs)
+ statuses := make([]CloudScanStatus, 0, len(libs))
+ queue := make([]string, 0, len(libs))
+ for _, lib := range libs {
+ if !lib.Enabled {
+ continue
+ }
+ mount, ok := ParseCloudLibraryMount(lib.Path)
+ if !ok {
+ continue
+ }
+ status, queued := s.queueCloudLibraryScan(lib, mount)
+ if queued {
+ queue = append(queue, lib.ID)
+ }
+ statuses = append(statuses, status)
+ }
+ if len(queue) > 0 {
+ go s.runQueuedCloudLibraryScans(queue)
+ }
+ return statuses, nil
+}
+
+func (s *ScannerService) queueCloudLibraryScan(lib model.Library, mount CloudMountInfo) (CloudScanStatus, bool) {
+ status := newCloudScanEntry(lib.ID, mount.Provider, nil).status.withQueuedState()
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ if s.cloudScans == nil {
+ s.cloudScans = make(map[string]*cloudScanEntry)
+ }
+ if entry := s.cloudScans[lib.ID]; cloudScanActive(entry) {
+ return entry.status, false
+ }
+ s.cloudScans[lib.ID] = &cloudScanEntry{status: status}
+ return status, true
+}
+
+func (s *ScannerService) runQueuedCloudLibraryScans(libraryIDs []string) {
+ ctx, cancel := cloudScanContext(context.Background(), cloudScanTimeout(context.Background(), s.repo, 24*time.Hour))
+ defer cancel()
+ for _, libraryID := range libraryIDs {
+ if ctx.Err() != nil {
+ return
+ }
+ if s.cloudScanWasCanceled(libraryID) {
+ continue
+ }
+ if _, err := s.ScanLibrary(ctx, libraryID); err != nil && !errors.Is(err, ErrCloudScanAlreadyRunning) && !errors.Is(err, context.Canceled) && s.log != nil {
+ s.log.Warn("cloud library queued scan failed", zap.String("library_id", libraryID), zap.Error(err))
+ }
+ }
+}
+
+func (s *ScannerService) cloudScanWasCanceled(libraryID string) bool {
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ entry := s.cloudScans[libraryID]
+ return entry != nil && entry.status.State == "canceled"
+}
diff --git a/internal/service/scanner_cloud_metadata_test.go b/internal/service/scanner_cloud_metadata_test.go
new file mode 100644
index 0000000..88bf9f5
--- /dev/null
+++ b/internal/service/scanner_cloud_metadata_test.go
@@ -0,0 +1,831 @@
+package service
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestScanCloudLibraryReadsRemoteSTRMTarget(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ if r.URL.Path != "/dav/Links" {
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ _, _ = w.Write([]byte(`
+
+
+ /dav/Links/
+
+
+
+ /dav/Links/Movie.strm
+ Movie.strm32
+
+`))
+ case http.MethodGet:
+ if r.URL.Path != "/dav/Links/Movie.strm" {
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ _, _ = w.Write([]byte("https://cdn.example.com/Movie.mkv\n"))
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Links", Path: "cloud://openlist/Links", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 {
+ t.Fatalf("scan result = %#v, want added=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Path != "cloud://openlist/Links/Movie.strm" {
+ t.Fatalf("path = %q", media.Path)
+ }
+ if media.STRMURL != "https://cdn.example.com/Movie.mkv" {
+ t.Fatalf("strm target = %q", media.STRMURL)
+ }
+}
+
+func TestScanCloudLibraryCachesFileLevelRemoteArtwork(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ if r.URL.Path != "/dav/Movies" {
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/Movie.mkvMovie.mkv4096
+ /dav/Movies/Movie.nfoMovie.nfo128
+ /dav/Movies/Movie.jpgMovie.jpg1024
+`))
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Movies/Movie.nfo":
+ _, _ = w.Write([]byte(`Sidecar Movie2026`))
+ case "/dav/Movies/Movie.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("file-level-poster"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "Sidecar Movie" || media.Year != 2026 {
+ t.Fatalf("metadata not applied: %#v", media)
+ }
+ if media.PosterURL != "/api/img/cloud/openlist?ref=%2FMovies%2FMovie.jpg" {
+ t.Fatalf("poster url = %q", media.PosterURL)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Movies/Movie.jpg") {
+ t.Fatal("file-level cloud poster should be cached locally during scan before media is exposed")
+ }
+ if got := rec.Body.String(); got != "file-level-poster" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
+
+func TestScanCloudLibraryUsesArtworkReferencedByRemoteNFO(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ if r.URL.Path != "/dav/Movies" {
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/Movie.mkvMovie.mkv4096
+ /dav/Movies/Movie.nfoMovie.nfo256
+ /dav/Movies/Artwork.Custom.tbnArtwork.Custom.tbn1024
+ /dav/Movies/Scene.Still.pngScene.Still.png1024
+`))
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Movies/Movie.nfo":
+ _, _ = w.Write([]byte(`NFO Custom ArtworkArtwork.Custom.tbnScene.Still.png?version=1`))
+ case "/dav/Movies/Artwork.Custom.tbn":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("custom-poster"))
+ case "/dav/Movies/Scene.Still.png":
+ w.Header().Set("Content-Type", "image/png")
+ _, _ = w.Write([]byte("custom-backdrop"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "NFO Custom Artwork" {
+ t.Fatalf("metadata title = %q", media.Title)
+ }
+ if media.PosterURL != "/api/img/cloud/openlist?ref=%2FMovies%2FArtwork.Custom.tbn" {
+ t.Fatalf("poster url = %q", media.PosterURL)
+ }
+ if media.BackdropURL != "/api/img/cloud/openlist?ref=%2FMovies%2FScene.Still.png" {
+ t.Fatalf("backdrop url = %q", media.BackdropURL)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Movies/Artwork.Custom.tbn") {
+ t.Fatal("NFO-referenced cloud poster should be cached locally during scan")
+ }
+ if got := rec.Body.String(); got != "custom-poster" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
+
+func TestScanCloudLibraryReadsRemoteNFOAndArtwork(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Anime/JianLai":
+ _, _ = w.Write([]byte(`
+
+ /dav/Anime/JianLai/
+ /dav/Anime/JianLai/tvshow.nfotvshow.nfo64
+ /dav/Anime/JianLai/poster.jpgposter.jpg1024
+ /dav/Anime/JianLai/Season1/Season1
+`))
+ case "/dav/Anime/JianLai/Season1":
+ _, _ = w.Write([]byte(`
+
+ /dav/Anime/JianLai/Season1/
+ /dav/Anime/JianLai/Season1/JianLai.S01E01.mkvJianLai.S01E01.mkv2048
+ /dav/Anime/JianLai/Season1/JianLai.S01E01.nfoJianLai.S01E01.nfo128
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Anime/JianLai/tvshow.nfo":
+ _, _ = w.Write([]byte(`剑来2024天地有剑气`))
+ case "/dav/Anime/JianLai/Season1/JianLai.S01E01.nfo":
+ _, _ = w.Write([]byte(`剑来第一集11`))
+ case "/dav/Anime/JianLai/poster.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("cloud-poster-bytes"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · 国漫 · 剑来", Path: "cloud://openlist/Anime/JianLai", Type: "anime", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ // 单集名(episode 「第一集」)不得写入 OriginalName(整剧原名/分组键)。
+ // tvshow.nfo 未提供 originaltitle, 故 OriginalName 应为空。
+ if media.Title != "剑来" || media.OriginalName != "" || media.Year != 2024 {
+ t.Fatalf("metadata not applied: %#v", media)
+ }
+ if media.SeasonNum != 1 || media.EpisodeNum != 1 {
+ t.Fatalf("episode numbers = %d/%d", media.SeasonNum, media.EpisodeNum)
+ }
+ if media.PosterURL != "/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg" {
+ t.Fatalf("poster url = %q", media.PosterURL)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Anime/JianLai/poster.jpg") {
+ t.Fatal("cloud poster should be cached locally during scan before media is exposed")
+ }
+ if got := rec.Body.String(); got != "cloud-poster-bytes" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+ if media.ScrapeStatus != "matched" {
+ t.Fatalf("scrape status = %q", media.ScrapeStatus)
+ }
+}
+
+func TestScanCloudLibraryRefreshesExistingRemoteNFOAndArtwork(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Anime/JianLai":
+ _, _ = w.Write([]byte(`
+
+ /dav/Anime/JianLai/
+ /dav/Anime/JianLai/tvshow.nfotvshow.nfo64
+ /dav/Anime/JianLai/poster.jpgposter.jpg1024
+ /dav/Anime/JianLai/Season1/Season1
+`))
+ case "/dav/Anime/JianLai/Season1":
+ _, _ = w.Write([]byte(`
+
+ /dav/Anime/JianLai/Season1/
+ /dav/Anime/JianLai/Season1/JianLai.S01E01.mkvJianLai.S01E01.mkv2048
+ /dav/Anime/JianLai/Season1/JianLai.S01E01.nfoJianLai.S01E01.nfo128
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Anime/JianLai/tvshow.nfo":
+ _, _ = w.Write([]byte(`剑来2024天地有剑气296753`))
+ case "/dav/Anime/JianLai/Season1/JianLai.S01E01.nfo":
+ _, _ = w.Write([]byte(`剑来第一集11`))
+ case "/dav/Anime/JianLai/poster.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("cloud-poster-bytes"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · 国漫 · 剑来", Path: "cloud://openlist/Anime/JianLai", Type: "anime", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+
+ mediaPath := "cloud://openlist/Anime/JianLai/Season1/JianLai.S01E01.mkv"
+ old := model.Media{
+ LibraryID: lib.ID,
+ Title: "JianLai.S01E01",
+ Path: mediaPath,
+ SizeBytes: 2048,
+ Container: "mkv",
+ PosterURL: "https://image.tmdb.org/t/p/w500/old.jpg",
+ STRMURL: BuildRelativeCloudPlayURL("openlist", "/Anime/JianLai/Season1/JianLai.S01E01.mkv"),
+ ScrapeStatus: "no_match",
+ }
+ if err := repos.Media.Upsert(t.Context(), &old); err != nil {
+ t.Fatal(err)
+ }
+
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Updated != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want updated=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media, "path = ?", mediaPath).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "剑来" || media.Year != 2024 || media.TMDbID != 296753 {
+ t.Fatalf("metadata not refreshed: %#v", media)
+ }
+ if media.PosterURL != "/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg" {
+ t.Fatalf("poster url = %q", media.PosterURL)
+ }
+ if media.ScrapeStatus != "matched" {
+ t.Fatalf("scrape status = %q", media.ScrapeStatus)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Anime/JianLai/poster.jpg") {
+ t.Fatal("refreshed cloud poster should be cached locally during scan")
+ }
+ if got := rec.Body.String(); got != "cloud-poster-bytes" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
+
+func TestScanCloudLibraryReadsMovieDirectoryNFOAndCleanTitleArtwork(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Movies":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/Action Movie (2025) {tmdb-1197306}/Action Movie (2025) {tmdb-1197306}
+`))
+ case "/dav/Movies/Action Movie (2025) {tmdb-1197306}":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/Action%20Movie%20(2025)%20%7Btmdb-1197306%7D/
+ /dav/Movies/Action%20Movie%20(2025)%20%7Btmdb-1197306%7D/Action%20Movie%20(2025)%20-%202160p.WEB-DL.mkvAction Movie (2025) - 2160p.WEB-DL.mkv4096
+ /dav/Movies/Action%20Movie%20(2025)%20%7Btmdb-1197306%7D/movie.nfomovie.nfo128
+ /dav/Movies/Action%20Movie%20(2025)%20%7Btmdb-1197306%7D/action%20movie%20(2025)-poster.jpgaction movie (2025)-poster.jpg1024
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Movies/Action Movie (2025) {tmdb-1197306}/movie.nfo":
+ _, _ = w.Write([]byte(`Action Movie20251197306`))
+ case "/dav/Movies/Action Movie (2025) {tmdb-1197306}/action movie (2025)-poster.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("clean-title-poster"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "Action Movie" || media.Year != 2025 || media.TMDbID != 1197306 {
+ t.Fatalf("movie.nfo metadata not applied: %#v", media)
+ }
+ wantPoster := "/api/img/cloud/openlist?ref=%2FMovies%2FAction+Movie+%282025%29+%7Btmdb-1197306%7D%2Faction+movie+%282025%29-poster.jpg"
+ if media.PosterURL != wantPoster {
+ t.Fatalf("poster url = %q, want %q", media.PosterURL, wantPoster)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Movies/Action Movie (2025) {tmdb-1197306}/action movie (2025)-poster.jpg") {
+ t.Fatal("clean-title cloud poster should be cached locally during scan")
+ }
+ if got := rec.Body.String(); got != "clean-title-poster" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
+
+func TestScanCloudLibraryReadsRemoteJSONMetadataAndArtwork(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Movies":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/Sidecar Movie (2026) {tmdb-12345}
+`))
+ case "/dav/Movies/Sidecar Movie (2026) {tmdb-12345}":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/Sidecar%20Movie%20(2026).mkvSidecar Movie (2026).mkv4096
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/Sidecar%20Movie%20(2026)-mediainfo.jsonSidecar Movie (2026)-mediainfo.json256
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/poster.jpgposter.jpg1024
+ /dav/Movies/Sidecar%20Movie%20(2026)%20%7Btmdb-12345%7D/backdrop.jpgbackdrop.jpg1024
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ case http.MethodGet:
+ switch r.URL.Path {
+ case "/dav/Movies/Sidecar Movie (2026) {tmdb-12345}/Sidecar Movie (2026)-mediainfo.json":
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{"title":"JSON Sidecar Movie","year":2026,"tmdb_id":12345,"overview":"metadata from cloud json","poster":"poster.jpg","backdrop":"backdrop.jpg","genres":["Action","Drama"]}`))
+ case "/dav/Movies/Sidecar Movie (2026) {tmdb-12345}/poster.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("json-poster"))
+ case "/dav/Movies/Sidecar Movie (2026) {tmdb-12345}/backdrop.jpg":
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("json-backdrop"))
+ default:
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "JSON Sidecar Movie" || media.Year != 2026 || media.TMDbID != 12345 || media.ScrapeStatus != "matched" {
+ t.Fatalf("json metadata not applied: %#v", media)
+ }
+ wantPoster := "/api/img/cloud/openlist?ref=%2FMovies%2FSidecar+Movie+%282026%29+%7Btmdb-12345%7D%2Fposter.jpg"
+ if media.PosterURL != wantPoster {
+ t.Fatalf("poster url = %q, want %q", media.PosterURL, wantPoster)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Movies/Sidecar Movie (2026) {tmdb-12345}/poster.jpg") {
+ t.Fatal("JSON cloud poster should be cached locally during scan")
+ }
+ if got := rec.Body.String(); got != "json-poster" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
+
+func TestScanCloudLibraryEnrichesPathHintTMDbArtwork(t *testing.T) {
+ tmdb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/movie/755679" {
+ t.Fatalf("unexpected tmdb path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{
+ "id": 755679,
+ "title": "速度与激情11",
+ "original_title": "Fast X: Part 2",
+ "overview": "Exact metadata by TMDb ID",
+ "poster_path": "/poster-fast11.jpg",
+ "backdrop_path": "/backdrop-fast11.jpg",
+ "release_date": "2028-04-07",
+ "vote_average": 7.2,
+ "genres": [{"name":"Action"}],
+ "production_countries": [{"iso_3166_1":"US"}],
+ "spoken_languages": [{"iso_639_1":"en"}]
+ }`))
+ }))
+ defer tmdb.Close()
+
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Movies":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/速度与激情11 (2028) {tmdb-755679}
+`))
+ case "/dav/Movies/速度与激情11 (2028) {tmdb-755679}":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028).mkv速度与激情11 (2028).mkv4096
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}, &model.APIConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = tmdb.URL
+ cfg.Secrets.TMDbImageProxy = "https://image.tmdb.org/t/p"
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+ scanner := NewScannerService(cfg, log, repos, NewHub(log), nil, scraper)
+ scanner.SetStorageConfig(storage)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.ScrapeStatus != "matched" || media.TMDbID != 755679 || media.PosterURL == "" || media.BackdropURL == "" || media.Overview == "" {
+ t.Fatalf("path-hint tmdb metadata not enriched: %#v", media)
+ }
+ if media.PosterURL != "https://image.tmdb.org/t/p/w500/poster-fast11.jpg" {
+ t.Fatalf("poster url = %q", media.PosterURL)
+ }
+}
+
+func TestScanCloudLibraryKeepsCloudArtworkWhenEnrichingPathHint(t *testing.T) {
+ tmdb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/movie/755679" {
+ t.Fatalf("unexpected tmdb path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{
+ "id": 755679,
+ "title": "速度与激情11",
+ "overview": "Exact metadata by TMDb ID",
+ "poster_path": "/remote-poster.jpg",
+ "backdrop_path": "/remote-backdrop.jpg",
+ "release_date": "2028-04-07"
+ }`))
+ }))
+ defer tmdb.Close()
+
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.Method {
+ case "PROPFIND":
+ w.Header().Set("Content-Type", "application/xml")
+ w.WriteHeader(http.StatusMultiStatus)
+ switch r.URL.Path {
+ case "/dav/Movies":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/速度与激情11 (2028) {tmdb-755679}
+`))
+ case "/dav/Movies/速度与激情11 (2028) {tmdb-755679}":
+ _, _ = w.Write([]byte(`
+
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028).mkv速度与激情11 (2028).mkv4096
+ /dav/Movies/%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511%20(2028)%20%7Btmdb-755679%7D/poster.jpgposter.jpg1024
+`))
+ default:
+ t.Fatalf("unexpected propfind path %s", r.URL.Path)
+ }
+ case http.MethodGet:
+ if r.URL.Path != "/dav/Movies/速度与激情11 (2028) {tmdb-755679}/poster.jpg" {
+ t.Fatalf("unexpected get path %s", r.URL.Path)
+ }
+ w.Header().Set("Content-Type", "image/jpeg")
+ _, _ = w.Write([]byte("local-cloud-poster"))
+ default:
+ t.Fatalf("unexpected method %s", r.Method)
+ }
+ }))
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}, &model.APIConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "url": upstream.URL,
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = tmdb.URL
+ cfg.Secrets.TMDbImageProxy = "https://image.tmdb.org/t/p"
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+ scanner := NewScannerService(cfg, log, repos, NewHub(log), nil, scraper)
+ scanner.SetStorageConfig(storage)
+ imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
+ scanner.SetImageProxy(imageProxy)
+
+ res, err := scanner.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("scan cloud: %v", err)
+ }
+ if res.Added != 1 || res.LocalMetadata != 1 {
+ t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ wantPoster := "/api/img/cloud/openlist?ref=%2FMovies%2F%E9%80%9F%E5%BA%A6%E4%B8%8E%E6%BF%80%E6%83%8511+%282028%29+%7Btmdb-755679%7D%2Fposter.jpg"
+ if media.PosterURL != wantPoster {
+ t.Fatalf("poster url = %q, want local cloud poster %q", media.PosterURL, wantPoster)
+ }
+ if media.BackdropURL != "https://image.tmdb.org/t/p/w1280/remote-backdrop.jpg" || media.Overview == "" {
+ t.Fatalf("external enrichment should still fill missing fields: %#v", media)
+ }
+ rec := httptest.NewRecorder()
+ if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Movies/速度与激情11 (2028) {tmdb-755679}/poster.jpg") {
+ t.Fatal("local cloud poster should be cached during enriched scan")
+ }
+ if got := rec.Body.String(); got != "local-cloud-poster" {
+ t.Fatalf("cached poster body = %q", got)
+ }
+}
diff --git a/internal/service/scanner_cloud_paths.go b/internal/service/scanner_cloud_paths.go
new file mode 100644
index 0000000..556671f
--- /dev/null
+++ b/internal/service/scanner_cloud_paths.go
@@ -0,0 +1,61 @@
+package service
+
+import (
+ "fmt"
+ "path/filepath"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func parseCloudLibraryPath(raw string) (typ, dirID string, ok bool) {
+ info, ok := ParseCloudLibraryMount(raw)
+ if !ok {
+ return "", "", false
+ }
+ return info.Provider, info.ScanDir, true
+}
+
+func cloudEntryRef(typ, id, pickCode string) string {
+ if typ == "cloud115" && strings.TrimSpace(pickCode) != "" {
+ return strings.TrimSpace(pickCode)
+ }
+ return strings.TrimSpace(id)
+}
+
+func cloudMediaPath(typ, ref string) string {
+ return "cloud://" + strings.TrimSpace(typ) + "/" + strings.TrimLeft(strings.TrimSpace(ref), "/")
+}
+
+func cloudMediaDedupeKey(lib *model.Library, dirID, name string, size int64) string {
+ base := strings.TrimSpace(strings.TrimSuffix(filepath.Base(name), filepath.Ext(name)))
+ if base == "" {
+ return ""
+ }
+ season, episode := ParseEpisode(name)
+ title, year := CleanQuery(name)
+ title = normalizeCloudDedupeText(title)
+ if (season > 0 || episode > 0) && title != "" {
+ return fmt.Sprintf("episode:%s:%s:%d:%d:%d", strings.ToLower(strings.TrimSpace(lib.Type)), title, year, season, episode)
+ }
+ if (season > 0 || episode > 0) && title == "" {
+ return fmt.Sprintf("episode-dir:%s:%s:%d:%d:%d", strings.ToLower(strings.TrimSpace(lib.Type)), normalizeCloudDedupeText(dirID), season, episode, size)
+ }
+ return fmt.Sprintf("file:%s:%d", normalizeCloudDedupeText(base), size)
+}
+
+func normalizeCloudDedupeText(value string) string {
+ value = strings.ToLower(strings.TrimSpace(value))
+ if value == "" {
+ return ""
+ }
+ fields := strings.FieldsFunc(value, func(r rune) bool {
+ switch r {
+ case '.', '_', '-', ' ', '\t', '/', '\\', '[', ']', '(', ')':
+ return true
+ default:
+ return false
+ }
+ })
+ return strings.Join(fields, " ")
+}
diff --git a/internal/service/scanner_cloud_probe.go b/internal/service/scanner_cloud_probe.go
new file mode 100644
index 0000000..df0f1c9
--- /dev/null
+++ b/internal/service/scanner_cloud_probe.go
@@ -0,0 +1,122 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) probeCloudMediaAsync(task cloudMediaProbeTask) {
+ defer func() {
+ s.cloudMediaProbeMu.Lock()
+ delete(s.cloudMediaProbing, task.path)
+ s.cloudMediaProbeMu.Unlock()
+ }()
+ ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
+ defer cancel()
+ probe, err := s.probeCloudFileMetadata(ctx, task.typ, task.ref)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("cloud media async probe failed", zap.String("provider", task.typ), zap.String("path", task.path), zap.Error(err))
+ }
+ s.cloudMediaProbeMu.Lock()
+ if s.cloudMediaProbeBackoff == nil {
+ s.cloudMediaProbeBackoff = make(map[string]time.Time)
+ }
+ s.cloudMediaProbeBackoff[task.path] = time.Now().Add(cloudMediaProbeFailureBackoff)
+ s.cloudMediaProbeMu.Unlock()
+ return
+ }
+ updates := probeResultUpdates(probe)
+ if len(updates) == 0 {
+ return
+ }
+ if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("path = ?", task.path).Updates(updates).Error; err != nil {
+ if s.log != nil {
+ s.log.Debug("update cloud media track metadata failed", zap.String("path", task.path), zap.Error(err))
+ }
+ return
+ }
+ s.cloudMediaProbeMu.Lock()
+ delete(s.cloudMediaProbeBackoff, task.path)
+ s.cloudMediaProbeMu.Unlock()
+ if s.hub != nil {
+ s.hub.Publish("scan", map[string]any{
+ "path": task.path,
+ "cloud": true,
+ "track_probed": true,
+ "duration_sec": probe.DurationSec,
+ "video_codec": probe.VideoCodec,
+ "audio_codec": probe.AudioCodec,
+ "width": probe.Width,
+ "height": probe.Height,
+ "probe_message": "云盘媒体轨道元数据已后台补齐",
+ })
+ }
+}
+
+func (s *ScannerService) ffprobeWorkerCount() int {
+ if s == nil || s.cfg == nil {
+ return 1
+ }
+ return normalizeFFprobeMaxConcurrent(s.cfg.App.FFprobeMaxConcurrent)
+}
+
+func (s *ScannerService) cloudScanWorkerCount() int {
+ if s == nil || s.cfg == nil {
+ return 4
+ }
+ return normalizeCloudScanMaxConcurrent(s.cfg.App.CloudScanMaxConcurrent)
+}
+
+func normalizeCloudScanMaxConcurrent(n int) int {
+ if n <= 0 {
+ return 1
+ }
+ if n > 16 {
+ return 16
+ }
+ return n
+}
+
+func (s *ScannerService) probeCloudFileMetadata(ctx context.Context, typ, ref string) (*ProbeResult, error) {
+ if s == nil || s.probe == nil || s.storage == nil {
+ return nil, errors.New("cloud probe unavailable")
+ }
+ link, err := s.storage.CloudResolve(ctx, typ, ref, "")
+ if err != nil {
+ return nil, err
+ }
+ return s.probe.ProbeHTTP(ctx, link.URL, link.Headers)
+}
+
+func probeResultUpdates(probe *ProbeResult) map[string]any {
+ updates := map[string]any{}
+ if probe == nil {
+ return updates
+ }
+ if probe.DurationSec > 0 {
+ updates["duration_sec"] = probe.DurationSec
+ }
+ if probe.Width > 0 {
+ updates["width"] = probe.Width
+ }
+ if probe.Height > 0 {
+ updates["height"] = probe.Height
+ }
+ if strings.TrimSpace(probe.VideoCodec) != "" {
+ updates["video_codec"] = probe.VideoCodec
+ }
+ if strings.TrimSpace(probe.AudioCodec) != "" {
+ updates["audio_codec"] = probe.AudioCodec
+ }
+ if probe.Container != "" {
+ updates["container"] = probe.Container
+ }
+ return updates
+}
diff --git a/internal/service/scanner_cloud_scan.go b/internal/service/scanner_cloud_scan.go
new file mode 100644
index 0000000..c4bf692
--- /dev/null
+++ b/internal/service/scanner_cloud_scan.go
@@ -0,0 +1,231 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "path/filepath"
+ "strings"
+ "sync"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) scanCloudLibrary(ctx context.Context, lib *model.Library, mount CloudMountInfo, autoScrape bool) (*ScanResult, error) {
+ res := &ScanResult{LibraryID: lib.ID}
+ if s.storage == nil {
+ return res, fmt.Errorf("cloud storage service unavailable")
+ }
+
+ cfg, err := s.repo.StorageConfig.Get(ctx, mount.Provider)
+ if err != nil || cfg == nil {
+ return res, fmt.Errorf("storage config not found: %s", mount.Provider)
+ }
+ if !cfg.Enabled {
+ return res, fmt.Errorf("storage %s is disabled", mount.Provider)
+ }
+ typ := mount.Provider
+ rootDir := mount.ScanDir
+ rootDisplayDir := mount.DisplayDir
+ autoCategoryRoot := cloudRootMountNeedsAutoCategory(mount)
+ scopeIDs := s.cloudScanLibraryScopeIDs(ctx, lib, mount)
+ seen := make(map[string]struct{})
+ seenRefs := make(map[string]struct{})
+ candidates := make([]cloudCandidate, 0, 256)
+ candidateByKey := make(map[string]int)
+ visitedDirs := map[string]struct{}{}
+ progress := newCloudScanProgressState()
+ var stateMu sync.Mutex
+ progress.publish(s, lib.ID, res, "listing", true)
+ var walkWG sync.WaitGroup
+ var walkErr error
+ var walkErrOnce sync.Once
+ setWalkErr := func(err error) {
+ if err != nil {
+ walkErrOnce.Do(func() {
+ walkErr = err
+ })
+ }
+ }
+ listSlots := make(chan struct{}, s.cloudScanWorkerCount())
+ var walkCloud func(dirID, displayDir string, inheritedMeta *LocalMetadata) error
+ walkCloud = func(dirID, displayDir string, inheritedMeta *LocalMetadata) error {
+ defer walkWG.Done()
+ if err := ctx.Err(); err != nil {
+ setWalkErr(err)
+ return err
+ }
+ stateMu.Lock()
+ if _, ok := visitedDirs[dirID]; ok {
+ stateMu.Unlock()
+ return nil
+ }
+ visitedDirs[dirID] = struct{}{}
+ stateMu.Unlock()
+
+ select {
+ case listSlots <- struct{}{}:
+ defer func() { <-listSlots }()
+ case <-ctx.Done():
+ setWalkErr(ctx.Err())
+ return ctx.Err()
+ }
+ entries, err := s.storage.CloudList(ctx, typ, dirID)
+ if err != nil {
+ if dirID != rootDir {
+ progress.addSkipped(res)
+ s.log.Warn("skip inaccessible cloud directory",
+ zap.String("library_id", lib.ID),
+ zap.String("provider", typ),
+ zap.String("dir", dirID),
+ zap.Error(err))
+ return nil
+ }
+ setWalkErr(err)
+ return err
+ }
+ progress.publish(s, lib.ID, res, "listing", progress.markDirVisited())
+ sidecars := newCloudSidecarSet(typ, entries)
+ dirMeta := s.cloudDirectoryMetadata(ctx, typ, displayDir, sidecars, inheritedMeta)
+ s.cacheCloudMetadataArtworkNow(ctx, dirMeta)
+ for _, entry := range entries {
+ select {
+ case <-ctx.Done():
+ setWalkErr(ctx.Err())
+ return ctx.Err()
+ default:
+ }
+ if entry.IsDir {
+ if strings.TrimSpace(entry.ID) != "" {
+ walkWG.Add(1)
+ go func(childID, childDisplay string, childMeta *LocalMetadata) {
+ _ = walkCloud(childID, childDisplay, childMeta)
+ }(entry.ID, joinCloudDisplayPath(displayDir, entry.Name), dirMeta)
+ }
+ continue
+ }
+ ext := strings.ToLower(filepath.Ext(entry.Name))
+ if _, ok := videoExtensions[ext]; !ok {
+ continue
+ }
+ ref := cloudEntryRef(typ, entry.ID, entry.PickCode)
+ if ref == "" {
+ progress.addSkipped(res)
+ continue
+ }
+ stateMu.Lock()
+ if _, ok := seenRefs[ref]; ok {
+ stateMu.Unlock()
+ progress.addSkipped(res)
+ continue
+ }
+ seenRefs[ref] = struct{}{}
+ stateMu.Unlock()
+ progress.publish(s, lib.ID, res, "listing", progress.markFileDiscovered())
+ displayPath := joinCloudDisplayPath(displayDir, entry.Name)
+ path := cloudMediaPath(typ, displayPath)
+ localMeta := s.cloudFileMetadata(ctx, typ, displayPath, entry.Name, sidecars, dirMeta, librarySupportsSeasons(lib))
+ localMeta = s.enrichCloudMetadataFromExternalIDs(ctx, lib, path, localMeta)
+ if localMeta != nil {
+ s.cacheCloudMetadataArtworkNow(ctx, localMeta)
+ }
+ candidate := cloudCandidate{
+ ref: ref,
+ name: entry.Name,
+ size: entry.Size,
+ path: path,
+ localMeta: localMeta,
+ }
+ if autoCategoryRoot {
+ candidate.categoryDisplayDir = cloudAutoCategoryDisplayDirForMediaPath(path)
+ }
+ key := cloudMediaDedupeKey(lib, displayDir, entry.Name, entry.Size)
+ stateMu.Lock()
+ if key != "" {
+ if prevIndex, ok := candidateByKey[key]; ok {
+ if candidate.size > candidates[prevIndex].size {
+ candidates[prevIndex] = candidate
+ }
+ stateMu.Unlock()
+ progress.addSkipped(res)
+ continue
+ }
+ candidateByKey[key] = len(candidates)
+ }
+ candidates = append(candidates, candidate)
+ stateMu.Unlock()
+ }
+ return nil
+ }
+ walkWG.Add(1)
+ go func() {
+ _ = walkCloud(rootDir, rootDisplayDir, nil)
+ }()
+ walkWG.Wait()
+ if walkErr != nil {
+ return res, walkErr
+ }
+ if err := ctx.Err(); err != nil {
+ return res, err
+ }
+ existingMedia, err := s.existingCloudMediaSnapshotForLibraries(ctx, scopeIDs)
+ if err != nil {
+ s.log.Warn("load existing cloud media snapshot failed", zap.String("library_id", lib.ID), zap.Error(err))
+ existingMedia = nil
+ }
+ sortCloudCandidatesByRefreshPriority(candidates, existingMedia)
+ writeBatch := newLocalMediaWriteBatch(s, ctx, res, 100)
+ probeBudget := maxCloudMediaProbeQueuePerScan
+ targetLibs := map[string]*model.Library{"": lib}
+ touchedLibraryIDs := []string{}
+ for _, candidate := range candidates {
+ select {
+ case <-ctx.Done():
+ return res, ctx.Err()
+ default:
+ }
+ targetLib := lib
+ if candidate.categoryDisplayDir != "" {
+ if cached, ok := targetLibs[candidate.categoryDisplayDir]; ok {
+ targetLib = cached
+ } else if categoryLib, err := s.ensureCloudAutoCategoryLibrary(ctx, lib, typ, candidate.categoryDisplayDir); err == nil && categoryLib != nil {
+ targetLib = categoryLib
+ targetLibs[candidate.categoryDisplayDir] = categoryLib
+ scopeIDs = appendUniqueLibraryIDs(scopeIDs, categoryLib.ID)
+ } else if err != nil {
+ s.log.Warn("ensure cloud auto category library failed",
+ zap.String("library_id", lib.ID),
+ zap.String("provider", typ),
+ zap.String("category", candidate.categoryDisplayDir),
+ zap.Error(err))
+ }
+ }
+ touchedLibraryIDs = appendUniqueLibraryIDs(touchedLibraryIDs, targetLib.ID)
+ seen[candidate.path] = struct{}{}
+ s.ingestCloudFile(ctx, targetLib, typ, candidate.ref, candidate.path, candidate.name, candidate.size, candidate.localMeta, existingMedia, writeBatch, &probeBudget, res)
+ progress.publish(s, lib.ID, res, "importing", res.Visited == 1 || res.Visited%100 == 0)
+ }
+ writeBatch.Flush()
+ removed, err := s.pruneMissingCloudMediaForLibraries(ctx, scopeIDs, seen)
+ if err != nil {
+ s.log.Warn("prune missing cloud media failed", zap.String("library_id", lib.ID), zap.Error(err))
+ } else {
+ res.Removed = removed
+ }
+ publishCloudScanFinished(s, lib.ID, res, progress)
+ s.invalidateMediaCache(ctx)
+ for _, targetID := range appendUniqueLibraryIDs(touchedLibraryIDs, lib.ID) {
+ s.maybeGenerateSTRMAfterScan(targetID)
+ }
+ if scanHasImportChanges(res) && autoScrape && s.scraper != nil && s.scraper.AnyEnabled() && s.autoScrapeEnabled(ctx) {
+ for _, targetID := range appendUniqueLibraryIDs(touchedLibraryIDs, lib.ID) {
+ s.startAutoScrape(ctx, targetID)
+ }
+ }
+ return res, nil
+}
+
+func scanHasImportChanges(res *ScanResult) bool {
+ return res != nil && (res.Added > 0 || res.Updated > 0 || res.Removed > 0)
+}
diff --git a/internal/service/scanner_cloud_scan_progress.go b/internal/service/scanner_cloud_scan_progress.go
new file mode 100644
index 0000000..120152e
--- /dev/null
+++ b/internal/service/scanner_cloud_scan_progress.go
@@ -0,0 +1,169 @@
+package service
+
+import (
+ "sort"
+ "sync"
+ "time"
+)
+
+type cloudCandidate struct {
+ ref string
+ name string
+ size int64
+ path string
+ categoryDisplayDir string
+ localMeta *LocalMetadata
+}
+
+type cloudScanProgressState struct {
+ mu sync.Mutex
+ startedAt time.Time
+ lastProgress time.Time
+ dirsVisited int
+ filesDiscovered int
+}
+
+type cloudScanProgressSnapshot struct {
+ dirsVisited int
+ filesDiscovered int
+ visited int
+ added int
+ updated int
+ skipped int
+ removed int64
+ elapsed time.Duration
+}
+
+func newCloudScanProgressState() *cloudScanProgressState {
+ return &cloudScanProgressState{startedAt: time.Now()}
+}
+
+func (p *cloudScanProgressState) markDirVisited() bool {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ p.dirsVisited++
+ return p.dirsVisited == 1 || p.dirsVisited%20 == 0
+}
+
+func (p *cloudScanProgressState) markFileDiscovered() bool {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ p.filesDiscovered++
+ return p.filesDiscovered%100 == 0
+}
+
+func (p *cloudScanProgressState) addSkipped(res *ScanResult) {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ res.Skipped++
+}
+
+func (p *cloudScanProgressState) publish(s *ScannerService, libraryID string, res *ScanResult, stage string, force bool) {
+ if s == nil || s.hub == nil {
+ return
+ }
+ snap, ok := p.snapshotForProgress(res, force)
+ if !ok {
+ return
+ }
+ filesPerSecond := snap.filesPerSecond()
+ s.updateCloudScanProgress(libraryID, stage, snap.dirsVisited, snap.filesDiscovered, snap.visited, snap.added, snap.updated, snap.skipped, snap.removed, filesPerSecond)
+ s.hub.Publish("scan", map[string]any{
+ "library_id": libraryID,
+ "cloud": true,
+ "stage": stage,
+ "dirs": snap.dirsVisited,
+ "discovered": snap.filesDiscovered,
+ "visited": snap.visited,
+ "added": snap.added,
+ "updated": snap.updated,
+ "skipped": snap.skipped,
+ "elapsed_seconds": int(snap.elapsed.Seconds()),
+ "files_per_second": filesPerSecond,
+ "estimate_message": "云盘接口不提供总文件数,剩余时间会随目录大小和网盘响应速度变化",
+ })
+}
+
+func (p *cloudScanProgressState) snapshotForProgress(res *ScanResult, force bool) (cloudScanProgressSnapshot, bool) {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ if !force && time.Since(p.lastProgress) < 2*time.Second {
+ return cloudScanProgressSnapshot{}, false
+ }
+ p.lastProgress = time.Now()
+ return p.snapshotLocked(res), true
+}
+
+func (p *cloudScanProgressState) finalSnapshot(res *ScanResult) cloudScanProgressSnapshot {
+ p.mu.Lock()
+ defer p.mu.Unlock()
+ return p.snapshotLocked(res)
+}
+
+func (p *cloudScanProgressState) snapshotLocked(res *ScanResult) cloudScanProgressSnapshot {
+ snap := cloudScanProgressSnapshot{
+ dirsVisited: p.dirsVisited,
+ filesDiscovered: p.filesDiscovered,
+ elapsed: time.Since(p.startedAt),
+ }
+ if res != nil {
+ snap.visited = res.Visited
+ snap.added = res.Added
+ snap.updated = res.Updated
+ snap.skipped = res.Skipped
+ snap.removed = res.Removed
+ }
+ return snap
+}
+
+func (s cloudScanProgressSnapshot) filesPerSecond() float64 {
+ processed := s.filesDiscovered
+ if s.visited > processed {
+ processed = s.visited
+ }
+ if s.elapsed.Seconds() <= 0 {
+ return 0
+ }
+ return float64(processed) / s.elapsed.Seconds()
+}
+
+func publishCloudScanFinished(s *ScannerService, libraryID string, res *ScanResult, progress *cloudScanProgressState) {
+ if s == nil || s.hub == nil || progress == nil {
+ return
+ }
+ snap := progress.finalSnapshot(res)
+ s.hub.Publish("scan", map[string]any{
+ "library_id": libraryID,
+ "finished": true,
+ "visited": res.Visited,
+ "added": res.Added,
+ "updated": res.Updated,
+ "skipped": res.Skipped,
+ "removed": res.Removed,
+ "error_count": res.ErrorCount,
+ "errors": res.Errors,
+ "discovered": snap.filesDiscovered,
+ "dirs": snap.dirsVisited,
+ "elapsed_seconds": int(snap.elapsed.Seconds()),
+ "cloud": true,
+ })
+}
+
+func sortCloudCandidatesByRefreshPriority(candidates []cloudCandidate, existingMedia map[string]existingCloudMedia) {
+ if existingMedia == nil {
+ return
+ }
+ priority := func(candidate cloudCandidate) int {
+ existing, ok := existingMedia[candidate.path]
+ if !ok {
+ return 2
+ }
+ if cloudTrackMetadataMissing(existing) || cloudMetadataNeedsRefresh(existing, candidate.localMeta) {
+ return 0
+ }
+ return 1
+ }
+ sort.SliceStable(candidates, func(i, j int) bool {
+ return priority(candidates[i]) < priority(candidates[j])
+ })
+}
diff --git a/internal/service/scanner_cloud_status.go b/internal/service/scanner_cloud_status.go
new file mode 100644
index 0000000..32ff11e
--- /dev/null
+++ b/internal/service/scanner_cloud_status.go
@@ -0,0 +1,265 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) beginCloudScan(ctx context.Context, lib *model.Library, mount CloudMountInfo) (context.Context, func(*ScanResult, error), error) {
+ if s == nil || lib == nil {
+ return ctx, func(*ScanResult, error) {}, nil
+ }
+ s.cloudScanMu.Lock()
+ if s.cloudScans == nil {
+ s.cloudScans = make(map[string]*cloudScanEntry)
+ }
+ if entry := s.cloudScans[lib.ID]; cloudScanBlocksBegin(entry) {
+ s.cloudScanMu.Unlock()
+ return ctx, nil, ErrCloudScanAlreadyRunning
+ }
+ runCtx, cancel := context.WithCancel(ctx)
+ s.cloudScans[lib.ID] = newCloudScanEntry(lib.ID, mount.Provider, cancel)
+ s.cloudScanMu.Unlock()
+
+ finish := func(res *ScanResult, err error) {
+ s.finishCloudScan(lib, mount, res, err)
+ }
+ return runCtx, finish, nil
+}
+
+func newCloudScanEntry(libraryID, provider string, cancel context.CancelFunc) *cloudScanEntry {
+ now := time.Now()
+ return &cloudScanEntry{
+ status: CloudScanStatus{
+ LibraryID: libraryID,
+ Provider: provider,
+ Stage: "listing",
+ State: "running",
+ StartedAt: now,
+ UpdatedAt: now,
+ ResumeHint: "中断后再次点击扫描会从头遍历,但已入库媒体会去重更新,只补齐缺失项。",
+ Estimate: "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度。",
+ },
+ cancel: cancel,
+ }
+}
+
+func (s *ScannerService) finishCloudScan(lib *model.Library, mount CloudMountInfo, res *ScanResult, err error) {
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ current := s.cloudScans[lib.ID]
+ if current == nil {
+ return
+ }
+ applyCloudScanResult(¤t.status, res)
+ current.status.UpdatedAt = time.Now()
+ current.status.FinishedAt = current.status.UpdatedAt
+ current.cancel = nil
+ applyCloudScanCompletion(¤t.status, err)
+ s.publishCloudScanFinished(lib.ID, mount.Provider, current.status)
+ s.notifyScanFinished(lib, res, err, true)
+}
+
+func applyCloudScanResult(status *CloudScanStatus, res *ScanResult) {
+ if status == nil || res == nil {
+ return
+ }
+ status.Visited = res.Visited
+ status.Added = res.Added
+ status.Updated = res.Updated
+ status.Skipped = res.Skipped
+ status.Removed = res.Removed
+ status.ErrorCount = res.ErrorCount
+ status.Errors = append([]string(nil), res.Errors...)
+}
+
+func applyCloudScanCompletion(status *CloudScanStatus, err error) {
+ if status == nil {
+ return
+ }
+ switch {
+ case errors.Is(err, context.Canceled):
+ status.State = "canceled"
+ status.Stage = "canceled"
+ status.Error = ""
+ case errors.Is(err, context.DeadlineExceeded):
+ status.State = "error"
+ status.Stage = "error"
+ status.Error = "扫描超时:" + err.Error()
+ case err != nil:
+ status.State = "error"
+ status.Stage = "error"
+ status.Error = err.Error()
+ default:
+ status.State = "finished"
+ status.Stage = "finished"
+ if status.ErrorCount > 0 {
+ status.Error = fmt.Sprintf("部分文件入库失败:%d 个,详情见 errors", status.ErrorCount)
+ } else {
+ status.Error = ""
+ }
+ }
+}
+
+func (s *ScannerService) publishCloudScanFinished(libraryID, provider string, status CloudScanStatus) {
+ if s == nil || s.hub == nil {
+ return
+ }
+ s.hub.Publish("scan", map[string]any{
+ "library_id": libraryID,
+ "provider": provider,
+ "cloud": true,
+ "finished": true,
+ "state": status.State,
+ "stage": status.Stage,
+ "error": status.Error,
+ "visited": status.Visited,
+ "added": status.Added,
+ "updated": status.Updated,
+ "skipped": status.Skipped,
+ "removed": status.Removed,
+ "error_count": status.ErrorCount,
+ "errors": status.Errors,
+ })
+}
+
+func (s *ScannerService) updateCloudScanProgress(libraryID, stage string, dirs, discovered, visited, added, updated, skipped int, removed int64, filesPerSecond float64) {
+ if s == nil {
+ return
+ }
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ entry := s.cloudScans[libraryID]
+ if entry == nil {
+ return
+ }
+ entry.status.Stage = stage
+ entry.status.UpdatedAt = time.Now()
+ entry.status.Dirs = dirs
+ entry.status.Discovered = discovered
+ entry.status.Visited = visited
+ entry.status.Added = added
+ entry.status.Updated = updated
+ entry.status.Skipped = skipped
+ entry.status.Removed = removed
+ entry.status.FilesPerSecond = filesPerSecond
+}
+
+func (s *ScannerService) acquireCloudScanSlot(ctx context.Context, libraryID string) (func(), error) {
+ if s == nil {
+ return func() {}, nil
+ }
+ s.cloudScanMu.Lock()
+ if s.cloudSlots == nil {
+ s.cloudSlots = make(chan struct{}, 1)
+ }
+ slots := s.cloudSlots
+ if entry := s.cloudScans[libraryID]; entry != nil {
+ entry.status.Stage = "queued"
+ entry.status.UpdatedAt = time.Now()
+ }
+ s.cloudScanMu.Unlock()
+
+ select {
+ case slots <- struct{}{}:
+ s.cloudScanMu.Lock()
+ if entry := s.cloudScans[libraryID]; entry != nil && entry.status.State == "running" {
+ entry.status.Stage = "listing"
+ entry.status.UpdatedAt = time.Now()
+ }
+ s.cloudScanMu.Unlock()
+ return func() { <-slots }, nil
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ }
+}
+
+// CloudScanStatuses returns the current or most recent status per cloud library.
+func (s *ScannerService) CloudScanStatuses() []CloudScanStatus {
+ if s == nil {
+ return nil
+ }
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ out := make([]CloudScanStatus, 0, len(s.cloudScans))
+ for _, entry := range s.cloudScans {
+ out = append(out, entry.status)
+ }
+ return out
+}
+
+func (s *ScannerService) CancelCloudScan(libraryID string) bool {
+ if s == nil || strings.TrimSpace(libraryID) == "" {
+ return false
+ }
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ return cancelCloudScanEntry(s.cloudScans[libraryID])
+}
+
+func (s *ScannerService) CancelAllCloudScans() int {
+ if s == nil {
+ return 0
+ }
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ cancelled := 0
+ for _, entry := range s.cloudScans {
+ if cancelCloudScanEntry(entry) {
+ cancelled++
+ }
+ }
+ return cancelled
+}
+
+func (s *ScannerService) CancelCloudScansForProvider(provider string) int {
+ if s == nil {
+ return 0
+ }
+ provider = strings.TrimSpace(provider)
+ if provider == "" {
+ return 0
+ }
+ s.cloudScanMu.Lock()
+ defer s.cloudScanMu.Unlock()
+ cancelled := 0
+ for _, entry := range s.cloudScans {
+ if entry == nil || entry.status.Provider != provider {
+ continue
+ }
+ if cancelCloudScanEntry(entry) {
+ cancelled++
+ }
+ }
+ return cancelled
+}
+
+func cancelCloudScanEntry(entry *cloudScanEntry) bool {
+ if !cloudScanActive(entry) {
+ return false
+ }
+ entry.status.State = "canceling"
+ entry.status.Stage = "canceling"
+ entry.status.UpdatedAt = time.Now()
+ if entry.cancel != nil {
+ entry.cancel()
+ return true
+ }
+ entry.status.State = "canceled"
+ entry.status.Stage = "canceled"
+ entry.status.FinishedAt = time.Now()
+ return true
+}
+
+func cloudScanActive(entry *cloudScanEntry) bool {
+ return entry != nil && (entry.status.State == "running" || entry.status.State == "queued" || entry.status.State == "canceling")
+}
+
+func cloudScanBlocksBegin(entry *cloudScanEntry) bool {
+ return entry != nil && (entry.status.State == "running" || entry.status.State == "canceling")
+}
diff --git a/internal/service/scanner_cloud_status_test.go b/internal/service/scanner_cloud_status_test.go
new file mode 100644
index 0000000..5468e71
--- /dev/null
+++ b/internal/service/scanner_cloud_status_test.go
@@ -0,0 +1,57 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "testing"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestBeginCloudScanAllowsQueuedEntryToStart(t *testing.T) {
+ scanner := &ScannerService{
+ cloudScans: map[string]*cloudScanEntry{
+ "lib-1": {status: CloudScanStatus{LibraryID: "lib-1", Provider: "openlist", State: "queued", Stage: "queued"}},
+ },
+ }
+ lib := &model.Library{Base: model.Base{ID: "lib-1"}, Name: "Movies"}
+ mount := CloudMountInfo{Provider: "openlist"}
+
+ _, finish, err := scanner.beginCloudScan(context.Background(), lib, mount)
+ if err != nil {
+ t.Fatalf("queued scan should be allowed to start, got %v", err)
+ }
+ if finish == nil {
+ t.Fatal("finish callback should not be nil")
+ }
+ statuses := scanner.CloudScanStatuses()
+ if len(statuses) != 1 || statuses[0].State != "running" || statuses[0].Stage != "listing" {
+ t.Fatalf("status after begin = %#v, want running/listing", statuses)
+ }
+
+ finish(&ScanResult{Visited: 5, Added: 2, Updated: 1, ErrorCount: 1, Errors: []string{"bad file"}}, nil)
+ statuses = scanner.CloudScanStatuses()
+ if len(statuses) != 1 || statuses[0].State != "finished" || statuses[0].Visited != 5 || statuses[0].Added != 2 || statuses[0].ErrorCount != 1 {
+ t.Fatalf("status after finish = %#v", statuses)
+ }
+ if statuses[0].Error == "" {
+ t.Fatal("finished scan with error_count should keep summary error text")
+ }
+}
+
+func TestBeginCloudScanRejectsRunningEntry(t *testing.T) {
+ scanner := &ScannerService{
+ cloudScans: map[string]*cloudScanEntry{
+ "lib-1": {status: CloudScanStatus{LibraryID: "lib-1", Provider: "openlist", State: "running", Stage: "listing"}},
+ },
+ }
+ lib := &model.Library{Base: model.Base{ID: "lib-1"}, Name: "Movies"}
+
+ _, finish, err := scanner.beginCloudScan(context.Background(), lib, CloudMountInfo{Provider: "openlist"})
+ if !errors.Is(err, ErrCloudScanAlreadyRunning) {
+ t.Fatalf("err = %v, want ErrCloudScanAlreadyRunning", err)
+ }
+ if finish != nil {
+ t.Fatal("finish callback should be nil when begin is rejected")
+ }
+}
diff --git a/internal/service/scanner_cloud_test.go b/internal/service/scanner_cloud_test.go
index b759277..6de6fd6 100644
--- a/internal/service/scanner_cloud_test.go
+++ b/internal/service/scanner_cloud_test.go
@@ -21,53 +21,93 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
-func TestScanCloudLibraryImportsRecursivePlayableMedia(t *testing.T) {
- empty := false
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/sort" {
- t.Fatalf("unexpected path %s", r.URL.Path)
+type openListTestEntry struct {
+ Name string
+ Size int64
+ IsDir bool
+}
+
+func newOpenListAPIServer(t *testing.T, list func(path string, page, perPage int) ([]openListTestEntry, int)) *httptest.Server {
+ t.Helper()
+ return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/api/fs/list" {
+ t.Fatalf("unexpected openlist api request %s", r.URL.Path)
+ }
+ var in struct {
+ Path string `json:"path"`
+ Page int `json:"page"`
+ PerPage int `json:"per_page"`
+ }
+ if err := json.NewDecoder(r.Body).Decode(&in); err != nil {
+ t.Fatalf("decode openlist list request: %v", err)
+ }
+ if in.Path == "" {
+ in.Path = "/"
+ }
+ if in.Page <= 0 {
+ in.Page = 1
+ }
+ if in.PerPage <= 0 {
+ in.PerPage = 500
+ }
+ entries, total := list(in.Path, in.Page, in.PerPage)
+ content := make([]map[string]any, 0, len(entries))
+ for _, entry := range entries {
+ content = append(content, map[string]any{
+ "name": entry.Name,
+ "size": entry.Size,
+ "is_dir": entry.IsDir,
+ })
}
w.Header().Set("Content-Type", "application/json")
- if empty {
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[]}}`))
- return
- }
- switch r.URL.Query().Get("pdir_fid") {
- case "0":
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"d1","file_name":"Movies","dir":true,"size":0},
- {"fid":"f1","file_name":"Root.Movie.2024.mkv","dir":false,"size":123}
- ]}}`))
- case "d1":
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"f2","file_name":"Nested.Show.S01E02.mp4","dir":false,"size":456}
- ]}}`))
- default:
- t.Fatalf("unexpected pdir_fid %q", r.URL.Query().Get("pdir_fid"))
- }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "code": 200,
+ "message": "success",
+ "data": map[string]any{
+ "content": content,
+ "total": total,
+ },
+ })
}))
+}
+
+func TestScanCloudLibraryImportsRecursivePlayableMedia(t *testing.T) {
+ empty := false
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
+ if empty {
+ return nil, 0
+ }
+ switch path {
+ case "/":
+ return []openListTestEntry{
+ {Name: "Movies", IsDir: true},
+ {Name: "Root.Movie.2024.mkv", Size: 123},
+ }, 2
+ case "/Movies":
+ return []openListTestEntry{
+ {Name: "Nested.Show.S01E02.mp4", Size: 456},
+ }, 1
+ default:
+ t.Fatalf("unexpected openlist path %q", path)
+ return nil, 0
+ }
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "tv", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
@@ -88,13 +128,13 @@ func TestScanCloudLibraryImportsRecursivePlayableMedia(t *testing.T) {
if len(rows) != 2 {
t.Fatalf("media rows = %d, want 2: %#v", len(rows), rows)
}
- if rows[0].Path != "cloud://quark/Movies/Nested.Show.S01E02.mp4" || !strings.Contains(rows[0].STRMURL, "ref=f2") {
+ if rows[0].Path != "cloud://openlist/Movies/Nested.Show.S01E02.mp4" || !strings.Contains(rows[0].STRMURL, "ref=%2FMovies%2FNested.Show.S01E02.mp4") {
t.Fatalf("nested media path/strm wrong: path=%q strm=%q", rows[0].Path, rows[0].STRMURL)
}
if rows[0].SeasonNum != 1 || rows[0].EpisodeNum != 2 {
t.Fatalf("nested episode metadata wrong: %#v", rows[0])
}
- if rows[1].Path != "cloud://quark/Root.Movie.2024.mkv" || rows[1].STRMURL != "/api/cloud/play/quark?ref=f1" {
+ if rows[1].Path != "cloud://openlist/Root.Movie.2024.mkv" || rows[1].STRMURL != "/api/cloud/play/openlist?ref=%2FRoot.Movie.2024.mkv" {
t.Fatalf("root media path/strm wrong: path=%q strm=%q", rows[0].Path, rows[0].STRMURL)
}
@@ -126,24 +166,163 @@ func TestScanCloudLibraryImportsRecursivePlayableMedia(t *testing.T) {
}
}
+func TestScanRootCloudLibraryCreatesAutoCategoryLibraries(t *testing.T) {
+ empty := false
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
+ if empty {
+ return nil, 0
+ }
+ switch path {
+ case "/":
+ return []openListTestEntry{
+ {Name: "电视剧", IsDir: true},
+ {Name: "电影", IsDir: true},
+ {Name: "国漫", IsDir: true},
+ }, 3
+ case "/电视剧":
+ return []openListTestEntry{{Name: "欧美剧", IsDir: true}}, 1
+ case "/电视剧/欧美剧":
+ return []openListTestEntry{{Name: "The Show", IsDir: true}}, 1
+ case "/电视剧/欧美剧/The Show":
+ return []openListTestEntry{{Name: "The.Show.S01E01.mkv", Size: 101}}, 1
+ case "/电影":
+ return []openListTestEntry{{Name: "华语电影", IsDir: true}}, 1
+ case "/电影/华语电影":
+ return []openListTestEntry{{Name: "Movie.2024.mkv", Size: 202}}, 1
+ case "/国漫":
+ return []openListTestEntry{{Name: "剑来", IsDir: true}}, 1
+ case "/国漫/剑来":
+ return []openListTestEntry{{Name: "剑来.S01E01.mkv", Size: 303}}, 1
+ default:
+ t.Fatalf("unexpected openlist path %q", path)
+ return nil, 0
+ }
+ })
+ defer upstream.Close()
+
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
+ repos := repository.New(db)
+ log := zap.NewNop()
+ storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "server": upstream.URL,
+ "token": "openlist-token",
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ root := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &root); err != nil {
+ t.Fatal(err)
+ }
+ scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
+ scanner.SetStorageConfig(storage)
+
+ res, err := scanner.ScanLibrary(t.Context(), root.ID)
+ if err != nil {
+ t.Fatalf("scan root cloud: %v", err)
+ }
+ if res.Visited != 3 || res.Added != 3 {
+ t.Fatalf("scan result = %#v, want visited=3 added=3", res)
+ }
+
+ libs, err := repos.Library.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ byDisplayDir := map[string]model.Library{}
+ for _, lib := range libs {
+ if !CloudLibraryAutoCategory(lib) {
+ continue
+ }
+ info, ok := ParseCloudLibraryMount(lib.Path)
+ if !ok {
+ t.Fatalf("auto category path did not parse: %q", lib.Path)
+ }
+ byDisplayDir[info.DisplayDir] = lib
+ }
+ wantTypes := map[string]string{
+ "电视剧/欧美剧": "tv",
+ "电影/华语电影": "movie",
+ "动漫/国漫": "anime",
+ }
+ for dir, wantType := range wantTypes {
+ lib, ok := byDisplayDir[dir]
+ if !ok {
+ t.Fatalf("missing auto category library %q; got %#v", dir, byDisplayDir)
+ }
+ if lib.Type != wantType {
+ t.Fatalf("auto category %s type = %s, want %s", dir, lib.Type, wantType)
+ }
+ }
+
+ var rows []model.Media
+ if err := repos.DB.Order("path").Find(&rows).Error; err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 3 {
+ t.Fatalf("media rows = %d, want 3", len(rows))
+ }
+ wantLibraries := map[string]string{
+ "cloud://openlist/电视剧/欧美剧/The Show/The.Show.S01E01.mkv": byDisplayDir["电视剧/欧美剧"].ID,
+ "cloud://openlist/电影/华语电影/Movie.2024.mkv": byDisplayDir["电影/华语电影"].ID,
+ "cloud://openlist/国漫/剑来/剑来.S01E01.mkv": byDisplayDir["动漫/国漫"].ID,
+ }
+ for _, row := range rows {
+ if row.LibraryID != wantLibraries[row.Path] {
+ t.Fatalf("%s library_id = %s, want %s", row.Path, row.LibraryID, wantLibraries[row.Path])
+ }
+ }
+
+ res, err = scanner.ScanLibrary(t.Context(), root.ID)
+ if err != nil {
+ t.Fatalf("rescan root cloud: %v", err)
+ }
+ if res.Added != 0 || res.Updated != 0 || res.Skipped != 3 {
+ t.Fatalf("rescan should skip unchanged auto-category rows, got %#v", res)
+ }
+ libs, err = repos.Library.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ autoCount := 0
+ for _, lib := range libs {
+ if CloudLibraryAutoCategory(lib) {
+ autoCount++
+ }
+ }
+ if autoCount != 3 {
+ t.Fatalf("auto category library count after rescan = %d, want 3", autoCount)
+ }
+
+ empty = true
+ res, err = scanner.ScanLibrary(t.Context(), root.ID)
+ if err != nil {
+ t.Fatalf("empty rescan root cloud: %v", err)
+ }
+ if res.Removed != 3 {
+ t.Fatalf("removed = %d, want 3", res.Removed)
+ }
+ if got := countMedia(t, repos); got != 0 {
+ t.Fatalf("media count after auto-category prune = %d, want 0", got)
+ }
+}
+
func TestScanCloudLibraryListsChildDirectoriesConcurrently(t *testing.T) {
var active int32
var maxActive int32
var releaseOnce sync.Once
release := make(chan struct{})
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/sort" {
- t.Errorf("unexpected path %s", r.URL.Path)
- return
- }
- w.Header().Set("Content-Type", "application/json")
- switch r.URL.Query().Get("pdir_fid") {
- case "0":
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"d1","file_name":"A","dir":true,"size":0},
- {"fid":"d2","file_name":"B","dir":true,"size":0}
- ]}}`))
- case "d1", "d2":
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
+ switch path {
+ case "/":
+ return []openListTestEntry{
+ {Name: "A", IsDir: true},
+ {Name: "B", IsDir: true},
+ }, 2
+ case "/A", "/B":
cur := atomic.AddInt32(&active, 1)
defer atomic.AddInt32(&active, -1)
for {
@@ -157,18 +336,17 @@ func TestScanCloudLibraryListsChildDirectoriesConcurrently(t *testing.T) {
}
select {
case <-release:
- case <-r.Context().Done():
- return
case <-time.After(1500 * time.Millisecond):
t.Errorf("child directory requests were not concurrent")
- return
+ return nil, 0
}
- id := r.URL.Query().Get("pdir_fid")
- _, _ = fmt.Fprintf(w, `{"status":200,"code":0,"data":{"list":[{"fid":"f-%s","file_name":"Movie.%s.mkv","dir":false,"size":123}]}}`, id, id)
+ id := strings.TrimPrefix(path, "/")
+ return []openListTestEntry{{Name: fmt.Sprintf("Movie.%s.mkv", id), Size: 123}}, 1
default:
- t.Errorf("unexpected pdir_fid %q", r.URL.Query().Get("pdir_fid"))
+ t.Errorf("unexpected openlist path %q", path)
+ return nil, 0
}
- }))
+ })
defer upstream.Close()
db, err := gorm.Open(sqlite.Open("file:cloud_scan_concurrent?mode=memory&cache=shared"), &gorm.Config{})
@@ -182,15 +360,15 @@ func TestScanCloudLibraryListsChildDirectoriesConcurrently(t *testing.T) {
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "movie", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
@@ -218,8 +396,8 @@ func TestCloudLibraryPathParsing(t *testing.T) {
if !ok || typ != "cloud115" || dir != "abc 123" {
t.Fatalf("parse path got typ=%q dir=%q ok=%v", typ, dir, ok)
}
- typ, dir, ok = parseCloudLibraryPath("cloud://quark?dir=0")
- if !ok || typ != "quark" || dir != "" {
+ typ, dir, ok = parseCloudLibraryPath("cloud://openlist/Movies?dir=%2FMovies")
+ if !ok || typ != "openlist" || dir != "Movies" {
t.Fatalf("parse query got typ=%q dir=%q ok=%v", typ, dir, ok)
}
if ref := cloudEntryRef("cloud115", "fid", "pick"); ref != "pick" {
@@ -335,13 +513,7 @@ func TestScan115CloudLibraryKeepsDisplayHierarchyAndSeasonCounts(t *testing.T) {
}))
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
@@ -424,185 +596,6 @@ func TestCloudMetadataNeedsRefreshWhenPathHintConflicts(t *testing.T) {
}
}
-func TestScanCloudLibraryReadsRemoteSTRMTarget(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- switch r.Method {
- case "PROPFIND":
- if r.URL.Path != "/dav/Links" {
- t.Fatalf("unexpected propfind path %s", r.URL.Path)
- }
- w.Header().Set("Content-Type", "application/xml")
- w.WriteHeader(http.StatusMultiStatus)
- _, _ = w.Write([]byte(`
-
-
- /dav/Links/
-
-
-
- /dav/Links/Movie.strm
- Movie.strm32
-
-`))
- case http.MethodGet:
- if r.URL.Path != "/dav/Links/Movie.strm" {
- t.Fatalf("unexpected get path %s", r.URL.Path)
- }
- _, _ = w.Write([]byte("https://cdn.example.com/Movie.mkv\n"))
- default:
- t.Fatalf("unexpected method %s", r.Method)
- }
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- log := zap.NewNop()
- storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
- if _, err := storage.Save(t.Context(), StorageInput{
- Type: "openlist",
- Config: map[string]any{
- "url": upstream.URL,
- },
- }); err != nil {
- t.Fatal(err)
- }
- lib := model.Library{Name: "OpenList · Links", Path: "cloud://openlist/Links", Type: "movie", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatal(err)
- }
- scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
- scanner.SetStorageConfig(storage)
-
- res, err := scanner.ScanLibrary(t.Context(), lib.ID)
- if err != nil {
- t.Fatalf("scan cloud: %v", err)
- }
- if res.Added != 1 {
- t.Fatalf("scan result = %#v, want added=1", res)
- }
- var media model.Media
- if err := repos.DB.First(&media).Error; err != nil {
- t.Fatal(err)
- }
- if media.Path != "cloud://openlist/Links/Movie.strm" {
- t.Fatalf("path = %q", media.Path)
- }
- if media.STRMURL != "https://cdn.example.com/Movie.mkv" {
- t.Fatalf("strm target = %q", media.STRMURL)
- }
-}
-
-func TestScanCloudLibraryReadsRemoteNFOAndArtwork(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- switch r.Method {
- case "PROPFIND":
- w.Header().Set("Content-Type", "application/xml")
- w.WriteHeader(http.StatusMultiStatus)
- switch r.URL.Path {
- case "/dav/Anime/JianLai":
- _, _ = w.Write([]byte(`
-
- /dav/Anime/JianLai/
- /dav/Anime/JianLai/tvshow.nfotvshow.nfo64
- /dav/Anime/JianLai/poster.jpgposter.jpg1024
- /dav/Anime/JianLai/Season1/Season1
-`))
- case "/dav/Anime/JianLai/Season1":
- _, _ = w.Write([]byte(`
-
- /dav/Anime/JianLai/Season1/
- /dav/Anime/JianLai/Season1/JianLai.S01E01.mkvJianLai.S01E01.mkv2048
- /dav/Anime/JianLai/Season1/JianLai.S01E01.nfoJianLai.S01E01.nfo128
-`))
- default:
- t.Fatalf("unexpected propfind path %s", r.URL.Path)
- }
- case http.MethodGet:
- switch r.URL.Path {
- case "/dav/Anime/JianLai/tvshow.nfo":
- _, _ = w.Write([]byte(`剑来2024天地有剑气`))
- case "/dav/Anime/JianLai/Season1/JianLai.S01E01.nfo":
- _, _ = w.Write([]byte(`剑来第一集11`))
- case "/dav/Anime/JianLai/poster.jpg":
- w.Header().Set("Content-Type", "image/jpeg")
- _, _ = w.Write([]byte("cloud-poster-bytes"))
- default:
- t.Fatalf("unexpected get path %s", r.URL.Path)
- }
- default:
- t.Fatalf("unexpected method %s", r.Method)
- }
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- log := zap.NewNop()
- storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
- if _, err := storage.Save(t.Context(), StorageInput{
- Type: "openlist",
- Config: map[string]any{
- "url": upstream.URL,
- },
- }); err != nil {
- t.Fatal(err)
- }
- lib := model.Library{Name: "OpenList · 国漫 · 剑来", Path: "cloud://openlist/Anime/JianLai", Type: "anime", Enabled: true}
- if err := repos.Library.Create(t.Context(), &lib); err != nil {
- t.Fatal(err)
- }
- scanner := NewScannerService(&config.Config{}, log, repos, NewHub(log), nil, nil)
- scanner.SetStorageConfig(storage)
- imageProxy := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, log)
- scanner.SetImageProxy(imageProxy)
-
- res, err := scanner.ScanLibrary(t.Context(), lib.ID)
- if err != nil {
- t.Fatalf("scan cloud: %v", err)
- }
- if res.Added != 1 || res.LocalMetadata != 1 {
- t.Fatalf("scan result = %#v, want added=1 local_metadata=1", res)
- }
- var media model.Media
- if err := repos.DB.First(&media).Error; err != nil {
- t.Fatal(err)
- }
- // 单集名(episode 「第一集」)不得写入 OriginalName(整剧原名/分组键)。
- // tvshow.nfo 未提供 originaltitle, 故 OriginalName 应为空。
- if media.Title != "剑来" || media.OriginalName != "" || media.Year != 2024 {
- t.Fatalf("metadata not applied: %#v", media)
- }
- if media.SeasonNum != 1 || media.EpisodeNum != 1 {
- t.Fatalf("episode numbers = %d/%d", media.SeasonNum, media.EpisodeNum)
- }
- if media.PosterURL != "/api/cloud/play/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg" {
- t.Fatalf("poster url = %q", media.PosterURL)
- }
- rec := httptest.NewRecorder()
- if !imageProxy.ServeCloudCached(rec, httptest.NewRequest(http.MethodGet, media.PosterURL, nil), "openlist:/Anime/JianLai/poster.jpg") {
- t.Fatal("cloud poster should be cached locally during scan before media is exposed")
- }
- if got := rec.Body.String(); got != "cloud-poster-bytes" {
- t.Fatalf("cached poster body = %q", got)
- }
- if media.ScrapeStatus != "matched" {
- t.Fatalf("scrape status = %q", media.ScrapeStatus)
- }
-}
-
func TestScanOpenListCloudLibraryUsesAPIPaginationBeyondFirstPage(t *testing.T) {
const totalFiles = 125
requestedPages := map[int]bool{}
@@ -656,13 +649,7 @@ func TestScanOpenListCloudLibraryUsesAPIPaginationBeyondFirstPage(t *testing.T)
}))
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
@@ -699,54 +686,44 @@ func TestScanOpenListCloudLibraryUsesAPIPaginationBeyondFirstPage(t *testing.T)
func TestScanCloudLibraryQueuesMissingExistingTrackMetadataBeforeNewFiles(t *testing.T) {
const newFiles = maxCloudMediaProbeQueuePerScan + 5
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/sort" || r.URL.Query().Get("pdir_fid") != "0" {
- t.Fatalf("unexpected cloud list request %s?%s", r.URL.Path, r.URL.RawQuery)
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
+ if path != "/" {
+ t.Fatalf("unexpected openlist path %q", path)
}
- var b strings.Builder
- b.WriteString(`{"status":200,"code":0,"data":{"list":[`)
+ entries := make([]openListTestEntry, 0, newFiles+1)
for i := 0; i < newFiles; i++ {
- if i > 0 {
- b.WriteByte(',')
- }
- _, _ = fmt.Fprintf(&b, `{"fid":"new-%02d","file_name":"New.Movie.%02d.mkv","dir":false,"size":%d}`, i, i, 1000+i)
+ entries = append(entries, openListTestEntry{Name: fmt.Sprintf("New.Movie.%02d.mkv", i), Size: int64(1000 + i)})
}
- _, _ = fmt.Fprintf(&b, `,{"fid":"existing","file_name":"Existing.Show.S01E01.mkv","dir":false,"size":2048}]}}`)
- _, _ = w.Write([]byte(b.String()))
- }))
+ entries = append(entries, openListTestEntry{Name: "Existing.Show.S01E01.mkv", Size: 2048})
+ return entries, len(entries)
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "tv", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
- existingPath := "cloud://quark/Existing.Show.S01E01.mkv"
+ existingPath := "cloud://openlist/Existing.Show.S01E01.mkv"
if err := repos.DB.Create(&model.Media{
LibraryID: lib.ID,
Title: "Existing Show",
Path: existingPath,
SizeBytes: 2048,
Container: "mkv",
- STRMURL: "/api/cloud/play/quark?ref=existing",
+ STRMURL: "/api/cloud/play/openlist?ref=%2FExisting.Show.S01E01.mkv",
SeasonNum: 1,
EpisodeNum: 1,
}).Error; err != nil {
@@ -778,15 +755,23 @@ func TestScanCloudLibraryQueuesMissingExistingTrackMetadataBeforeNewFiles(t *tes
}
}
-func TestParseCloudImagePlaybackURL(t *testing.T) {
- typ, ref, ok := parseCloudImagePlaybackURL("http://nas.local/api/cloud/play/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg")
+func TestParseCloudArtworkURL(t *testing.T) {
+ typ, ref, ok := ParseCloudArtworkURL("http://nas.local/api/cloud/play/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg")
if !ok || typ != "openlist" || ref != "/Anime/JianLai/poster.jpg" {
t.Fatalf("parse cloud image url = typ=%q ref=%q ok=%v", typ, ref, ok)
}
- if _, _, ok := parseCloudImagePlaybackURL("/api/cloud/play/openlist?ref=%2FAnime%2FJianLai%2Fmovie.mkv"); ok {
+ typ, ref, ok = ParseCloudArtworkURL("/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg")
+ if !ok || typ != "openlist" || ref != "/Anime/JianLai/poster.jpg" {
+ t.Fatalf("parse cached cloud artwork url = typ=%q ref=%q ok=%v", typ, ref, ok)
+ }
+ typ, ref, ok = ParseCloudArtworkURL("/api/img/cloud/openlist?ref=%2FMovies%2FMovie.tbn")
+ if !ok || typ != "openlist" || ref != "/Movies/Movie.tbn" {
+ t.Fatalf("parse tbn cloud artwork url = typ=%q ref=%q ok=%v", typ, ref, ok)
+ }
+ if _, _, ok := ParseCloudArtworkURL("/api/cloud/play/openlist?ref=%2FAnime%2FJianLai%2Fmovie.mkv"); ok {
t.Fatal("video cloud url should not be treated as artwork")
}
- if _, _, ok := parseCloudImagePlaybackURL("https://image.tmdb.org/t/p/w500/poster.jpg"); ok {
+ if _, _, ok := ParseCloudArtworkURL("https://image.tmdb.org/t/p/w500/poster.jpg"); ok {
t.Fatal("remote HTTP poster should not be treated as cloud artwork")
}
}
diff --git a/internal/service/scanner_existing_media.go b/internal/service/scanner_existing_media.go
new file mode 100644
index 0000000..0419a01
--- /dev/null
+++ b/internal/service/scanner_existing_media.go
@@ -0,0 +1,112 @@
+package service
+
+import (
+ "context"
+ "path/filepath"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) existingCloudMediaSnapshot(ctx context.Context, libraryID string) (map[string]existingCloudMedia, error) {
+ return s.existingCloudMediaSnapshotForLibraries(ctx, []string{libraryID})
+}
+
+func (s *ScannerService) existingCloudMediaSnapshotForLibraries(ctx context.Context, libraryIDs []string) (map[string]existingCloudMedia, error) {
+ if len(libraryIDs) == 0 {
+ return map[string]existingCloudMedia{}, nil
+ }
+ var rows []model.Media
+ if err := s.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Select("library_id", "path", "title", "original_name", "episode_title", "size_bytes", "duration_sec", "width", "height", "video_codec", "audio_codec", "container", "poster_url", "backdrop_url", "strm_url", "overview", "year", "rating", "tm_db_id", "bangumi_id", "douban_id", "thetvdb_id", "season_num", "episode_num", "genres", "countries", "languages", "nsfw", "scrape_status").
+ Where("library_id IN ? AND path LIKE ?", libraryIDs, "cloud://%").
+ Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ snapshot := make(map[string]existingCloudMedia, len(rows))
+ for _, row := range rows {
+ if strings.TrimSpace(row.Path) == "" {
+ continue
+ }
+ snapshot[row.Path] = existingCloudMedia{
+ LibraryID: row.LibraryID,
+ Title: row.Title,
+ OriginalName: row.OriginalName,
+ EpisodeTitle: row.EpisodeTitle,
+ SizeBytes: row.SizeBytes,
+ DurationSec: row.DurationSec,
+ Width: row.Width,
+ Height: row.Height,
+ VideoCodec: row.VideoCodec,
+ AudioCodec: row.AudioCodec,
+ Container: row.Container,
+ PosterURL: row.PosterURL,
+ BackdropURL: row.BackdropURL,
+ STRMURL: row.STRMURL,
+ Overview: row.Overview,
+ Year: row.Year,
+ Rating: row.Rating,
+ TMDbID: row.TMDbID,
+ BangumiID: row.BangumiID,
+ DoubanID: row.DoubanID,
+ TheTVDBID: row.TheTVDBID,
+ SeasonNum: row.SeasonNum,
+ EpisodeNum: row.EpisodeNum,
+ Genres: row.Genres,
+ Countries: row.Countries,
+ Languages: row.Languages,
+ NSFW: row.NSFW,
+ ScrapeStatus: row.ScrapeStatus,
+ }
+ }
+ return snapshot, nil
+}
+
+func (s *ScannerService) existingLocalMediaSnapshot(ctx context.Context, libraryID string) (map[string]existingLocalMedia, error) {
+ var rows []model.Media
+ if err := s.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Select("path", "title", "original_name", "episode_title", "size_bytes", "duration_sec", "width", "height", "video_codec", "audio_codec", "container", "strm_url", "file_id", "poster_url", "backdrop_url", "overview", "year", "rating", "tm_db_id", "bangumi_id", "douban_id", "thetvdb_id", "season_num", "episode_num", "genres", "countries", "languages", "nsfw", "scrape_status").
+ Where("library_id = ? AND path NOT LIKE ?", libraryID, "cloud://%").
+ Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ snapshot := make(map[string]existingLocalMedia, len(rows))
+ for _, row := range rows {
+ if strings.TrimSpace(row.Path) == "" {
+ continue
+ }
+ snapshot[filepath.Clean(row.Path)] = existingLocalMedia{
+ Title: row.Title,
+ OriginalName: row.OriginalName,
+ EpisodeTitle: row.EpisodeTitle,
+ SizeBytes: row.SizeBytes,
+ DurationSec: row.DurationSec,
+ Width: row.Width,
+ Height: row.Height,
+ VideoCodec: row.VideoCodec,
+ AudioCodec: row.AudioCodec,
+ Container: row.Container,
+ STRMURL: row.STRMURL,
+ FileID: row.FileID,
+ PosterURL: row.PosterURL,
+ BackdropURL: row.BackdropURL,
+ Overview: row.Overview,
+ Year: row.Year,
+ Rating: row.Rating,
+ TMDbID: row.TMDbID,
+ BangumiID: row.BangumiID,
+ DoubanID: row.DoubanID,
+ TheTVDBID: row.TheTVDBID,
+ SeasonNum: row.SeasonNum,
+ EpisodeNum: row.EpisodeNum,
+ Genres: row.Genres,
+ Countries: row.Countries,
+ Languages: row.Languages,
+ NSFW: row.NSFW,
+ ScrapeStatus: row.ScrapeStatus,
+ }
+ }
+ return snapshot, nil
+}
diff --git a/internal/service/scanner_existing_media_test.go b/internal/service/scanner_existing_media_test.go
new file mode 100644
index 0000000..ed49328
--- /dev/null
+++ b/internal/service/scanner_existing_media_test.go
@@ -0,0 +1,115 @@
+package service
+
+import (
+ "path/filepath"
+ "testing"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+ "go.uber.org/zap"
+)
+
+func TestExistingCloudMediaSnapshotFiltersCloudRows(t *testing.T) {
+ db := newServiceTestDB(t, &model.Media{})
+ repos := repository.New(db)
+ scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
+
+ if err := db.Create(&[]model.Media{
+ {
+ LibraryID: "lib-1",
+ Path: "cloud://openlist/Movie.mkv",
+ SizeBytes: 2048,
+ DurationSec: 120,
+ Width: 1920,
+ Height: 1080,
+ VideoCodec: "h264",
+ AudioCodec: "aac",
+ Container: "mkv",
+ PosterURL: "/poster.jpg",
+ BackdropURL: "/backdrop.jpg",
+ STRMURL: "/api/cloud/play/openlist?ref=movie",
+ Year: 2026,
+ TMDbID: 123,
+ BangumiID: 456,
+ DoubanID: "douban-1",
+ TheTVDBID: "tvdb-1",
+ ScrapeStatus: "matched",
+ },
+ {LibraryID: "lib-1", Path: "/media/local.mkv", SizeBytes: 99},
+ {LibraryID: "lib-2", Path: "cloud://openlist/Other.mkv", SizeBytes: 88},
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ got, err := scanner.existingCloudMediaSnapshot(t.Context(), "lib-1")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(got) != 1 {
+ t.Fatalf("snapshot len = %d, want 1: %#v", len(got), got)
+ }
+ row := got["cloud://openlist/Movie.mkv"]
+ if row.SizeBytes != 2048 || row.DurationSec != 120 || row.Width != 1920 || row.Height != 1080 {
+ t.Fatalf("track fields not preserved: %#v", row)
+ }
+ if row.VideoCodec != "h264" || row.AudioCodec != "aac" || row.Container != "mkv" {
+ t.Fatalf("codec fields not preserved: %#v", row)
+ }
+ if row.PosterURL != "/poster.jpg" || row.BackdropURL != "/backdrop.jpg" || row.STRMURL == "" {
+ t.Fatalf("artwork/strm fields not preserved: %#v", row)
+ }
+ if row.Year != 2026 || row.TMDbID != 123 || row.BangumiID != 456 || row.DoubanID != "douban-1" || row.TheTVDBID != "tvdb-1" {
+ t.Fatalf("scraper ids not preserved: %#v", row)
+ }
+}
+
+func TestExistingLocalMediaSnapshotFiltersAndCleansLocalRows(t *testing.T) {
+ db := newServiceTestDB(t, &model.Media{})
+ repos := repository.New(db)
+ scanner := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
+
+ rawPath := filepath.Join("D:", "media", "Movies", "..", "Movie.mkv")
+ cleanPath := filepath.Clean(rawPath)
+ if err := db.Create(&[]model.Media{
+ {
+ LibraryID: "lib-1",
+ Path: rawPath,
+ SizeBytes: 4096,
+ DurationSec: 240,
+ Width: 3840,
+ Height: 2160,
+ VideoCodec: "hevc",
+ AudioCodec: "truehd",
+ Container: "mkv",
+ STRMURL: "https://cdn.example.com/movie.mkv",
+ FileID: "dev:inode",
+ ScrapeStatus: "matched",
+ },
+ {LibraryID: "lib-1", Path: "cloud://openlist/Movie.mkv", SizeBytes: 99},
+ {LibraryID: "lib-2", Path: filepath.Join("D:", "media", "Other.mkv"), SizeBytes: 88},
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ got, err := scanner.existingLocalMediaSnapshot(t.Context(), "lib-1")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(got) != 1 {
+ t.Fatalf("snapshot len = %d, want 1: %#v", len(got), got)
+ }
+ row, ok := got[cleanPath]
+ if !ok {
+ t.Fatalf("snapshot key %q not found in %#v", cleanPath, got)
+ }
+ if row.SizeBytes != 4096 || row.DurationSec != 240 || row.Width != 3840 || row.Height != 2160 {
+ t.Fatalf("track fields not preserved: %#v", row)
+ }
+ if row.VideoCodec != "hevc" || row.AudioCodec != "truehd" || row.Container != "mkv" {
+ t.Fatalf("codec fields not preserved: %#v", row)
+ }
+ if row.STRMURL == "" || row.FileID != "dev:inode" {
+ t.Fatalf("identity fields not preserved: %#v", row)
+ }
+}
diff --git a/internal/service/scanner_incremental_test.go b/internal/service/scanner_incremental_test.go
index 124ec05..3e1ea91 100644
--- a/internal/service/scanner_incremental_test.go
+++ b/internal/service/scanner_incremental_test.go
@@ -6,9 +6,7 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -17,13 +15,7 @@ import (
func newScannerTestEnv(t *testing.T) (*ScannerService, *repository.Container) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
sc := NewScannerService(&config.Config{}, zap.NewNop(), repos, NewHub(zap.NewNop()), nil, nil)
return sc, repos
@@ -126,6 +118,60 @@ func TestScanLibrarySkipsUnchangedExistingLocalMedia(t *testing.T) {
}
}
+func TestScanLibrarySkipsUnchangedLocalMetadata(t *testing.T) {
+ sc, repos := newScannerTestEnv(t)
+ root := t.TempDir()
+ lib := model.Library{Name: "Movies", Path: root, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ file := filepath.Join(root, "Local Metadata (2024).mkv")
+ if err := os.WriteFile(file, []byte("same-size"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ nfo := filepath.Join(root, "Local Metadata (2024).nfo")
+ writeTestMovieNFO(t, nfo, "Local Metadata", "2024", "12345")
+
+ first, err := sc.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("first scan: %v", err)
+ }
+ if first.Added != 1 || first.LocalMetadata != 1 {
+ t.Fatalf("first scan = %#v, want added=1 local_metadata=1", first)
+ }
+ second, err := sc.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("second scan: %v", err)
+ }
+ if second.Added != 0 || second.Updated != 0 || second.Skipped != 1 {
+ t.Fatalf("second scan = %#v, want unchanged local metadata skipped", second)
+ }
+
+ writeTestMovieNFO(t, nfo, "Local Metadata Updated", "2024", "12345")
+ third, err := sc.ScanLibrary(t.Context(), lib.ID)
+ if err != nil {
+ t.Fatalf("third scan: %v", err)
+ }
+ if third.Added != 0 || third.Updated != 1 || third.Skipped != 0 {
+ t.Fatalf("third scan = %#v, want local metadata update only", third)
+ }
+ var media model.Media
+ if err := repos.DB.First(&media, "path = ?", file).Error; err != nil {
+ t.Fatal(err)
+ }
+ if media.Title != "Local Metadata Updated" || media.ScrapeStatus != "matched" {
+ t.Fatalf("local metadata was not refreshed: title=%q status=%q", media.Title, media.ScrapeStatus)
+ }
+}
+
+func writeTestMovieNFO(t *testing.T, path, title, year, tmdbID string) {
+ t.Helper()
+ body := `` + title + `` + year + `` + tmdbID + ``
+ if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
+ t.Fatal(err)
+ }
+}
+
func TestScanLibrarySkipsUnchangedExistingLocalMediaWithMissingTrackMetadata(t *testing.T) {
sc, repos := newScannerTestEnv(t)
root := t.TempDir()
diff --git a/internal/service/scanner_local_ingest.go b/internal/service/scanner_local_ingest.go
new file mode 100644
index 0000000..6cd460e
--- /dev/null
+++ b/internal/service/scanner_local_ingest.go
@@ -0,0 +1,241 @@
+package service
+
+import (
+ "context"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// ingestFile upserts a single media file. seenInodes dedups hardlinks within a
+// single scan; pass a fresh map for one-off ingests. It mutates res counters.
+func (s *ScannerService) ingestFile(ctx context.Context, lib *model.Library, path string, size int64, seenInodes map[string]string, existingMedia map[string]existingLocalMedia, writeBatch *localMediaWriteBatch, res *ScanResult) {
+ res.Visited++
+ ext := strings.ToLower(filepath.Ext(path))
+ cleanPath := filepath.Clean(path)
+
+ fileID, skippedDuplicate := s.recordLocalFileIdentity(ctx, path, seenInodes, existingMedia, res)
+ if skippedDuplicate {
+ return
+ }
+
+ parsedSeason, parsedEpisode := ParseEpisode(path)
+ localMeta := s.readLocalScanMetadata(lib, path, parsedSeason, parsedEpisode)
+ isNewMedia, skipUnchanged := s.localMediaScanState(localMediaScanStateInput{
+ ctx: ctx,
+ path: path,
+ cleanPath: cleanPath,
+ ext: ext,
+ size: size,
+ localMeta: localMeta,
+ existingMedia: existingMedia,
+ })
+ if skipUnchanged {
+ res.Skipped++
+ return
+ }
+
+ media := s.buildLocalScanMedia(localScanMediaInput{
+ lib: lib,
+ path: path,
+ ext: ext,
+ fileID: fileID,
+ size: size,
+ parsedSeason: parsedSeason,
+ parsedEpisode: parsedEpisode,
+ localMeta: localMeta,
+ res: res,
+ })
+ s.writeLocalScanMedia(localScanWriteInput{
+ ctx: ctx,
+ path: path,
+ media: media,
+ isNewMedia: isNewMedia,
+ writeBatch: writeBatch,
+ after: s.localProbeAfter(path, ext),
+ res: res,
+ })
+}
+
+func (s *ScannerService) recordLocalFileIdentity(ctx context.Context, path string, seenInodes map[string]string, existingMedia map[string]existingLocalMedia, res *ScanResult) (string, bool) {
+ fileID, hasID := fileIdentity(path)
+ if !hasID {
+ return "", false
+ }
+ if first, ok := seenInodes[fileID]; ok && first != path {
+ res.Skipped++
+ s.log.Debug("scan skip hardlink duplicate",
+ zap.String("path", path), zap.String("primary", first))
+ return fileID, true
+ }
+ if existingMedia == nil {
+ if other, ok := s.duplicateByFileID(ctx, fileID, path); ok {
+ res.Skipped++
+ s.log.Debug("scan skip hardlink duplicate (existing)",
+ zap.String("path", path), zap.String("primary", other))
+ return fileID, true
+ }
+ }
+ seenInodes[fileID] = path
+ return fileID, false
+}
+
+func (s *ScannerService) readLocalScanMetadata(lib *model.Library, path string, parsedSeason, parsedEpisode int) *LocalMetadata {
+ localMeta, err := ReadLocalMetadata(path, lib.Path, librarySupportsSeasons(lib) || parsedSeason > 0 || parsedEpisode > 0)
+ if err != nil {
+ s.log.Warn("read local metadata failed", zap.String("path", path), zap.Error(err))
+ }
+ return localMeta
+}
+
+type localMediaScanStateInput struct {
+ ctx context.Context
+ path string
+ cleanPath string
+ ext string
+ size int64
+ localMeta *LocalMetadata
+ existingMedia map[string]existingLocalMedia
+}
+
+func (s *ScannerService) localMediaScanState(in localMediaScanStateInput) (bool, bool) {
+ if in.existingMedia == nil {
+ return !s.mediaPathExists(in.ctx, in.path), false
+ }
+ existing, exists := in.existingMedia[in.cleanPath]
+ isNewMedia := !exists
+ if exists && in.ext != ".strm" && existing.SizeBytes == in.size && !localMetadataNeedsRefresh(existing, in.localMeta) {
+ return isNewMedia, true
+ }
+ return isNewMedia, false
+}
+
+type localScanMediaInput struct {
+ lib *model.Library
+ path string
+ ext string
+ fileID string
+ size int64
+ parsedSeason int
+ parsedEpisode int
+ localMeta *LocalMetadata
+ res *ScanResult
+}
+
+func (s *ScannerService) buildLocalScanMedia(in localScanMediaInput) *model.Media {
+ title, year := CleanQuery(in.path)
+ if title == "" {
+ title = strings.TrimSuffix(filepath.Base(in.path), in.ext)
+ }
+
+ media := &model.Media{
+ LibraryID: in.lib.ID,
+ Title: title,
+ Year: year,
+ Path: in.path,
+ SizeBytes: in.size,
+ Container: strings.TrimPrefix(in.ext, "."),
+ FileID: in.fileID,
+ SeasonNum: in.parsedSeason,
+ EpisodeNum: in.parsedEpisode,
+ }
+ if in.ext == ".strm" {
+ media.Container = "strm"
+ if targetURL, err := readLocalSTRMTarget(in.path); err == nil && targetURL != "" {
+ media.STRMURL = targetURL
+ } else if err != nil {
+ s.log.Debug("read local strm failed", zap.String("path", in.path), zap.Error(err))
+ }
+ }
+ if in.localMeta != nil {
+ applyLocalMetadata(media, in.localMeta)
+ in.res.LocalMetadata++
+ }
+ return media
+}
+
+func (s *ScannerService) localProbeAfter(path, ext string) func() {
+ if ext == ".strm" || s.probe == nil {
+ return nil
+ }
+ return func() {
+ s.queueLocalMediaProbe(path)
+ }
+}
+
+type localScanWriteInput struct {
+ ctx context.Context
+ path string
+ media *model.Media
+ isNewMedia bool
+ writeBatch *localMediaWriteBatch
+ after func()
+ res *ScanResult
+}
+
+func (s *ScannerService) writeLocalScanMedia(in localScanWriteInput) {
+ if in.isNewMedia && in.writeBatch != nil {
+ in.writeBatch.AddWithAfter(in.path, in.media, in.after)
+ return
+ }
+ if err := s.repo.Media.Upsert(in.ctx, in.media); err != nil {
+ addScanError(in.res, in.path, err)
+ s.log.Warn("upsert media failed", zap.String("path", in.path), zap.Error(err))
+ return
+ }
+ if in.after != nil {
+ in.after()
+ }
+ if in.isNewMedia {
+ in.res.Added++
+ } else {
+ in.res.Updated++
+ }
+ s.publishLocalScanProgress(in.path, in.res)
+}
+
+func (s *ScannerService) publishLocalScanProgress(path string, res *ScanResult) {
+ s.hub.Publish("scan", map[string]any{
+ "library_id": res.LibraryID,
+ "path": path,
+ "visited": res.Visited,
+ "added": res.Added,
+ "updated": res.Updated,
+ "probed": res.Probed,
+ "local_meta": res.LocalMetadata,
+ })
+}
+
+// duplicateByFileID reports an existing media path that shares the given inode
+// identity but lives at a different path and still exists on disk.
+func (s *ScannerService) duplicateByFileID(ctx context.Context, fileID, path string) (string, bool) {
+ if fileID == "" {
+ return "", false
+ }
+ var rows []model.Media
+ if err := s.repo.DB.WithContext(ctx).
+ Where("file_id = ? AND path <> ?", fileID, path).
+ Limit(8).Find(&rows).Error; err != nil {
+ return "", false
+ }
+ for _, r := range rows {
+ if r.Path == "" {
+ continue
+ }
+ if _, err := os.Stat(r.Path); err == nil {
+ return r.Path, true
+ }
+ }
+ return "", false
+}
+
+func (s *ScannerService) mediaPathExists(ctx context.Context, path string) bool {
+ var count int64
+ err := s.repo.DB.WithContext(ctx).Unscoped().Model(&model.Media{}).
+ Where("path = ?", path).Count(&count).Error
+ return err == nil && count > 0
+}
diff --git a/internal/service/scanner_local_probe_queue.go b/internal/service/scanner_local_probe_queue.go
new file mode 100644
index 0000000..752f155
--- /dev/null
+++ b/internal/service/scanner_local_probe_queue.go
@@ -0,0 +1,123 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScannerService) localMediaProbeWorker() {
+ for task := range s.localMediaProbeQueue {
+ s.probeLocalMediaAsync(task)
+ }
+}
+
+func (s *ScannerService) queueLocalMediaProbe(path string) bool {
+ task, ok := s.newLocalMediaProbeTask(path)
+ if !ok {
+ return false
+ }
+ s.startLocalMediaProbeWorkers()
+ if !s.reserveLocalMediaProbe(task.path) {
+ return false
+ }
+ if s.enqueueLocalMediaProbe(task) {
+ return true
+ }
+ s.releaseLocalMediaProbe(task.path)
+ s.logLocalMediaProbeQueueFull(task)
+ return false
+}
+
+func (s *ScannerService) newLocalMediaProbeTask(path string) (localMediaProbeTask, bool) {
+ if s == nil || s.probe == nil {
+ return localMediaProbeTask{}, false
+ }
+ task := localMediaProbeTask{path: strings.TrimSpace(path)}
+ return task, task.path != ""
+}
+
+func (s *ScannerService) startLocalMediaProbeWorkers() {
+ s.localMediaProbeOnce.Do(func() {
+ workers := s.ffprobeWorkerCount()
+ for i := 0; i < workers; i++ {
+ go s.localMediaProbeWorker()
+ }
+ })
+}
+
+func (s *ScannerService) reserveLocalMediaProbe(path string) bool {
+ s.localMediaProbeMu.Lock()
+ defer s.localMediaProbeMu.Unlock()
+ if s.localMediaProbing == nil {
+ s.localMediaProbing = make(map[string]struct{})
+ }
+ if _, ok := s.localMediaProbing[path]; ok {
+ return false
+ }
+ s.localMediaProbing[path] = struct{}{}
+ return true
+}
+
+func (s *ScannerService) releaseLocalMediaProbe(path string) {
+ s.localMediaProbeMu.Lock()
+ delete(s.localMediaProbing, path)
+ s.localMediaProbeMu.Unlock()
+}
+
+func (s *ScannerService) enqueueLocalMediaProbe(task localMediaProbeTask) bool {
+ select {
+ case s.localMediaProbeQueue <- task:
+ return true
+ default:
+ return false
+ }
+}
+
+func (s *ScannerService) logLocalMediaProbeQueueFull(task localMediaProbeTask) {
+ if s != nil && s.log != nil {
+ s.log.Debug("local media probe queue full", zap.String("path", task.path))
+ }
+}
+
+func (s *ScannerService) probeLocalMediaAsync(task localMediaProbeTask) {
+ defer s.releaseLocalMediaProbe(task.path)
+ if s == nil || s.probe == nil || strings.TrimSpace(task.path) == "" {
+ return
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
+ defer cancel()
+ probe, err := s.probe.Probe(ctx, task.path)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("local media async probe failed", zap.String("path", task.path), zap.Error(err))
+ }
+ return
+ }
+ updates := probeResultUpdates(probe)
+ if len(updates) == 0 {
+ return
+ }
+ if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("path = ?", task.path).Updates(updates).Error; err != nil {
+ if s.log != nil {
+ s.log.Debug("update local media track metadata failed", zap.String("path", task.path), zap.Error(err))
+ }
+ return
+ }
+ if s.hub != nil {
+ s.hub.Publish("scan", map[string]any{
+ "path": task.path,
+ "track_probed": true,
+ "duration_sec": probe.DurationSec,
+ "video_codec": probe.VideoCodec,
+ "audio_codec": probe.AudioCodec,
+ "width": probe.Width,
+ "height": probe.Height,
+ "probe_message": "本地媒体轨道元数据已后台补齐",
+ })
+ }
+}
diff --git a/internal/service/scanner_local_probe_queue_test.go b/internal/service/scanner_local_probe_queue_test.go
new file mode 100644
index 0000000..53c620b
--- /dev/null
+++ b/internal/service/scanner_local_probe_queue_test.go
@@ -0,0 +1,64 @@
+package service
+
+import (
+ "testing"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+func newLocalProbeQueueTestScanner(capacity int) *ScannerService {
+ return &ScannerService{
+ log: zap.NewNop(),
+ probe: &FFprobeService{},
+ localMediaProbeQueue: make(chan localMediaProbeTask, capacity),
+ localMediaProbing: make(map[string]struct{}),
+ }
+}
+
+func TestNewLocalMediaProbeTaskTrimsAndRejectsInvalidInput(t *testing.T) {
+ scanner := newLocalProbeQueueTestScanner(1)
+ task, ok := scanner.newLocalMediaProbeTask(" C:/media/movie.mkv ")
+ if !ok || task.path != "C:/media/movie.mkv" {
+ t.Fatalf("task = %#v ok=%v, want trimmed valid task", task, ok)
+ }
+ if _, ok := scanner.newLocalMediaProbeTask(" \t "); ok {
+ t.Fatal("blank path should be rejected")
+ }
+ scanner.probe = nil
+ if _, ok := scanner.newLocalMediaProbeTask("C:/media/movie.mkv"); ok {
+ t.Fatal("scanner without probe should reject local probe task")
+ }
+}
+
+func TestReserveLocalMediaProbeRejectsDuplicateAndReleaseAllowsRetry(t *testing.T) {
+ scanner := newLocalProbeQueueTestScanner(1)
+ if !scanner.reserveLocalMediaProbe("C:/media/movie.mkv") {
+ t.Fatal("first reserve should succeed")
+ }
+ if scanner.reserveLocalMediaProbe("C:/media/movie.mkv") {
+ t.Fatal("duplicate reserve should fail")
+ }
+ scanner.releaseLocalMediaProbe("C:/media/movie.mkv")
+ if !scanner.reserveLocalMediaProbe("C:/media/movie.mkv") {
+ t.Fatal("reserve should succeed after release")
+ }
+}
+
+func TestEnqueueLocalMediaProbeReportsFullQueue(t *testing.T) {
+ scanner := newLocalProbeQueueTestScanner(1)
+ if !scanner.enqueueLocalMediaProbe(localMediaProbeTask{path: "C:/media/a.mkv"}) {
+ t.Fatal("first enqueue should fit buffer")
+ }
+ if scanner.enqueueLocalMediaProbe(localMediaProbeTask{path: "C:/media/b.mkv"}) {
+ t.Fatal("second enqueue should report full queue")
+ }
+ select {
+ case task := <-scanner.localMediaProbeQueue:
+ if task.path != "C:/media/a.mkv" {
+ t.Fatalf("queued task = %#v, want first path", task)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("expected queued task")
+ }
+}
diff --git a/internal/service/scanner_local_write_batch.go b/internal/service/scanner_local_write_batch.go
new file mode 100644
index 0000000..da97cb6
--- /dev/null
+++ b/internal/service/scanner_local_write_batch.go
@@ -0,0 +1,104 @@
+package service
+
+import (
+ "context"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type localMediaWriteBatch struct {
+ scanner *ScannerService
+ ctx context.Context
+ res *ScanResult
+ limit int
+ items []localMediaWriteItem
+}
+
+type localMediaWriteItem struct {
+ path string
+ media *model.Media
+ after func()
+}
+
+func newLocalMediaWriteBatch(scanner *ScannerService, ctx context.Context, res *ScanResult, limit int) *localMediaWriteBatch {
+ if limit <= 0 {
+ limit = 100
+ }
+ return &localMediaWriteBatch{scanner: scanner, ctx: ctx, res: res, limit: limit}
+}
+
+func (b *localMediaWriteBatch) Add(path string, media *model.Media) {
+ b.AddWithAfter(path, media, nil)
+}
+
+func (b *localMediaWriteBatch) AddWithAfter(path string, media *model.Media, after func()) {
+ if b == nil || b.scanner == nil || media == nil {
+ return
+ }
+ if media.ScrapeStatus == "" {
+ media.ScrapeStatus = "pending"
+ }
+ b.items = append(b.items, localMediaWriteItem{path: path, media: media, after: after})
+ if len(b.items) >= b.limit {
+ b.Flush()
+ }
+}
+
+func (b *localMediaWriteBatch) Flush() {
+ if b == nil || len(b.items) == 0 || b.scanner == nil || b.scanner.repo == nil || b.scanner.repo.DB == nil {
+ return
+ }
+ items := b.items
+ b.items = nil
+ media := make([]model.Media, 0, len(items))
+ for _, item := range items {
+ if item.media != nil {
+ media = append(media, *item.media)
+ }
+ }
+ if len(media) == 0 {
+ return
+ }
+ if err := b.scanner.repo.DB.WithContext(b.ctx).CreateInBatches(&media, b.limit).Error; err == nil {
+ b.res.Added += len(media)
+ for _, item := range items {
+ if item.after != nil {
+ item.after()
+ }
+ }
+ b.publish()
+ return
+ }
+ for _, item := range items {
+ if item.media == nil {
+ continue
+ }
+ if err := b.scanner.repo.Media.Upsert(b.ctx, item.media); err != nil {
+ addScanError(b.res, item.path, err)
+ b.scanner.log.Warn("upsert media failed", zap.String("path", item.path), zap.Error(err))
+ continue
+ }
+ b.res.Added++
+ if item.after != nil {
+ item.after()
+ }
+ }
+ b.publish()
+}
+
+func (b *localMediaWriteBatch) publish() {
+ if b == nil || b.scanner == nil || b.scanner.hub == nil || b.res == nil {
+ return
+ }
+ b.scanner.hub.Publish("scan", map[string]any{
+ "library_id": b.res.LibraryID,
+ "visited": b.res.Visited,
+ "added": b.res.Added,
+ "updated": b.res.Updated,
+ "probed": b.res.Probed,
+ "local_meta": b.res.LocalMetadata,
+ "batched": true,
+ })
+}
diff --git a/internal/service/scanner_metadata_refresh.go b/internal/service/scanner_metadata_refresh.go
new file mode 100644
index 0000000..1b1285c
--- /dev/null
+++ b/internal/service/scanner_metadata_refresh.go
@@ -0,0 +1,200 @@
+package service
+
+import "strings"
+
+func cloudMetadataNeedsRefresh(existing existingCloudMedia, localMeta *LocalMetadata) bool {
+ if localMeta == nil {
+ return false
+ }
+ if localMeta.PathHint && !localMeta.HasNFO && !localMeta.HasArtwork {
+ return cloudPathHintNeedsRefresh(existing, localMeta)
+ }
+ if localMetadataMarksMatched(localMeta) && strings.TrimSpace(existing.ScrapeStatus) != "matched" {
+ return true
+ }
+ if localMeta.Title != "" && strings.TrimSpace(existing.Title) != strings.TrimSpace(localMeta.Title) {
+ return true
+ }
+ if localMeta.OriginalName != "" && strings.TrimSpace(existing.OriginalName) != strings.TrimSpace(localMeta.OriginalName) {
+ return true
+ }
+ if localMeta.EpisodeTitle != "" && strings.TrimSpace(existing.EpisodeTitle) != strings.TrimSpace(localMeta.EpisodeTitle) {
+ return true
+ }
+ if localMeta.AdultCode != "" && !strings.EqualFold(strings.TrimSpace(existing.OriginalName), strings.TrimSpace(localMeta.AdultCode)) {
+ return true
+ }
+ if localMeta.Year > 0 && existing.Year != localMeta.Year {
+ return true
+ }
+ if localMeta.Overview != "" && strings.TrimSpace(existing.Overview) != strings.TrimSpace(localMeta.Overview) {
+ return true
+ }
+ if localMeta.Rating > 0 && existing.Rating != localMeta.Rating {
+ return true
+ }
+ if localMeta.TMDbID > 0 && existing.TMDbID != localMeta.TMDbID {
+ return true
+ }
+ if localMeta.BangumiID > 0 && existing.BangumiID != localMeta.BangumiID {
+ return true
+ }
+ if strings.TrimSpace(localMeta.DoubanID) != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(localMeta.DoubanID) {
+ return true
+ }
+ if strings.TrimSpace(localMeta.TheTVDBID) != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(localMeta.TheTVDBID) {
+ return true
+ }
+ if strings.TrimSpace(localMeta.PosterURL) != "" && strings.TrimSpace(existing.PosterURL) != strings.TrimSpace(localMeta.PosterURL) {
+ return true
+ }
+ if strings.TrimSpace(localMeta.BackdropURL) != "" && strings.TrimSpace(existing.BackdropURL) != strings.TrimSpace(localMeta.BackdropURL) {
+ return true
+ }
+ if (localMeta.SeasonNum > 0 || localMeta.EpisodeNum > 0) && existing.SeasonNum != localMeta.SeasonNum {
+ return true
+ }
+ if localMeta.EpisodeNum > 0 && existing.EpisodeNum != localMeta.EpisodeNum {
+ return true
+ }
+ if localMeta.Genres != "" && strings.TrimSpace(existing.Genres) != strings.TrimSpace(localMeta.Genres) {
+ return true
+ }
+ if localMeta.Countries != "" && strings.TrimSpace(existing.Countries) != strings.TrimSpace(localMeta.Countries) {
+ return true
+ }
+ if localMeta.Languages != "" && strings.TrimSpace(existing.Languages) != strings.TrimSpace(localMeta.Languages) {
+ return true
+ }
+ if localMeta.NSFW && !existing.NSFW {
+ return true
+ }
+ return false
+}
+
+func cloudPathHintNeedsRefresh(existing existingCloudMedia, localMeta *LocalMetadata) bool {
+ if localMeta.TMDbID > 0 && existing.TMDbID != localMeta.TMDbID {
+ return true
+ }
+ if localMeta.BangumiID > 0 && existing.BangumiID != localMeta.BangumiID {
+ return true
+ }
+ if strings.TrimSpace(localMeta.DoubanID) != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(localMeta.DoubanID) {
+ return true
+ }
+ return strings.TrimSpace(localMeta.TheTVDBID) != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(localMeta.TheTVDBID)
+}
+
+func cloudTrackMetadataMissing(existing existingCloudMedia) bool {
+ return existing.DurationSec <= 0 ||
+ existing.Width <= 0 ||
+ existing.Height <= 0 ||
+ strings.TrimSpace(existing.VideoCodec) == "" ||
+ strings.TrimSpace(existing.AudioCodec) == ""
+}
+
+func localMetadataNeedsRefresh(existing existingLocalMedia, local *LocalMetadata) bool {
+ if local == nil {
+ return false
+ }
+ if localMetadataMarksMatched(local) && strings.TrimSpace(existing.ScrapeStatus) != "matched" {
+ return true
+ }
+ if local.Title != "" && strings.TrimSpace(existing.Title) != strings.TrimSpace(local.Title) {
+ return true
+ }
+ if local.OriginalName != "" && strings.TrimSpace(existing.OriginalName) != strings.TrimSpace(local.OriginalName) {
+ return true
+ }
+ if local.EpisodeTitle != "" && strings.TrimSpace(existing.EpisodeTitle) != strings.TrimSpace(local.EpisodeTitle) {
+ return true
+ }
+ if local.AdultCode != "" && !strings.EqualFold(strings.TrimSpace(existing.OriginalName), strings.TrimSpace(local.AdultCode)) {
+ return true
+ }
+ if local.Year > 0 && existing.Year != local.Year {
+ return true
+ }
+ if local.Overview != "" && strings.TrimSpace(existing.Overview) != strings.TrimSpace(local.Overview) {
+ return true
+ }
+ if local.Rating > 0 && existing.Rating != local.Rating {
+ return true
+ }
+ if local.PosterURL != "" && strings.TrimSpace(existing.PosterURL) != strings.TrimSpace(local.PosterURL) {
+ return true
+ }
+ if local.BackdropURL != "" && strings.TrimSpace(existing.BackdropURL) != strings.TrimSpace(local.BackdropURL) {
+ return true
+ }
+ if local.TMDbID > 0 && existing.TMDbID != local.TMDbID {
+ return true
+ }
+ if local.BangumiID > 0 && existing.BangumiID != local.BangumiID {
+ return true
+ }
+ if local.DoubanID != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(local.DoubanID) {
+ return true
+ }
+ if local.TheTVDBID != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(local.TheTVDBID) {
+ return true
+ }
+ if (local.SeasonNum > 0 || local.EpisodeNum > 0) && existing.SeasonNum != local.SeasonNum {
+ return true
+ }
+ if local.EpisodeNum > 0 && existing.EpisodeNum != local.EpisodeNum {
+ return true
+ }
+ if local.Genres != "" && strings.TrimSpace(existing.Genres) != strings.TrimSpace(local.Genres) {
+ return true
+ }
+ if local.Countries != "" && strings.TrimSpace(existing.Countries) != strings.TrimSpace(local.Countries) {
+ return true
+ }
+ if local.Languages != "" && strings.TrimSpace(existing.Languages) != strings.TrimSpace(local.Languages) {
+ return true
+ }
+ return local.NSFW && !existing.NSFW
+}
+
+func cloudSeriesTitleFromMediaPath(mediaPath string) (string, int) {
+ displayPath := strings.TrimSpace(mediaPath)
+ if strings.HasPrefix(strings.ToLower(displayPath), "cloud://") {
+ rest := strings.TrimPrefix(displayPath, "cloud://")
+ if idx := strings.Index(rest, "/"); idx >= 0 {
+ displayPath = rest[idx+1:]
+ } else {
+ return "", 0
+ }
+ }
+ displayPath = strings.Trim(strings.ReplaceAll(displayPath, "\\", "/"), "/")
+ if displayPath == "" {
+ return "", 0
+ }
+ parts := strings.Split(displayPath, "/")
+ if len(parts) < 2 {
+ return "", 0
+ }
+ dirs := parts[:len(parts)-1]
+ if len(dirs) == 0 {
+ return "", 0
+ }
+ base := strings.TrimSpace(dirs[len(dirs)-1])
+ usedSeasonFolder := false
+ if _, ok := seasonFromDir(base); ok {
+ usedSeasonFolder = true
+ dirs = dirs[:len(dirs)-1]
+ if len(dirs) == 0 {
+ return "", 0
+ }
+ base = strings.TrimSpace(dirs[len(dirs)-1])
+ }
+ if base == "" || (!usedSeasonFolder && len(dirs) < 2) {
+ return "", 0
+ }
+ title, year := CleanQuery(base)
+ if title == "" {
+ title = base
+ }
+ return strings.TrimSpace(title), year
+}
diff --git a/internal/service/scanner_post_scan.go b/internal/service/scanner_post_scan.go
new file mode 100644
index 0000000..bf7d923
--- /dev/null
+++ b/internal/service/scanner_post_scan.go
@@ -0,0 +1,25 @@
+package service
+
+import (
+ "context"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+func (s *ScannerService) invalidateMediaCache(ctx context.Context) {
+ if s != nil && s.cache != nil {
+ s.cache.DeletePrefix(ctx, "media:")
+ s.cache.DeletePrefix(ctx, "stats:")
+ }
+}
+
+func (s *ScannerService) startAutoScrape(ctx context.Context, libraryID string) {
+ scrapeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Minute)
+ go func() {
+ defer cancel()
+ if _, err := s.scraper.EnrichLibraryDetailedWithOptions(scrapeCtx, libraryID, skipEpisodeArtworkOptions(false)); err != nil {
+ s.log.Warn("scraper enrich failed", zap.Error(err))
+ }
+ }()
+}
diff --git a/internal/service/scanner_probe_queue.go b/internal/service/scanner_probe_queue.go
new file mode 100644
index 0000000..c9e48e8
--- /dev/null
+++ b/internal/service/scanner_probe_queue.go
@@ -0,0 +1,98 @@
+package service
+
+import (
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+func (s *ScannerService) cloudMediaProbeWorker() {
+ for task := range s.cloudMediaProbeQueue {
+ s.probeCloudMediaAsync(task)
+ }
+}
+
+func (s *ScannerService) queueCloudMediaProbe(typ, ref, path string) bool {
+ task, ok := s.newCloudMediaProbeTask(typ, ref, path)
+ if !ok || !s.reserveCloudMediaProbe(task, time.Now()) {
+ return false
+ }
+ select {
+ case s.cloudMediaProbeQueue <- task:
+ return true
+ default:
+ s.deferCloudMediaProbe(task, cloudMediaProbeQueueFullBackoff)
+ s.logCloudMediaProbeQueueFull(task)
+ return false
+ }
+}
+
+func (s *ScannerService) newCloudMediaProbeTask(typ, ref, path string) (cloudMediaProbeTask, bool) {
+ if s == nil || s.storage == nil || s.probe == nil {
+ return cloudMediaProbeTask{}, false
+ }
+ task := cloudMediaProbeTask{
+ typ: strings.TrimSpace(typ),
+ ref: strings.TrimSpace(ref),
+ path: strings.TrimSpace(path),
+ }
+ return task, task.typ != "" && task.ref != "" && task.path != ""
+}
+
+func (s *ScannerService) reserveCloudMediaProbe(task cloudMediaProbeTask, now time.Time) bool {
+ s.cloudMediaProbeMu.Lock()
+ defer s.cloudMediaProbeMu.Unlock()
+ if until, ok := s.cloudMediaProbeBackoff[task.path]; ok {
+ if now.Before(until) {
+ return false
+ }
+ delete(s.cloudMediaProbeBackoff, task.path)
+ }
+ if _, ok := s.cloudMediaProbing[task.path]; ok {
+ return false
+ }
+ s.cloudMediaProbing[task.path] = struct{}{}
+ return true
+}
+
+func (s *ScannerService) deferCloudMediaProbe(task cloudMediaProbeTask, backoff time.Duration) {
+ s.cloudMediaProbeMu.Lock()
+ defer s.cloudMediaProbeMu.Unlock()
+ delete(s.cloudMediaProbing, task.path)
+ if s.cloudMediaProbeBackoff == nil {
+ s.cloudMediaProbeBackoff = make(map[string]time.Time)
+ }
+ s.cloudMediaProbeBackoff[task.path] = time.Now().Add(backoff)
+}
+
+func (s *ScannerService) logCloudMediaProbeQueueFull(task cloudMediaProbeTask) {
+ if s == nil || s.log == nil {
+ return
+ }
+ now := time.Now()
+ s.cloudMediaProbeWarnMu.Lock()
+ shouldWarn := now.Sub(s.cloudMediaProbeLastWarn) >= time.Minute
+ if shouldWarn {
+ s.cloudMediaProbeLastWarn = now
+ }
+ s.cloudMediaProbeWarnMu.Unlock()
+ if shouldWarn {
+ s.log.Warn("cloud media probe queue full; deferring remaining probes (logged at most once per minute)",
+ zap.String("provider", task.typ), zap.String("path", task.path))
+ return
+ }
+ s.log.Debug("cloud media probe queue full", zap.String("provider", task.typ), zap.String("path", task.path))
+}
+
+func (s *ScannerService) queueCloudMediaProbeWithBudget(typ, ref, path string, budget *int) bool {
+ if budget != nil {
+ if *budget <= 0 {
+ return false
+ }
+ // Budget is consumed per attempt, not only per successful enqueue, so a
+ // full probe queue cannot generate unbounded repeated attempts/logging.
+ *budget--
+ }
+ return s.queueCloudMediaProbe(typ, ref, path)
+}
diff --git a/internal/service/scanner_probe_queue_test.go b/internal/service/scanner_probe_queue_test.go
new file mode 100644
index 0000000..467f62f
--- /dev/null
+++ b/internal/service/scanner_probe_queue_test.go
@@ -0,0 +1,69 @@
+package service
+
+import (
+ "testing"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+func newProbeQueueTestScanner(capacity int) *ScannerService {
+ return &ScannerService{
+ log: zap.NewNop(),
+ storage: &StorageConfigService{},
+ probe: &FFprobeService{},
+ cloudMediaProbeQueue: make(chan cloudMediaProbeTask, capacity),
+ cloudMediaProbing: make(map[string]struct{}),
+ cloudMediaProbeBackoff: make(map[string]time.Time),
+ }
+}
+
+func TestQueueCloudMediaProbeTrimsTaskAndRejectsDuplicate(t *testing.T) {
+ scanner := newProbeQueueTestScanner(1)
+ if !scanner.queueCloudMediaProbe(" openlist ", " /Movies/a.mkv ", " cloud://openlist/Movies/a.mkv ") {
+ t.Fatal("first cloud probe should enqueue")
+ }
+ if scanner.queueCloudMediaProbe("openlist", "/Movies/a.mkv", "cloud://openlist/Movies/a.mkv") {
+ t.Fatal("duplicate cloud probe should be rejected while in flight")
+ }
+
+ task := <-scanner.cloudMediaProbeQueue
+ if task.typ != "openlist" || task.ref != "/Movies/a.mkv" || task.path != "cloud://openlist/Movies/a.mkv" {
+ t.Fatalf("task was not normalized: %#v", task)
+ }
+}
+
+func TestQueueCloudMediaProbeFullQueueBacksOffAndReleases(t *testing.T) {
+ scanner := newProbeQueueTestScanner(0)
+ if scanner.queueCloudMediaProbe("openlist", "/Movies/a.mkv", "cloud://openlist/Movies/a.mkv") {
+ t.Fatal("unbuffered queue without receiver should reject enqueue")
+ }
+
+ scanner.cloudMediaProbeMu.Lock()
+ _, probing := scanner.cloudMediaProbing["cloud://openlist/Movies/a.mkv"]
+ until, backedOff := scanner.cloudMediaProbeBackoff["cloud://openlist/Movies/a.mkv"]
+ scanner.cloudMediaProbeMu.Unlock()
+ if probing {
+ t.Fatal("queue-full path should release in-flight marker")
+ }
+ if !backedOff || !until.After(time.Now()) {
+ t.Fatalf("queue-full path should receive future backoff, got %v", until)
+ }
+ if scanner.queueCloudMediaProbe("openlist", "/Movies/a.mkv", "cloud://openlist/Movies/a.mkv") {
+ t.Fatal("backed-off path should not be retried immediately")
+ }
+}
+
+func TestQueueCloudMediaProbeBudgetConsumesAttempts(t *testing.T) {
+ scanner := newProbeQueueTestScanner(0)
+ budget := 1
+ if scanner.queueCloudMediaProbeWithBudget("openlist", "/Movies/a.mkv", "cloud://openlist/Movies/a.mkv", &budget) {
+ t.Fatal("unbuffered queue without receiver should reject enqueue")
+ }
+ if budget != 0 {
+ t.Fatalf("budget = %d, want 0 after attempted enqueue", budget)
+ }
+ if scanner.queueCloudMediaProbeWithBudget("openlist", "/Movies/b.mkv", "cloud://openlist/Movies/b.mkv", &budget) {
+ t.Fatal("zero budget should prevent enqueue")
+ }
+}
diff --git a/internal/service/scanner_prune.go b/internal/service/scanner_prune.go
new file mode 100644
index 0000000..0e29e19
--- /dev/null
+++ b/internal/service/scanner_prune.go
@@ -0,0 +1,128 @@
+package service
+
+import (
+ "context"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// RemovePath deletes the media row for a path that has disappeared from disk
+// (incremental delete used by the watcher on Remove/Rename events).
+func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, error) {
+ if _, err := os.Stat(path); err == nil {
+ return 0, nil // still exists; nothing to remove
+ }
+ res := s.repo.DB.WithContext(ctx).
+ Where("path = ?", path).
+ Delete(&model.Media{})
+ if res.Error == nil && res.RowsAffected > 0 {
+ s.invalidateMediaCache(ctx)
+ }
+ return res.RowsAffected, res.Error
+}
+
+func (s *ScannerService) pruneMissingMedia(ctx context.Context, libraryID string, seen map[string]struct{}) (int64, error) {
+ // 只取 id/path,并把删除按批提交:此前整表载入完整 Media 结构体、
+ // 每行一条 DELETE,大库 prune 既费内存又长期占用写锁。
+ var rows []struct {
+ ID string
+ Path string
+ }
+ if err := s.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Select("id, path").
+ Where("library_id = ?", libraryID).
+ Find(&rows).Error; err != nil {
+ return 0, err
+ }
+ stale := make([]string, 0)
+ for _, row := range rows {
+ if row.Path == "" {
+ continue
+ }
+ if _, ok := seen[filepath.Clean(row.Path)]; ok {
+ continue
+ }
+ if _, err := os.Stat(row.Path); err == nil {
+ continue
+ } else if !os.IsNotExist(err) {
+ continue
+ }
+ stale = append(stale, row.ID)
+ }
+ return s.deleteMediaByIDs(ctx, stale, false)
+}
+
+// deleteMediaByIDs removes media rows in fixed-size batches so each write
+// transaction stays short and the global write gate is released frequently.
+func (s *ScannerService) deleteMediaByIDs(ctx context.Context, ids []string, hard bool) (int64, error) {
+ const batch = 500
+ var removed int64
+ for i := 0; i < len(ids); i += batch {
+ end := i + batch
+ if end > len(ids) {
+ end = len(ids)
+ }
+ q := s.repo.DB.WithContext(ctx)
+ if hard {
+ q = q.Unscoped()
+ }
+ res := q.Where("id IN ?", ids[i:end]).Delete(&model.Media{})
+ if res.Error != nil {
+ return removed, res.Error
+ }
+ removed += res.RowsAffected
+ }
+ return removed, nil
+}
+
+func (s *ScannerService) pruneMissingCloudMedia(ctx context.Context, libraryID string, seen map[string]struct{}) (int64, error) {
+ return s.pruneMissingCloudMediaForLibraries(ctx, []string{libraryID}, seen)
+}
+
+func (s *ScannerService) pruneMissingCloudMediaForLibraries(ctx context.Context, libraryIDs []string, seen map[string]struct{}) (int64, error) {
+ if len(libraryIDs) == 0 {
+ return 0, nil
+ }
+ var rows []struct {
+ ID string
+ Path string
+ }
+ if err := s.repo.DB.WithContext(ctx).
+ Model(&model.Media{}).
+ Select("id, path").
+ Where("library_id IN ? AND path LIKE ?", libraryIDs, "cloud://%").
+ Find(&rows).Error; err != nil {
+ return 0, err
+ }
+ stale := make([]string, 0)
+ for _, row := range rows {
+ if _, ok := seen[row.Path]; ok {
+ continue
+ }
+ stale = append(stale, row.ID)
+ }
+ return s.deleteMediaByIDs(ctx, stale, true)
+}
+
+func (s *ScannerService) autoScrapeEnabled(ctx context.Context) bool {
+ if s.repo == nil || s.repo.Setting == nil {
+ return false
+ }
+ value, err := s.repo.Setting.Get(ctx, "scrape.auto_on_scan")
+ if err != nil {
+ s.log.Warn("read scrape.auto_on_scan failed", zap.Error(err))
+ return false
+ }
+ switch strings.ToLower(strings.TrimSpace(value)) {
+ case "1", "true", "yes", "on", "enabled":
+ return true
+ default:
+ return false
+ }
+}
diff --git a/internal/service/scanner_scan.go b/internal/service/scanner_scan.go
new file mode 100644
index 0000000..83ec2f8
--- /dev/null
+++ b/internal/service/scanner_scan.go
@@ -0,0 +1,274 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "os"
+ "path/filepath"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// ScanLibrary walks the library root and persists discovered media files.
+func (s *ScannerService) ScanLibrary(ctx context.Context, libraryID string) (*ScanResult, error) {
+ return s.scanLibrary(ctx, libraryID, true)
+}
+
+// ScanLibraryWithoutAutoScrape walks a library without kicking off online
+// metadata enrichment. Cloud mounts can contain very large trees; keeping mount
+// scans import-only prevents scraper bursts from overwhelming small NAS boxes.
+func (s *ScannerService) ScanLibraryWithoutAutoScrape(ctx context.Context, libraryID string) (*ScanResult, error) {
+ return s.scanLibrary(ctx, libraryID, false)
+}
+
+func (s *ScannerService) TryBeginLocalScan(libraryID string) (func(), bool) {
+ if s == nil || strings.TrimSpace(libraryID) == "" {
+ return func() {}, true
+ }
+ s.localScanMu.Lock()
+ if s.localScans == nil {
+ s.localScans = make(map[string]struct{})
+ }
+ if _, ok := s.localScans[libraryID]; ok {
+ s.localScanMu.Unlock()
+ return nil, false
+ }
+ s.localScans[libraryID] = struct{}{}
+ s.localScanMu.Unlock()
+ return func() {
+ s.localScanMu.Lock()
+ delete(s.localScans, libraryID)
+ s.localScanMu.Unlock()
+ }, true
+}
+
+func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, autoScrape bool) (*ScanResult, error) {
+ lib, err := s.repo.Library.FindByID(ctx, libraryID)
+ if err != nil || lib == nil {
+ return nil, err
+ }
+ if mount, ok := ParseCloudLibraryMount(lib.Path); ok {
+ return s.scanMountedCloudLibrary(ctx, lib, mount, autoScrape)
+ }
+ if err := s.resolveLocalLibraryPath(ctx, lib); err != nil {
+ return &ScanResult{LibraryID: lib.ID}, err
+ }
+ res := &ScanResult{LibraryID: lib.ID}
+ seen := make(map[string]struct{})
+ seenInodes := make(map[string]string)
+ writeBatch := newLocalMediaWriteBatch(s, ctx, res, 100)
+ existingMedia, err := s.existingLocalMediaSnapshot(ctx, lib.ID)
+ if err != nil {
+ s.log.Warn("load existing local media snapshot failed", zap.String("library_id", lib.ID), zap.Error(err))
+ existingMedia = nil
+ } else {
+ for path, existing := range existingMedia {
+ if existing.FileID != "" {
+ seenInodes[existing.FileID] = path
+ }
+ }
+ }
+
+ walkFn := func(path string, info walkInfo) error {
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ default:
+ }
+ if info.isDir {
+ return nil
+ }
+ ext := strings.ToLower(filepath.Ext(path))
+ if _, ok := videoExtensions[ext]; !ok {
+ return nil
+ }
+ seen[filepath.Clean(path)] = struct{}{}
+ s.ingestFile(ctx, lib, path, info.size, seenInodes, existingMedia, writeBatch, res)
+ return nil
+ }
+
+ walkErr := walk(lib.Path, walkFn)
+ writeBatch.Flush()
+ if walkErr != nil {
+ addScanError(res, lib.Path, walkErr)
+ if res.Added+res.Updated > 0 {
+ s.invalidateMediaCache(ctx)
+ }
+ return res, walkErr
+ }
+ removed, err := s.pruneMissingMedia(ctx, lib.ID, seen)
+ if err != nil {
+ s.log.Warn("prune missing media failed", zap.String("library_id", lib.ID), zap.Error(err))
+ } else {
+ res.Removed = removed
+ }
+
+ s.hub.Publish("scan", map[string]any{
+ "library_id": lib.ID,
+ "finished": true,
+ "visited": res.Visited,
+ "added": res.Added,
+ "updated": res.Updated,
+ "probed": res.Probed,
+ "local_meta": res.LocalMetadata,
+ "removed": res.Removed,
+ "error_count": res.ErrorCount,
+ "errors": res.Errors,
+ })
+ s.notifyScanFinished(lib, res, nil, false)
+ s.invalidateMediaCache(ctx)
+ s.maybeGenerateSTRMAfterScan(lib.ID)
+
+ if scanHasImportChanges(res) && autoScrape && s.scraper != nil && s.scraper.AnyEnabled() && s.autoScrapeEnabled(ctx) {
+ s.startAutoScrape(ctx, lib.ID)
+ }
+ return res, nil
+}
+
+func (s *ScannerService) scanMountedCloudLibrary(ctx context.Context, lib *model.Library, mount CloudMountInfo, autoScrape bool) (*ScanResult, error) {
+ if IsDeprecatedNativeCloudProvider(mount.Provider) {
+ return &ScanResult{LibraryID: lib.ID, Skipped: 1}, nil
+ }
+ if CloudLibraryAutoCategory(*lib) {
+ res := &ScanResult{LibraryID: lib.ID, Skipped: 1}
+ s.log.Info("skip auto category cloud library scan",
+ zap.String("library_id", lib.ID),
+ zap.String("provider", mount.Provider))
+ s.hub.Publish("scan", map[string]any{
+ "library_id": lib.ID,
+ "finished": true,
+ "skipped": res.Skipped,
+ "cloud": true,
+ "auto_category": true,
+ })
+ return res, nil
+ }
+ if shadow := s.shadowedCloudLibrary(ctx, lib); shadow != nil {
+ res := &ScanResult{LibraryID: lib.ID, Skipped: 1}
+ s.log.Warn("skip shadowed cloud library scan",
+ zap.String("library_id", lib.ID),
+ zap.String("shadowed_by", shadow.Library.ID),
+ zap.String("provider", mount.Provider))
+ s.hub.Publish("scan", map[string]any{
+ "library_id": lib.ID,
+ "finished": true,
+ "skipped": res.Skipped,
+ "cloud": true,
+ "shadowed": true,
+ })
+ return res, nil
+ }
+ scanCtx, finish, err := s.beginCloudScan(ctx, lib, mount)
+ if err != nil {
+ if errors.Is(err, ErrCloudScanAlreadyRunning) {
+ return &ScanResult{LibraryID: lib.ID, Skipped: 1}, nil
+ }
+ return nil, err
+ }
+ release, err := s.acquireCloudScanSlot(scanCtx, lib.ID)
+ if err != nil {
+ res := &ScanResult{LibraryID: lib.ID}
+ if finish != nil {
+ finish(res, err)
+ }
+ return res, err
+ }
+ defer release()
+ res, err := s.scanCloudLibrary(scanCtx, lib, mount, autoScrape)
+ if finish != nil {
+ finish(res, err)
+ }
+ return res, err
+}
+
+func (s *ScannerService) notifyScanFinished(lib *model.Library, res *ScanResult, err error, cloud bool) {
+ if s == nil || s.notify == nil || lib == nil || res == nil {
+ return
+ }
+ if err != nil {
+ go func() {
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ s.notify.Broadcast(ctx, "MediaStationGo 扫描异常", fmt.Sprintf("媒体库:%s\n错误:%s", lib.Name, err.Error()), EventSystemAlert)
+ }()
+ return
+ }
+ if res.Added+res.Updated <= 0 {
+ return
+ }
+ source := "本地媒体库"
+ if cloud {
+ source = "网盘媒体库"
+ }
+ body := fmt.Sprintf("%s:%s\n新增:%d\n更新:%d\n跳过:%d\n移除:%d", source, lib.Name, res.Added, res.Updated, res.Skipped, res.Removed)
+ go func() {
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ s.notify.Broadcast(ctx, "MediaStationGo 入库完成", body, EventLibraryIngest)
+ }()
+}
+
+// IngestPath ingests a single file into the given library without walking the
+// whole tree. Used by the watcher for incremental, event-driven additions so
+// adding one new file no longer triggers a full library re-scan (减少硬盘损耗).
+// Non-video files and directories are ignored. Returns true if a media row was
+// added or updated.
+func (s *ScannerService) IngestPath(ctx context.Context, libraryID, path string) (bool, error) {
+ lib, err := s.repo.Library.FindByID(ctx, libraryID)
+ if err != nil || lib == nil {
+ return false, err
+ }
+ if err := s.resolveLocalLibraryPath(ctx, lib); err != nil {
+ return false, err
+ }
+ fi, err := os.Stat(path)
+ if err != nil || fi.IsDir() {
+ return false, err
+ }
+ ext := strings.ToLower(filepath.Ext(path))
+ if _, ok := videoExtensions[ext]; !ok {
+ return false, nil
+ }
+ res := &ScanResult{LibraryID: lib.ID}
+ s.ingestFile(ctx, lib, path, fi.Size(), make(map[string]string), nil, nil, res)
+ if res.Added+res.Updated > 0 {
+ s.invalidateMediaCache(ctx)
+ }
+ return res.Added+res.Updated > 0, nil
+}
+
+func (s *ScannerService) resolveLocalLibraryPath(ctx context.Context, lib *model.Library) error {
+ if lib == nil || strings.TrimSpace(lib.Path) == "" {
+ return nil
+ }
+ resolved, err := resolveAccessibleLibraryPath(lib.Path)
+ if err != nil {
+ return err
+ }
+ if sameLibraryPath(resolved, lib.Path) {
+ lib.Path = filepath.Clean(lib.Path)
+ return nil
+ }
+ if s.repo != nil && s.repo.DB != nil {
+ if updateErr := s.repo.DB.WithContext(ctx).Model(&model.Library{}).Where("id = ?", lib.ID).Update("path", resolved).Error; updateErr != nil && s.log != nil {
+ s.log.Warn("update mapped library path failed",
+ zap.String("library_id", lib.ID),
+ zap.String("from", lib.Path),
+ zap.String("to", resolved),
+ zap.Error(updateErr))
+ }
+ }
+ if s.log != nil {
+ s.log.Info("mapped library path for scan",
+ zap.String("library_id", lib.ID),
+ zap.String("from", lib.Path),
+ zap.String("to", resolved))
+ }
+ lib.Path = resolved
+ return nil
+}
diff --git a/internal/service/scanner_strm.go b/internal/service/scanner_strm.go
new file mode 100644
index 0000000..2641a00
--- /dev/null
+++ b/internal/service/scanner_strm.go
@@ -0,0 +1,104 @@
+package service
+
+import (
+ "context"
+ "net/url"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "go.uber.org/zap"
+)
+
+func (s *ScannerService) resolveCloudSTRMTarget(ctx context.Context, typ, ref string) (string, error) {
+ if s.storage == nil {
+ return "", nil
+ }
+ content, err := s.storage.CloudReadText(ctx, typ, ref, 64<<10)
+ if err != nil {
+ return "", err
+ }
+ for _, line := range strings.Split(content, "\n") {
+ candidate := strings.TrimSpace(strings.TrimPrefix(line, "\ufeff"))
+ if candidate == "" || strings.HasPrefix(candidate, "#") {
+ continue
+ }
+ u, err := url.Parse(candidate)
+ if err != nil {
+ continue
+ }
+ switch strings.ToLower(u.Scheme) {
+ case "http", "https", "webdav", "davs", "alist", "alists", "openlist", "openlists":
+ return candidate, nil
+ }
+ }
+ return "", nil
+}
+
+func readLocalSTRMTarget(path string) (string, error) {
+ data, err := os.ReadFile(path) // #nosec G304 -- path is a discovered .strm file under the configured library root.
+ if err != nil {
+ return "", err
+ }
+ for _, line := range strings.Split(string(data), "\n") {
+ candidate := strings.TrimSpace(strings.TrimPrefix(line, "\ufeff"))
+ if candidate == "" || strings.HasPrefix(candidate, "#") {
+ continue
+ }
+ if strings.HasPrefix(candidate, "/api/") || strings.HasPrefix(candidate, "/Videos/") || strings.HasPrefix(candidate, "/videos/") {
+ return candidate, nil
+ }
+ u, err := url.Parse(candidate)
+ if err != nil {
+ continue
+ }
+ switch strings.ToLower(u.Scheme) {
+ case "http", "https", "webdav", "davs", "alist", "alists", "openlist", "openlists":
+ return candidate, nil
+ }
+ }
+ return "", nil
+}
+
+func (s *ScannerService) maybeGenerateSTRMAfterScan(libraryID string) {
+ if s == nil || s.repo == nil || s.repo.Setting == nil {
+ return
+ }
+ value, err := s.repo.Setting.Get(context.Background(), "strm.auto_generate_enabled")
+ if err != nil || !parseBoolSetting(value, false) {
+ return
+ }
+ go func() {
+ ctx := context.Background()
+ strmSvc := NewSTRMService(s.log, s.repo, s.cfg)
+ opts := GenerateSTRMOptions{
+ LibraryID: libraryID,
+ Enabled: true,
+ IncludeLocal: true,
+ Overwrite: true,
+ }
+ if outDir, scope := s.autoSTRMOutputDir(ctx); outDir != "" {
+ opts.OutputDir = outDir
+ if scope == "all" {
+ if lib, err := s.repo.Library.FindByID(ctx, libraryID); err == nil && lib != nil {
+ opts.OutputDir = filepath.Join(outDir, strmLibraryOutputSubdir(*lib))
+ }
+ }
+ }
+ if _, err := strmSvc.GenerateForLibrary(ctx, opts); err != nil && s.log != nil {
+ s.log.Warn("auto generate strm failed", zap.String("library_id", libraryID), zap.Error(err))
+ }
+ }()
+}
+
+func (s *ScannerService) autoSTRMOutputDir(ctx context.Context) (string, string) {
+ if s == nil || s.repo == nil || s.repo.Setting == nil {
+ return "", ""
+ }
+ outDir, err := s.repo.Setting.Get(ctx, "strm.output_dir")
+ if err != nil {
+ return "", ""
+ }
+ scope, _ := s.repo.Setting.Get(ctx, "strm.output_scope")
+ return resolveMappedDestinationPath(strings.TrimSpace(outDir)), strings.ToLower(strings.TrimSpace(scope))
+}
diff --git a/internal/service/scheduler.go b/internal/service/scheduler.go
index 45c0c8a..890b59d 100644
--- a/internal/service/scheduler.go
+++ b/internal/service/scheduler.go
@@ -84,12 +84,7 @@ type scheduledJob struct {
type schedulerManualRunKey struct{}
const (
- cloudAutoSyncEnabledKey = "cloud.auto_sync_enabled"
- cloudSyncIntervalSecondsKey = "cloud.sync_interval_seconds"
- cloudLastAutoSyncDateKey = "cloud.last_auto_sync_date"
localLastPeriodicScanDateKey = "scan.last_periodic_date"
- cloudAutoSyncWindowStartHour = 23
- cloudAutoSyncWindowEndHour = 5
cloudAutoSyncCompletedDateForm = "2006-01-02"
)
@@ -369,189 +364,6 @@ func (s *SchedulerService) jobScanLibraries(ctx context.Context) error {
return nil
}
-// jobUploadLocalToCloud copies local media files into the configured external
-// storage backend. It is opt-in and never deletes the local source files.
-func (s *SchedulerService) jobUploadLocalToCloud(ctx context.Context) error {
- manual, _ := ctx.Value(schedulerManualRunKey{}).(bool)
- if s.storageCfg == nil || (!manual && !s.autoCloudUploadEnabled(ctx)) {
- return nil
- }
- input := s.cloudUploadInput(ctx)
- if strings.TrimSpace(input.Type) == "" || strings.TrimSpace(input.SourcePath) == "" {
- return nil
- }
- res, err := s.storageCfg.UploadLocal(ctx, input)
- if s.log != nil && res != nil {
- s.log.Info("cloud upload finished",
- zap.String("type", input.Type),
- zap.String("source", res.SourcePath),
- zap.String("dest", res.DestPath),
- zap.Int("uploaded", res.Uploaded),
- zap.Int("skipped", res.Skipped),
- zap.Int64("bytes", res.Bytes),
- zap.Int("errors", len(res.Errors)),
- )
- }
- return err
-}
-
-func (s *SchedulerService) cloudUploadInput(ctx context.Context) CloudUploadInput {
- get := func(key string) string {
- if s.repo == nil || s.repo.Setting == nil {
- return ""
- }
- v, _ := s.repo.Setting.Get(ctx, key)
- return strings.TrimSpace(v)
- }
- return CloudUploadInput{
- Type: get(CloudUploadProviderKey),
- SourcePath: get(CloudUploadSourceDirKey),
- DestPath: get(CloudUploadDestPathKey),
- Recursive: parseBoolSetting(get(CloudUploadRecursiveKey), true),
- IncludeSidecars: parseBoolSetting(get(CloudUploadSidecarsKey), true),
- Overwrite: parseBoolSetting(get(CloudUploadOverwriteKey), false),
- TransferMode: get(CloudUploadTransferModeKey),
- }
-}
-
-func (s *SchedulerService) autoCloudUploadEnabled(ctx context.Context) bool {
- if s.repo == nil || s.repo.Setting == nil {
- return false
- }
- v, err := s.repo.Setting.Get(ctx, CloudUploadAutoEnabledKey)
- if err != nil {
- return false
- }
- return parseBoolSetting(v, false)
-}
-
-func (s *SchedulerService) cloudUploadInterval(ctx context.Context) time.Duration {
- const fallback = time.Hour
- if s.repo == nil || s.repo.Setting == nil {
- return fallback
- }
- v, err := s.repo.Setting.Get(ctx, CloudUploadIntervalSecondsKey)
- if err != nil {
- return fallback
- }
- seconds, err := strconv.Atoi(strings.TrimSpace(v))
- if err != nil || seconds <= 0 {
- return fallback
- }
- if seconds < 300 {
- seconds = 300
- }
- return time.Duration(seconds) * time.Second
-}
-
-// jobSyncCloudLibraries keeps mounted cloud:// libraries refreshed without
-// enabling full disk scans. It imports remote cloud files as STRM-backed media
-// rows; the actual bytes stay on the provider and playback continues through
-// /api/cloud/play 302/proxy.
-func (s *SchedulerService) jobSyncCloudLibraries(ctx context.Context) error {
- manual, _ := ctx.Value(schedulerManualRunKey{}).(bool)
- if s.scanner == nil || (!manual && !s.autoCloudSyncDue(ctx, s.currentTime())) {
- return nil
- }
- libs, err := s.repo.Library.List(ctx)
- if err != nil {
- return err
- }
- libs = FilterScannableCloudLibraries(ctx, s.repo, libs)
- var firstErr error
- for _, l := range libs {
- if !l.Enabled {
- continue
- }
- if _, ok := ParseCloudLibraryMount(l.Path); !ok {
- continue
- }
- if _, err := s.scanner.ScanLibraryWithoutAutoScrape(ctx, l.ID); err != nil {
- s.log.Warn("cloud sync failed", zap.String("library", l.ID), zap.Error(err))
- if firstErr == nil {
- firstErr = err
- }
- }
- }
- if firstErr != nil {
- return firstErr
- }
- if !manual {
- _ = s.markCloudAutoSyncCompleted(ctx, s.currentTime())
- }
- return nil
-}
-
-func (s *SchedulerService) autoCloudSyncEnabled(ctx context.Context) bool {
- if s.repo == nil || s.repo.Setting == nil {
- return false
- }
- v, err := s.repo.Setting.Get(ctx, cloudAutoSyncEnabledKey)
- if err != nil {
- return false
- }
- return parseBoolSetting(v, false)
-}
-
-func (s *SchedulerService) autoCloudSyncDue(ctx context.Context, now time.Time) bool {
- if !s.autoCloudSyncEnabled(ctx) || !cloudAutoSyncInWindow(now) {
- return false
- }
- if s.repo == nil || s.repo.Setting == nil {
- return true
- }
- last, err := s.repo.Setting.Get(ctx, cloudLastAutoSyncDateKey)
- if err != nil {
- return true
- }
- return strings.TrimSpace(last) != cloudAutoSyncWindowDate(now)
-}
-
-func cloudAutoSyncInWindow(now time.Time) bool {
- hour := now.In(time.Local).Hour()
- if cloudAutoSyncWindowStartHour == cloudAutoSyncWindowEndHour {
- return true
- }
- if cloudAutoSyncWindowStartHour < cloudAutoSyncWindowEndHour {
- return hour >= cloudAutoSyncWindowStartHour && hour < cloudAutoSyncWindowEndHour
- }
- return hour >= cloudAutoSyncWindowStartHour || hour < cloudAutoSyncWindowEndHour
-}
-
-func cloudAutoSyncWindowDate(now time.Time) string {
- local := now.In(time.Local)
- if cloudAutoSyncWindowStartHour > cloudAutoSyncWindowEndHour && local.Hour() < cloudAutoSyncWindowEndHour {
- local = local.AddDate(0, 0, -1)
- }
- return local.Format(cloudAutoSyncCompletedDateForm)
-}
-
-func (s *SchedulerService) markCloudAutoSyncCompleted(ctx context.Context, now time.Time) error {
- if s.repo == nil || s.repo.Setting == nil {
- return nil
- }
- return s.repo.Setting.Set(ctx, cloudLastAutoSyncDateKey, cloudAutoSyncWindowDate(now))
-}
-
-func (s *SchedulerService) cloudSyncInterval(ctx context.Context) time.Duration {
- const fallback = 30 * time.Minute
- if s.repo == nil || s.repo.Setting == nil {
- return fallback
- }
- v, err := s.repo.Setting.Get(ctx, cloudSyncIntervalSecondsKey)
- if err != nil {
- return fallback
- }
- seconds, err := strconv.Atoi(strings.TrimSpace(v))
- if err != nil || seconds <= 0 {
- return fallback
- }
- if seconds < 300 {
- seconds = 300
- }
- return time.Duration(seconds) * time.Second
-}
-
func (s *SchedulerService) currentTime() time.Time {
if s != nil && s.now != nil {
return s.now()
diff --git a/internal/service/scheduler_cloud.go b/internal/service/scheduler_cloud.go
new file mode 100644
index 0000000..2bd643b
--- /dev/null
+++ b/internal/service/scheduler_cloud.go
@@ -0,0 +1,201 @@
+package service
+
+import (
+ "context"
+ "strconv"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+const (
+ cloudAutoSyncEnabledKey = "cloud.auto_sync_enabled"
+ cloudSyncIntervalSecondsKey = "cloud.sync_interval_seconds"
+ cloudLastAutoSyncDateKey = "cloud.last_auto_sync_date"
+ cloudAutoSyncWindowStartHour = 23
+ cloudAutoSyncWindowEndHour = 5
+)
+
+// jobUploadLocalToCloud copies local media files into the configured external
+// storage backend. It is opt-in and never deletes the local source files.
+func (s *SchedulerService) jobUploadLocalToCloud(ctx context.Context) error {
+ manual, _ := ctx.Value(schedulerManualRunKey{}).(bool)
+ if s.storageCfg == nil || (!manual && !s.autoCloudUploadEnabled(ctx)) {
+ return nil
+ }
+ input := s.cloudUploadInput(ctx)
+ if strings.TrimSpace(input.Type) == "" || strings.TrimSpace(input.SourcePath) == "" {
+ return nil
+ }
+ res, err := s.storageCfg.UploadLocal(ctx, input)
+ if s.log != nil && res != nil {
+ s.log.Info("cloud upload finished",
+ zap.String("type", input.Type),
+ zap.String("source", res.SourcePath),
+ zap.String("dest", res.DestPath),
+ zap.Int("uploaded", res.Uploaded),
+ zap.Int("skipped", res.Skipped),
+ zap.Int64("bytes", res.Bytes),
+ zap.Int("errors", len(res.Errors)),
+ )
+ }
+ return err
+}
+
+func (s *SchedulerService) cloudUploadInput(ctx context.Context) CloudUploadInput {
+ get := func(key string) string {
+ if s.repo == nil || s.repo.Setting == nil {
+ return ""
+ }
+ v, _ := s.repo.Setting.Get(ctx, key)
+ return strings.TrimSpace(v)
+ }
+ return CloudUploadInput{
+ Type: get(CloudUploadProviderKey),
+ SourcePath: get(CloudUploadSourceDirKey),
+ DestPath: get(CloudUploadDestPathKey),
+ Recursive: parseBoolSetting(get(CloudUploadRecursiveKey), true),
+ IncludeSidecars: parseBoolSetting(get(CloudUploadSidecarsKey), true),
+ Overwrite: parseBoolSetting(get(CloudUploadOverwriteKey), false),
+ TransferMode: get(CloudUploadTransferModeKey),
+ }
+}
+
+func (s *SchedulerService) autoCloudUploadEnabled(ctx context.Context) bool {
+ if s.repo == nil || s.repo.Setting == nil {
+ return false
+ }
+ v, err := s.repo.Setting.Get(ctx, CloudUploadAutoEnabledKey)
+ if err != nil {
+ return false
+ }
+ return parseBoolSetting(v, false)
+}
+
+func (s *SchedulerService) cloudUploadInterval(ctx context.Context) time.Duration {
+ const fallback = time.Hour
+ if s.repo == nil || s.repo.Setting == nil {
+ return fallback
+ }
+ v, err := s.repo.Setting.Get(ctx, CloudUploadIntervalSecondsKey)
+ if err != nil {
+ return fallback
+ }
+ seconds, err := strconv.Atoi(strings.TrimSpace(v))
+ if err != nil || seconds <= 0 {
+ return fallback
+ }
+ if seconds < 300 {
+ seconds = 300
+ }
+ return time.Duration(seconds) * time.Second
+}
+
+// jobSyncCloudLibraries keeps mounted cloud:// libraries refreshed without
+// enabling full disk scans. It imports remote cloud files as STRM-backed media
+// rows; the actual bytes stay on the provider and playback continues through
+// /api/cloud/play 302/proxy.
+func (s *SchedulerService) jobSyncCloudLibraries(ctx context.Context) error {
+ manual, _ := ctx.Value(schedulerManualRunKey{}).(bool)
+ if s.scanner == nil || (!manual && !s.autoCloudSyncDue(ctx, s.currentTime())) {
+ return nil
+ }
+ libs, err := s.repo.Library.List(ctx)
+ if err != nil {
+ return err
+ }
+ libs = FilterScannableCloudLibraries(ctx, s.repo, libs)
+ var firstErr error
+ for _, l := range libs {
+ if !l.Enabled {
+ continue
+ }
+ if _, ok := ParseCloudLibraryMount(l.Path); !ok {
+ continue
+ }
+ if _, err := s.scanner.ScanLibraryWithoutAutoScrape(ctx, l.ID); err != nil {
+ s.log.Warn("cloud sync failed", zap.String("library", l.ID), zap.Error(err))
+ if firstErr == nil {
+ firstErr = err
+ }
+ }
+ }
+ if firstErr != nil {
+ return firstErr
+ }
+ if !manual {
+ _ = s.markCloudAutoSyncCompleted(ctx, s.currentTime())
+ }
+ return nil
+}
+
+func (s *SchedulerService) autoCloudSyncEnabled(ctx context.Context) bool {
+ if s.repo == nil || s.repo.Setting == nil {
+ return false
+ }
+ v, err := s.repo.Setting.Get(ctx, cloudAutoSyncEnabledKey)
+ if err != nil {
+ return false
+ }
+ return parseBoolSetting(v, false)
+}
+
+func (s *SchedulerService) autoCloudSyncDue(ctx context.Context, now time.Time) bool {
+ if !s.autoCloudSyncEnabled(ctx) || !cloudAutoSyncInWindow(now) {
+ return false
+ }
+ if s.repo == nil || s.repo.Setting == nil {
+ return true
+ }
+ last, err := s.repo.Setting.Get(ctx, cloudLastAutoSyncDateKey)
+ if err != nil {
+ return true
+ }
+ return strings.TrimSpace(last) != cloudAutoSyncWindowDate(now)
+}
+
+func cloudAutoSyncInWindow(now time.Time) bool {
+ hour := now.In(time.Local).Hour()
+ if cloudAutoSyncWindowStartHour == cloudAutoSyncWindowEndHour {
+ return true
+ }
+ if cloudAutoSyncWindowStartHour < cloudAutoSyncWindowEndHour {
+ return hour >= cloudAutoSyncWindowStartHour && hour < cloudAutoSyncWindowEndHour
+ }
+ return hour >= cloudAutoSyncWindowStartHour || hour < cloudAutoSyncWindowEndHour
+}
+
+func cloudAutoSyncWindowDate(now time.Time) string {
+ local := now.In(time.Local)
+ if cloudAutoSyncWindowStartHour > cloudAutoSyncWindowEndHour && local.Hour() < cloudAutoSyncWindowEndHour {
+ local = local.AddDate(0, 0, -1)
+ }
+ return local.Format(cloudAutoSyncCompletedDateForm)
+}
+
+func (s *SchedulerService) markCloudAutoSyncCompleted(ctx context.Context, now time.Time) error {
+ if s.repo == nil || s.repo.Setting == nil {
+ return nil
+ }
+ return s.repo.Setting.Set(ctx, cloudLastAutoSyncDateKey, cloudAutoSyncWindowDate(now))
+}
+
+func (s *SchedulerService) cloudSyncInterval(ctx context.Context) time.Duration {
+ const fallback = 30 * time.Minute
+ if s.repo == nil || s.repo.Setting == nil {
+ return fallback
+ }
+ v, err := s.repo.Setting.Get(ctx, cloudSyncIntervalSecondsKey)
+ if err != nil {
+ return fallback
+ }
+ seconds, err := strconv.Atoi(strings.TrimSpace(v))
+ if err != nil || seconds <= 0 {
+ return fallback
+ }
+ if seconds < 300 {
+ seconds = 300
+ }
+ return time.Duration(seconds) * time.Second
+}
diff --git a/internal/service/scheduler_test.go b/internal/service/scheduler_test.go
index f798517..20ad44b 100644
--- a/internal/service/scheduler_test.go
+++ b/internal/service/scheduler_test.go
@@ -3,17 +3,13 @@ package service
import (
"context"
"errors"
- "net/http"
- "net/http/httptest"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -152,7 +148,18 @@ func TestSchedulerRunNowAsyncSurvivesCallerCancellation(t *testing.T) {
case <-time.After(time.Second):
t.Fatal("manual scheduled job did not finish after release")
}
- status := scheduler.Status()
+ var status []JobStatus
+ deadline := time.Now().Add(time.Second)
+ for {
+ status = scheduler.Status()
+ if len(status) == 1 && !status[0].Running {
+ break
+ }
+ if time.Now().After(deadline) {
+ break
+ }
+ time.Sleep(10 * time.Millisecond)
+ }
if len(status) != 1 || status[0].Running || status[0].LastErr != "" {
t.Fatalf("unexpected status after async run: %+v", status)
}
@@ -231,31 +238,23 @@ func TestSchedulerOrganizeSourceSyncsVisibilityWhenTargetAlreadyExists(t *testin
}
func TestSchedulerCloudSyncImportsMountedCloudLibrary(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/sort" || r.URL.Query().Get("pdir_fid") != "0" {
- t.Fatalf("unexpected cloud list request %s?%s", r.URL.Path, r.URL.RawQuery)
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
+ if path != "/" {
+ t.Fatalf("unexpected openlist path %q", path)
}
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"f1","file_name":"Cloud.Movie.2026.mkv","dir":false,"size":1024}
- ]}}`))
- }))
+ return []openListTestEntry{{Name: "Cloud.Movie.2026.mkv", Size: 1024}}, 1
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
@@ -264,7 +263,7 @@ func TestSchedulerCloudSyncImportsMountedCloudLibrary(t *testing.T) {
if err := repos.Library.Create(t.Context(), &local); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘 · 电影", Path: "cloud://quark/0", Type: "movie", Enabled: true}
+ lib := model.Library{Name: "OpenList · 电影", Path: "cloud://openlist", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
@@ -280,47 +279,39 @@ func TestSchedulerCloudSyncImportsMountedCloudLibrary(t *testing.T) {
t.Fatalf("cloud sync: %v", err)
}
var media model.Media
- if err := repos.DB.First(&media, "path = ?", "cloud://quark/Cloud.Movie.2026.mkv").Error; err != nil {
+ if err := repos.DB.First(&media, "path = ?", "cloud://openlist/Cloud.Movie.2026.mkv").Error; err != nil {
t.Fatalf("cloud media not imported: %v", err)
}
- if media.STRMURL != "/api/cloud/play/quark?ref=f1" {
+ if media.STRMURL != "/api/cloud/play/openlist?ref=%2FCloud.Movie.2026.mkv" {
t.Fatalf("strm url = %q", media.STRMURL)
}
}
func TestSchedulerCloudSyncRunsOnlyOnceInsideNightlyWindow(t *testing.T) {
var requests atomic.Int32
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
requests.Add(1)
- if r.URL.Path != "/file/sort" || r.URL.Query().Get("pdir_fid") != "0" {
- t.Fatalf("unexpected cloud list request %s?%s", r.URL.Path, r.URL.RawQuery)
+ if path != "/" {
+ t.Fatalf("unexpected openlist path %q", path)
}
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"f1","file_name":"Nightly.Cloud.Movie.2026.mkv","dir":false,"size":1024}
- ]}}`))
- }))
+ return []openListTestEntry{{Name: "Nightly.Cloud.Movie.2026.mkv", Size: 1024}}, 1
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "movie", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
@@ -378,34 +369,26 @@ func TestSchedulerCloudSyncRunsOnlyOnceInsideNightlyWindow(t *testing.T) {
func TestSchedulerRunNowCloudSyncBypassesNightlyWindow(t *testing.T) {
var requests atomic.Int32
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
requests.Add(1)
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"f1","file_name":"Manual.Cloud.Movie.2026.mkv","dir":false,"size":1024}
- ]}}`))
- }))
+ return []openListTestEntry{{Name: "Manual.Cloud.Movie.2026.mkv", Size: 1024}}, 1
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "movie", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
@@ -434,13 +417,7 @@ func TestSchedulerPeriodicLocalScanRunsAtMostOncePerDay(t *testing.T) {
libraryPath := filepath.Join(root, "library")
writeOrgFile(t, filepath.Join(libraryPath, "Daily.Show.S01E01.mkv"), "episode 1")
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "scan.periodic_enabled", "true"); err != nil {
t.Fatal(err)
@@ -487,13 +464,7 @@ func TestSchedulerManualLocalScanBypassesDailyPeriodicLimit(t *testing.T) {
libraryPath := filepath.Join(root, "library")
writeOrgFile(t, filepath.Join(libraryPath, "Manual.Show.S01E01.mkv"), "episode 1")
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{})
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), "scan.periodic_enabled", "true"); err != nil {
t.Fatal(err)
@@ -528,34 +499,26 @@ func TestSchedulerManualLocalScanBypassesDailyPeriodicLimit(t *testing.T) {
func TestSchedulerCloudSyncDisabledByDefault(t *testing.T) {
var requests atomic.Int32
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ upstream := newOpenListAPIServer(t, func(path string, page, perPage int) ([]openListTestEntry, int) {
requests.Add(1)
- _, _ = w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
- {"fid":"f1","file_name":"Cloud.Movie.2026.mkv","dir":false,"size":1024}
- ]}}`))
- }))
+ return []openListTestEntry{{Name: "Cloud.Movie.2026.mkv", Size: 1024}}, 1
+ })
defer upstream.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.Setting{}, &model.StorageConfig{})
repos := repository.New(db)
log := zap.NewNop()
storage := NewStorageConfigService(log, repos, NewCryptoService("", log))
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "openlist-token",
},
}); err != nil {
t.Fatal(err)
}
- lib := model.Library{Name: "夸克网盘", Path: "cloud://quark/0", Type: "movie", Enabled: true}
+ lib := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
diff --git a/internal/service/scraper.go b/internal/service/scraper.go
index 12d18ea..f26123a 100644
--- a/internal/service/scraper.go
+++ b/internal/service/scraper.go
@@ -8,7 +8,6 @@ package service
import (
"context"
"path/filepath"
- "regexp"
"strconv"
"strings"
"time"
@@ -33,6 +32,24 @@ type ScraperService struct {
adult *AdultProvider
hub *Hub
notify *NotifyChannelService
+ cache *RuntimeCacheService
+ images *ImageProxy
+}
+
+type ScrapeOptions struct {
+ RetryNoMatch bool
+ IncludeMatched bool
+ EpisodeArtwork *bool
+ DeferEpisodeDetails bool
+}
+
+func (o ScrapeOptions) episodeArtworkEnabled() bool {
+ return o.EpisodeArtwork == nil || *o.EpisodeArtwork
+}
+
+func skipEpisodeArtworkOptions(retryNoMatch bool) ScrapeOptions {
+ episodeArtwork := false
+ return ScrapeOptions{RetryNoMatch: retryNoMatch, EpisodeArtwork: &episodeArtwork}
}
// NewScraperService is the constructor.
@@ -67,148 +84,34 @@ func (s *ScraperService) SetNotifyChannels(notify *NotifyChannelService) {
}
}
-// yearPattern extracts a 4-digit year (1900-2099).
-var yearPattern = regexp.MustCompile(`(?:^|[^\d])(19\d{2}|20\d{2})(?:[^\d]|$)`)
+func (s *ScraperService) SetRuntimeCache(cache *RuntimeCacheService) *ScraperService {
+ if s != nil {
+ s.cache = cache
+ }
+ return s
+}
+
+func (s *ScraperService) SetImageProxy(images *ImageProxy) *ScraperService {
+ if s != nil {
+ s.images = images
+ }
+ return s
+}
var tmdbDetailsTimeout = 8 * time.Second
-// noiseTokens are stripped before search.
-var noiseTokens = []string{
- // 视频规格
- "1080p", "2160p", "4k", "720p", "480p", "uhd", "ds4k", "fhd",
- "bd", "bdrip", "brrip", "dvd", "dvdrip", "hdtv", "pdtv", "webdl",
- "hdrip", "bluray", "blu-ray", "webrip", "web-dl", "web",
- "x264", "x265", "h264", "h265", "hevc", "avc", "10bit", "8bit", "hi10p", "hi10",
- "hdr", "hdr10", "sdr", "dts", "ddp", "ddp5", "dd5", "dd2", "eac3", "truehd",
- "dovi", "atmos", "aac", "ac3", "flac",
- "remux", "extended", "uncut", "remastered", "repack", "proper", "internal",
- "limited", "imax", "directors-cut", "directors_cut",
- "hkfree", "yify", "rarbg", "ettv", "fgt", "tgx", "ctrlhd", "ntb", "flux",
-
- // 流媒体平台 / 字幕组 / 国家版本(动漫常见)
- "netflix", "nf", "amzn", "hulu", "disney", "max", "hbo",
- "linetv", "ourtv", "iqiyi", "youku", "bilibili", "qiyi", "krj",
- "crunchyroll", "funimation", "anidb", "horriblesubs", "subsplease",
- "erai-raws", "judas", "asw", "smcat", "leopard-raws", "ohys-raws", "colortv",
-
- // 中文字幕标记
- "zm", "zw", "ch", "chs", "cht", "cn", "tc", "sc",
- "中字", "繁字", "简中", "繁中", "国语", "粤语", "日语",
-
- // 季数前缀残留 — ParseEpisode 已抽取过
- "season", "264", "265",
-}
-
-var noiseTokenSet = func() map[string]struct{} {
- set := make(map[string]struct{}, len(noiseTokens)+1)
- for _, token := range noiseTokens {
- set[token] = struct{}{}
- }
- set["dl"] = struct{}{}
- return set
-}()
-
-var releaseBoundaryTokenSet = map[string]struct{}{
- "1080p": {}, "2160p": {}, "4k": {}, "720p": {}, "480p": {}, "uhd": {}, "fhd": {},
- "bd": {}, "bdrip": {}, "brrip": {}, "dvd": {}, "dvdrip": {}, "hdtv": {}, "pdtv": {},
- "webdl": {}, "hdrip": {}, "bluray": {}, "webrip": {}, "web": {}, "remux": {},
- "x264": {}, "x265": {}, "h264": {}, "h265": {}, "hevc": {}, "avc": {},
-}
-var strictSeasonFolderPatterns = []*regexp.Regexp{
- regexp.MustCompile(`(?i)^(?:s|season)\.?\s*(\d{1,2})$`),
- regexp.MustCompile(`^第\s*([0-9一二三四五六七八九十百零两]+)\s*季$`),
-}
-
-// bracketedTag matches "[anything]", "(anything)" or "{anything}" segments.
-var bracketedTag = regexp.MustCompile(`[\[\(\{][^\]\)\}]*[\]\)\}]`)
-var multiWordNoise = []*regexp.Regexp{
- regexp.MustCompile(`(?i)\bweb[\s._-]*dl\b`),
- regexp.MustCompile(`(?i)\bblu[\s._-]*ray\b`),
- regexp.MustCompile(`(?i)\bdirectors[\s._-]*cut\b`),
- regexp.MustCompile(`(?i)\berai[\s._-]*raws\b`),
- regexp.MustCompile(`(?i)\bohys[\s._-]*raws\b`),
-}
-
const (
defaultScrapeDelayMinMS = 250
defaultScrapeDelayMaxMS = 500
maxScrapeDelayMS = 5 * 60 * 1000
)
-// CleanQuery converts a filename like "Inception.2010.1080p.BluRay.x264.mkv"
-// into a TMDb-friendly title plus an optional year hint.
-func CleanQuery(raw string) (title string, year int) {
- name := strings.TrimSuffix(filepath.Base(raw), filepath.Ext(raw))
- lower := strings.ToLower(name)
-
- if m := yearPattern.FindStringSubmatch(lower); len(m) >= 2 {
- if v, err := strconv.Atoi(m[1]); err == nil {
- year = v
- lower = strings.ReplaceAll(lower, m[1], " ")
- }
- }
-
- lower = bracketedTag.ReplaceAllString(lower, " ")
-
- lower = patSEnE.ReplaceAllString(lower, " ")
- lower = patNxE.ReplaceAllString(lower, " ")
- lower = patEP.ReplaceAllString(lower, " ")
- lower = patCN.ReplaceAllString(lower, " ")
- // 去掉中文季/部标记(如「第二季」「第2部」),避免残留在标题里既污染
- // 搜索查询又导致整理后的目录名重复季信息。
- lower = patSeasonOnly.ReplaceAllString(lower, " ")
- lower = patCNSeason.ReplaceAllString(lower, " ")
-
- for _, pat := range multiWordNoise {
- lower = pat.ReplaceAllString(lower, " ")
- }
- for _, sep := range []string{".", "_", "-", "[", "]", "(", ")", "×"} {
- lower = strings.ReplaceAll(lower, sep, " ")
- }
- // 拆分后丢掉过短(≤1)且全为 ASCII 数字 / 字母的"碎片",避免
- // 「2」「0」「v」之类残留干扰 TMDb 搜索。中文字符不算碎片。
- out := make([]string, 0, 8)
- seenReleaseBoundary := false
- for _, w := range strings.Fields(lower) {
- if _, ok := noiseTokenSet[w]; ok {
- if _, boundary := releaseBoundaryTokenSet[w]; boundary {
- seenReleaseBoundary = true
- }
- continue
- }
- if seenReleaseBoundary && isASCIIWord(w) {
- continue
- }
- if len(w) <= 1 {
- r := []rune(w)
- if len(r) == 1 && r[0] < 128 {
- continue
- }
- }
- out = append(out, w)
- }
- title = strings.TrimSpace(strings.Join(out, " "))
- return title, year
-}
-
-func isASCIIWord(s string) bool {
- if s == "" {
- return false
- }
- for _, r := range s {
- if r >= 128 {
- return false
- }
- if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') {
- continue
- }
- return false
- }
- return true
-}
-
// EnrichOne runs the provider chain for a single media row.
func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
+ return s.EnrichOneWithOptions(ctx, m, ScrapeOptions{})
+}
+
+func (s *ScraperService) EnrichOneWithOptions(ctx context.Context, m *model.Media, options ScrapeOptions) error {
lib, err := s.repo.Library.FindByID(ctx, m.LibraryID)
if err != nil {
return err
@@ -220,12 +123,14 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
if !cloudMedia {
if found, err := ReadLocalMetadata(m.Path, lib.Path, seriesLike); err == nil && found != nil {
local = found
- applyLocalMetadata(m, local)
} else if err != nil {
s.log.Warn("read local metadata before scrape failed", zap.String("media_id", m.ID), zap.Error(err))
}
- } else if hinted, _ := pathHintMetadata(m.Path, seriesLike); hinted != nil {
- local = mergeCloudMetadata(local, hinted)
+ }
+ if hinted, _ := pathHintMetadata(m.Path, seriesLike); hinted != nil {
+ local = mergeScrapePathHintMetadata(local, hinted)
+ }
+ if local != nil {
applyLocalMetadata(m, local)
}
@@ -235,7 +140,7 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
if code := firstText(localAdultCode(local), AdultCodeFromMediaPath(m.Path), normalizeAdultCode(m.OriginalName), normalizeAdultCode(m.Title)); code != "" {
if adultMatch, err := s.adult.Search(ctx, code); err == nil && adultMatch != nil {
mergeLocalMetadataIntoMatch(adultMatch, local)
- return s.applyProviderMatch(ctx, m, lib, adultMatch)
+ return s.applyProviderMatchWithOptions(ctx, m, lib, adultMatch, options)
} else if err != nil {
s.log.Debug("adult metadata search failed", zap.String("media_id", m.ID), zap.String("code", code), zap.Error(err))
}
@@ -245,7 +150,7 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
if match := s.matchFromMediaExternalIDs(ctx, m, lib); match != nil {
s.applyFanartArtwork(ctx, match)
mergeLocalMetadataIntoMatch(match, local)
- return s.applyProviderMatch(ctx, m, lib, match)
+ return s.applyProviderMatchWithOptions(ctx, m, lib, match, options)
}
candidates := scrapeQueryCandidates(m, lib)
@@ -253,7 +158,7 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
match := (*Match)(nil)
for _, candidate := range candidates {
query = candidate
- candidateMatch := s.lookup(ctx, lib, candidate, year)
+ candidateMatch := s.lookup(ctx, lib, m, candidate, year)
if candidateMatch == nil {
continue
}
@@ -276,11 +181,12 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
}
}
if match == nil {
- if local != nil {
+ if local != nil && !local.PathHint {
return s.applyLocalMetadataMatch(ctx, m, local)
}
_ = s.repo.DB.Model(&model.Media{}).Where("id = ?", m.ID).
Update("scrape_status", "no_match").Error
+ s.invalidateMediaCache(ctx)
s.log.Info("metadata scrape no match",
zap.String("media_id", m.ID),
zap.String("query", query),
@@ -290,7 +196,7 @@ func (s *ScraperService) EnrichOne(ctx context.Context, m *model.Media) error {
s.applyFanartArtwork(ctx, match)
mergeLocalMetadataIntoMatch(match, local)
- return s.applyProviderMatch(ctx, m, lib, match)
+ return s.applyProviderMatchWithOptions(ctx, m, lib, match, options)
}
func (s *ScraperService) matchFromMediaExternalIDs(ctx context.Context, m *model.Media, lib *model.Library) *Match {
@@ -400,6 +306,10 @@ func mergeLocalMetadataIntoMatch(match *Match, local *LocalMetadata) {
if match == nil || local == nil {
return
}
+ if local.PathHint {
+ mergePathHintIDsIntoMatch(match, local)
+ return
+ }
if local.Title != "" {
match.Title = local.Title
}
@@ -451,12 +361,41 @@ func mergeLocalMetadataIntoMatch(match *Match, local *LocalMetadata) {
}
}
+func mergePathHintIDsIntoMatch(match *Match, local *LocalMetadata) {
+ if match == nil || local == nil {
+ return
+ }
+ if local.TMDbID > 0 {
+ match.TMDbID = local.TMDbID
+ }
+ if local.BangumiID > 0 {
+ match.BangumiID = local.BangumiID
+ }
+ if local.DoubanID != "" {
+ match.DoubanID = local.DoubanID
+ }
+ if local.TheTVDBID != "" {
+ match.TheTVDBID = local.TheTVDBID
+ }
+ if match.Year <= 0 && local.Year > 0 {
+ match.Year = local.Year
+ }
+}
+
func (s *ScraperService) applyProviderMatch(ctx context.Context, m *model.Media, lib *model.Library, match *Match) error {
+ return s.applyProviderMatchWithOptions(ctx, m, lib, match, ScrapeOptions{})
+}
+
+func (s *ScraperService) applyProviderMatchWithOptions(ctx context.Context, m *model.Media, lib *model.Library, match *Match, options ScrapeOptions) error {
+ posterCandidate := match.PosterURL
+ backdropCandidate := match.BackdropURL
+ posterURL, removePoster := s.prepareScrapedArtworkURL(ctx, m.ID, "poster_url", m.PosterURL, posterCandidate)
+ backdropURL, removeBackdrop := s.prepareScrapedArtworkURL(ctx, m.ID, "backdrop_url", m.BackdropURL, backdropCandidate)
updates := map[string]any{
"title": match.Title,
"overview": match.Overview,
- "poster_url": match.PosterURL,
- "backdrop_url": match.BackdropURL,
+ "poster_url": posterURL,
+ "backdrop_url": backdropURL,
"rating": match.Rating,
"year": match.Year,
"scrape_status": "matched",
@@ -464,6 +403,9 @@ func (s *ScraperService) applyProviderMatch(ctx context.Context, m *model.Media,
if match.OriginalName != "" {
updates["original_name"] = match.OriginalName
}
+ if strings.TrimSpace(m.EpisodeTitle) != "" {
+ updates["episode_title"] = strings.TrimSpace(m.EpisodeTitle)
+ }
if match.TMDbID > 0 {
updates["tm_db_id"] = match.TMDbID
}
@@ -493,104 +435,22 @@ func (s *ScraperService) applyProviderMatch(ctx context.Context, m *model.Media,
Updates(updates).Error; err != nil {
return err
}
+ s.removeCachedScrapedArtwork(removePoster, removeBackdrop)
// Fetch extended metadata after the selected match is already saved.
// Manual cloud/batch applies must not fail just because an optional provider
// details request is slow or unavailable.
if match.TMDbID > 0 && s.tmdb != nil && s.tmdb.Enabled() {
mediaType := s.determineMediaTypeForMedia(lib, m, match)
- detailCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), tmdbDetailsTimeout)
- details, err := s.tmdb.GetDetails(detailCtx, match.TMDbID, mediaType)
- cancel()
- if err != nil {
- s.log.Warn("failed to get details from tmdb",
- zap.Int("tmdb_id", match.TMDbID),
- zap.String("type", mediaType),
- zap.Error(err))
- } else if details != nil {
- detailUpdates := map[string]any{}
- if len(details.Languages) > 0 {
- detailUpdates["languages"] = strings.Join(details.Languages, ",")
- }
- if len(details.Countries) > 0 {
- detailUpdates["countries"] = strings.Join(details.Countries, ",")
- }
- if len(details.Genres) > 0 {
- detailUpdates["genres"] = strings.Join(details.Genres, ",")
- }
- if len(detailUpdates) > 0 {
- if err := s.repo.DB.Model(&model.Media{}).Where("id = ?", m.ID).
- Updates(detailUpdates).Error; err != nil {
- s.log.Warn("failed to save tmdb extended metadata",
- zap.String("media_id", m.ID),
- zap.Int("tmdb_id", match.TMDbID),
- zap.Error(err))
- }
- }
- s.log.Debug("enrich: saved extended metadata",
- zap.String("media_id", m.ID),
- zap.Strings("languages", details.Languages),
- zap.Strings("countries", details.Countries),
- zap.Strings("genres", details.Genres))
- }
- if mediaType == "tv" && m != nil && m.EpisodeNum > 0 {
- episodeCtx, episodeCancel := context.WithTimeout(context.WithoutCancel(ctx), tmdbDetailsTimeout)
- episode, err := s.tmdb.GetTVEpisodeDetails(episodeCtx, match.TMDbID, m.SeasonNum, m.EpisodeNum)
- episodeCancel()
- if err != nil {
- s.log.Debug("failed to get tmdb episode details",
- zap.String("media_id", m.ID),
- zap.Int("tmdb_id", match.TMDbID),
- zap.Int("season", m.SeasonNum),
- zap.Int("episode", m.EpisodeNum),
- zap.Error(err))
- } else if episode != nil {
- episodeUpdates := map[string]any{}
- // 注意: 不要把 episode.Name(单集名,如"觉醒"/"Pilot"/"第1集")写入
- // original_name —— 该字段是「整剧原名」,是合集分组键的回退依据。
- // 若每集都写成各自的单集名,同一部剧的各集 original_name 互不相同,
- // 会被前端 getSeriesKey / 后端 mediaVersionGroupKey 拆成多个独立卡片,
- // 导致同剧无法合并成合集。单集名属于单集信息,这里只回填不影响合集
- // 分组的单集字段(简介/剧照/评分/时长)。
- if strings.TrimSpace(episode.Overview) != "" {
- episodeUpdates["overview"] = strings.TrimSpace(episode.Overview)
- }
- if strings.TrimSpace(episode.StillURL) != "" {
- episodeUpdates["backdrop_url"] = strings.TrimSpace(episode.StillURL)
- }
- if episode.Rating > 0 {
- episodeUpdates["rating"] = episode.Rating
- }
- if episode.AirYear > 0 && match.Year <= 0 {
- episodeUpdates["year"] = episode.AirYear
- }
- if episode.Runtime > 0 && m.DurationSec <= 0 {
- episodeUpdates["duration_sec"] = episode.Runtime * 60
- }
- if len(episodeUpdates) > 0 {
- if err := s.repo.DB.Model(&model.Media{}).Where("id = ?", m.ID).
- Updates(episodeUpdates).Error; err != nil {
- s.log.Warn("failed to save tmdb episode metadata",
- zap.String("media_id", m.ID),
- zap.Int("tmdb_id", match.TMDbID),
- zap.Int("season", m.SeasonNum),
- zap.Int("episode", m.EpisodeNum),
- zap.Error(err))
- }
- }
- }
+ s.fetchAndSaveTMDbExtendedMetadata(ctx, m.ID, match.TMDbID, mediaType)
+ if mediaType == "tv" && !options.DeferEpisodeDetails {
+ s.fetchAndSaveTMDbEpisodeDetails(ctx, m, match.TMDbID, match.Year, options)
}
}
- cloudMedia := isCloudMediaPath(m.Path) || (lib != nil && isCloudMediaPath(lib.Path))
- if !cloudMedia {
- if refreshed, err := s.repo.Media.FindByID(ctx, m.ID); err == nil && refreshed != nil {
- if path, err := WriteMediaNFO(refreshed); err != nil {
- s.log.Warn("write nfo after scrape failed", zap.String("media_id", m.ID), zap.Error(err))
- } else {
- s.log.Debug("write nfo after scrape", zap.String("media_id", m.ID), zap.String("path", path))
- }
- }
+ if !(options.DeferEpisodeDetails && m != nil && m.EpisodeNum > 0) {
+ s.writeMediaNFOAfterScrape(ctx, m, lib)
}
+ s.invalidateMediaCache(ctx)
s.hub.Publish("scrape", map[string]any{
"media_id": m.ID,
"title": match.Title,
@@ -603,6 +463,38 @@ func (s *ScraperService) applyProviderMatch(ctx context.Context, m *model.Media,
return nil
}
+func mergeScrapePathHintMetadata(dst, src *LocalMetadata) *LocalMetadata {
+ if src == nil {
+ return dst
+ }
+ if dst == nil {
+ return cloneLocalMetadata(src)
+ }
+ hasLocalMetadata := localMetadataMarksMatched(dst)
+ if dst.Title == "" && src.Title != "" {
+ dst.Title = src.Title
+ }
+ if dst.Year <= 0 && src.Year > 0 {
+ dst.Year = src.Year
+ }
+ if src.TMDbID > 0 {
+ dst.TMDbID = src.TMDbID
+ }
+ if src.BangumiID > 0 {
+ dst.BangumiID = src.BangumiID
+ }
+ if strings.TrimSpace(src.DoubanID) != "" {
+ dst.DoubanID = strings.TrimSpace(src.DoubanID)
+ }
+ if strings.TrimSpace(src.TheTVDBID) != "" {
+ dst.TheTVDBID = strings.TrimSpace(src.TheTVDBID)
+ }
+ if !hasLocalMetadata {
+ dst.PathHint = dst.PathHint || src.PathHint
+ }
+ return dst
+}
+
func isCloudMediaPath(value string) bool {
return strings.HasPrefix(strings.ToLower(strings.TrimSpace(value)), "cloud://")
}
@@ -621,6 +513,9 @@ func (s *ScraperService) applyLocalMetadataMatch(ctx context.Context, m *model.M
if next.OriginalName != "" {
updates["original_name"] = next.OriginalName
}
+ if next.EpisodeTitle != "" {
+ updates["episode_title"] = next.EpisodeTitle
+ }
if next.Overview != "" {
updates["overview"] = next.Overview
}
@@ -670,6 +565,7 @@ func (s *ScraperService) applyLocalMetadataMatch(ctx context.Context, m *model.M
Where("id = ?", m.ID).Updates(updates).Error; err != nil {
return err
}
+ s.invalidateMediaCache(ctx)
s.hub.Publish("scrape", map[string]any{
"media_id": m.ID,
"title": next.Title,
@@ -680,316 +576,11 @@ func (s *ScraperService) applyLocalMetadataMatch(ctx context.Context, m *model.M
return nil
}
-func scrapeQueryCandidates(m *model.Media, lib *model.Library) []string {
- seen := map[string]struct{}{}
- var out []string
- add := func(raw string) {
- cleaned, _ := CleanQuery(raw)
- if cleaned == "" {
- cleaned = strings.TrimSpace(raw)
- }
- for _, candidate := range titleCandidates(cleaned) {
- key := strings.ToLower(candidate)
- if _, ok := seen[key]; ok || candidate == "" {
- continue
- }
- seen[key] = struct{}{}
- out = append(out, candidate)
- }
+func (s *ScraperService) invalidateMediaCache(ctx context.Context) {
+ if s != nil && s.cache != nil {
+ s.cache.DeletePrefix(ctx, "media:")
+ s.cache.DeletePrefix(ctx, "stats:")
}
- if lib != nil && mediaIsEpisodic(m, lib) {
- add(seriesFolderTitle(m.Path, lib.Path))
- }
- add(m.Title)
- add(m.Path)
- if len(out) == 0 {
- out = append(out, strings.TrimSuffix(filepath.Base(m.Path), filepath.Ext(m.Path)))
- }
- return out
-}
-
-func seriesFolderTitle(mediaPath, libraryRoot string) string {
- dir := filepath.Dir(mediaPath)
- if strictSeasonFolderMatched(filepath.Base(dir)) {
- dir = filepath.Dir(dir)
- }
- if libraryRoot != "" && samePath(dir, filepath.Clean(libraryRoot)) {
- return ""
- }
- base := filepath.Base(dir)
- if base == "." || base == string(filepath.Separator) {
- return ""
- }
- if isGenericMediaCategoryFolder(base) {
- return ""
- }
- return base
-}
-
-func isGenericMediaCategoryFolder(name string) bool {
- key := strings.ToLower(strings.TrimSpace(name))
- key = strings.Trim(key, `\/`)
- switch key {
- case "",
- "电影", "movies", "movie",
- "电视剧", "剧集", "tv", "shows", "series",
- "动漫", "动画", "anime", "bangumi",
- "国产剧", "国剧", "大陆剧", "国产电视剧",
- "欧美剧", "欧美电视剧",
- "日韩剧", "日剧", "韩剧",
- "华语电影", "国产电影", "大陆电影",
- "外语电影", "欧美电影", "日韩电影",
- "动画电影", "动漫电影",
- "国漫", "国产动漫", "日番", "日漫", "日本动漫", "日本动画",
- "综艺", "真人秀",
- "纪录片", "纪录",
- "儿童", "少儿",
- "成人", "番号", "9kg",
- "未分类", "uncategorized":
- return true
- default:
- return false
- }
-}
-
-func strictSeasonFolder(name string) int {
- if season, ok := seasonFromDir(name); ok {
- return season
- }
- return 0
-}
-
-func strictSeasonFolderMatched(name string) bool {
- _, ok := seasonFromDir(name)
- return ok
-}
-
-func titleCandidates(title string) []string {
- title = strings.Join(strings.Fields(strings.TrimSpace(title)), " ")
- if title == "" {
- return nil
- }
- out := make([]string, 0, 2)
- if cjk := cjkTitleOnly(title); cjk != "" {
- out = append(out, cjk)
- if cjk != title {
- return out
- }
- }
- out = append(out, title)
- return out
-}
-
-func cjkTitleOnly(title string) string {
- parts := make([]string, 0, 4)
- for _, field := range strings.Fields(title) {
- if containsCJK(field) {
- parts = append(parts, field)
- }
- }
- return strings.Join(parts, " ")
-}
-
-func containsCJK(s string) bool {
- for _, r := range s {
- switch {
- case r >= '\u3400' && r <= '\u4dbf':
- return true
- case r >= '\u4e00' && r <= '\u9fff':
- return true
- case r >= '\uf900' && r <= '\ufaff':
- return true
- }
- }
- return false
-}
-
-func mediaIsEpisodic(m *model.Media, lib *model.Library) bool {
- if m != nil && (m.SeasonNum > 0 || m.EpisodeNum > 0) {
- return true
- }
- return librarySupportsSeasons(lib)
-}
-
-func librarySupportsSeasons(lib *model.Library) bool {
- if lib == nil {
- return false
- }
- switch strings.ToLower(strings.TrimSpace(lib.Type)) {
- case "tv", "anime", "variety", "show", "shows":
- return true
- default:
- return false
- }
-}
-
-// lookup runs the provider chain after local NFO has been considered:
-// TMDb -> Douban -> Bangumi -> TheTVDB. Douban and Bangumi do not require API
-// keys; providers that are unavailable or return an error are skipped.
-func (s *ScraperService) lookup(ctx context.Context, lib *model.Library, query string, year int) *Match {
- kind := ""
- if lib != nil {
- kind = lib.Type
- }
- if s.tmdb != nil && s.tmdb.Enabled() {
- // anime / tv 先用 TMDb /search/tv(剧名通常是 TV 类目)。
- if kind == "anime" || kind == "tv" || kind == "variety" || kind == "show" || kind == "shows" {
- if m, err := s.tmdb.SearchTV(ctx, query, year); err == nil && m != nil {
- return m
- } else if err != nil {
- s.log.Debug("tmdb tv search failed", zap.String("query", query), zap.Error(err))
- }
- }
- if m, err := s.tmdb.SearchMovie(ctx, query, year); err == nil && m != nil {
- return m
- } else if err != nil {
- s.log.Debug("tmdb movie search failed", zap.String("query", query), zap.Error(err))
- }
- }
- if s.douban != nil && s.douban.Enabled() {
- if m, err := s.douban.SearchMatch(ctx, query); err == nil && m != nil {
- return m
- } else if err != nil {
- s.log.Debug("douban search failed", zap.String("query", query), zap.Error(err))
- }
- }
- if s.bangumi != nil && s.bangumi.Enabled() {
- if m, err := s.bangumi.Search(ctx, query); err == nil && m != nil {
- return m
- } else if err != nil {
- s.log.Debug("bangumi search failed", zap.String("query", query), zap.Error(err))
- }
- }
- if (kind == "anime" || kind == "tv" || kind == "variety" || kind == "show" || kind == "shows") && s.thetvdb != nil && s.thetvdb.Enabled() {
- if m, err := s.thetvdb.SearchSeries(ctx, query); err == nil && m != nil {
- return m
- } else if err != nil {
- s.log.Debug("thetvdb search failed", zap.String("query", query), zap.Error(err))
- }
- }
- return nil
-}
-
-// EnrichLibrary runs the provider chain for every pending media in a library.
-// When retryNoMatch is true it also retries rows previously marked no_match,
-// which is the expected behaviour for a manual "重新刮削" action. Scanner-driven
-// automatic enrichment keeps the default false path to avoid repeated scraping.
-//
-// Pending status includes both the canonical "pending" string and the
-// empty / NULL values, because MediaRepository.Upsert can wipe the GORM
-// default when re-running a scan over an already-existing row.
-func (s *ScraperService) EnrichLibrary(ctx context.Context, libraryID string, retryNoMatch ...bool) (int, error) {
- var rows []model.Media
- statusFilter := "scrape_status IS NULL OR scrape_status = '' OR scrape_status = ?"
- statusArgs := []any{"pending"}
- if len(retryNoMatch) > 0 && retryNoMatch[0] {
- statusFilter += " OR scrape_status = ?"
- statusArgs = append(statusArgs, "no_match")
- }
- q := s.repo.DB.Where(statusFilter, statusArgs...)
- if libraryID != "" {
- q = q.Where("library_id = ?", libraryID)
- }
- if err := q.Find(&rows).Error; err != nil {
- return 0, err
- }
- matched := 0
- processed := 0
- for i := range rows {
- select {
- case <-ctx.Done():
- return matched, ctx.Err()
- default:
- }
- if err := s.EnrichOne(ctx, &rows[i]); err != nil {
- s.log.Warn("enrich failed", zap.String("media", rows[i].ID), zap.Error(err))
- s.notifyScrapeFailed(rows[i], err)
- continue
- }
- processed++
- if s.mediaIsMatched(ctx, rows[i].ID) {
- matched++
- }
- if i < len(rows)-1 {
- if delay := s.scrapeDelay(ctx); delay > 0 {
- select {
- case <-ctx.Done():
- return matched, ctx.Err()
- case <-time.After(delay):
- }
- }
- }
- }
- s.hub.Publish("scrape", map[string]any{
- "library_id": libraryID,
- "finished": true,
- "matched": matched,
- "processed": processed,
- })
- return matched, nil
-}
-
-func (s *ScraperService) notifyScrapeFailed(m model.Media, err error) {
- if s == nil || s.notify == nil || err == nil {
- return
- }
- body := strings.TrimSpace(m.Title)
- if body == "" {
- body = m.Path
- }
- body = "媒体:" + body + "\n错误:" + err.Error()
- go func() {
- ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
- defer cancel()
- s.notify.Broadcast(ctx, "MediaStationGo 刮削失败", body, EventScrapeFailed)
- }()
-}
-
-func (s *ScraperService) scrapeDelay(ctx context.Context) time.Duration {
- minMS := s.scrapeDelaySetting(ctx, "scrape.delay_min_ms", defaultScrapeDelayMinMS)
- maxMS := s.scrapeDelaySetting(ctx, "scrape.delay_max_ms", defaultScrapeDelayMaxMS)
- if minMS < 0 {
- minMS = 0
- }
- if maxMS < 0 {
- maxMS = 0
- }
- if minMS > maxScrapeDelayMS {
- minMS = maxScrapeDelayMS
- }
- if maxMS > maxScrapeDelayMS {
- maxMS = maxScrapeDelayMS
- }
- if maxMS < minMS {
- maxMS = minMS
- }
- if maxMS == 0 {
- return 0
- }
- if maxMS == minMS {
- return time.Duration(minMS) * time.Millisecond
- }
- return time.Duration(minMS+secureRandomIntn(maxMS-minMS+1)) * time.Millisecond
-}
-
-func (s *ScraperService) scrapeDelaySetting(ctx context.Context, key string, fallback int) int {
- if s == nil || s.repo == nil || s.repo.Setting == nil {
- return fallback
- }
- value, err := s.repo.Setting.Get(ctx, key)
- if err != nil || strings.TrimSpace(value) == "" {
- return fallback
- }
- return parseIntSettingDefault(strings.TrimSpace(value), fallback)
-}
-
-func (s *ScraperService) mediaIsMatched(ctx context.Context, mediaID string) bool {
- var status string
- err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).
- Select("scrape_status").
- Where("id = ?", mediaID).
- Scan(&status).Error
- return err == nil && status == "matched"
}
// AnyEnabled reports whether at least one provider can run.
diff --git a/internal/service/scraper_artwork.go b/internal/service/scraper_artwork.go
new file mode 100644
index 0000000..eab2a72
--- /dev/null
+++ b/internal/service/scraper_artwork.go
@@ -0,0 +1,67 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+)
+
+func (s *ScraperService) prepareScrapedArtworkURL(ctx context.Context, mediaID, field, current, candidate string) (string, string) {
+ current = strings.TrimSpace(current)
+ candidate = strings.TrimSpace(candidate)
+ if candidate == "" {
+ return current, ""
+ }
+ if candidate == current {
+ return candidate, ""
+ }
+ if s == nil || s.images == nil || !isHTTPish(candidate) {
+ return candidate, ""
+ }
+ fetchCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second)
+ err := s.images.PrefetchRemote(fetchCtx, candidate)
+ cancel()
+ if err != nil {
+ if current != "" {
+ s.log.Warn("scrape artwork prefetch failed; keeping existing artwork",
+ zap.String("media_id", mediaID),
+ zap.String("field", field),
+ zap.String("candidate", candidate),
+ zap.String("existing", current),
+ zap.Error(err))
+ return current, ""
+ }
+ s.log.Warn("scrape artwork prefetch failed; keeping new artwork URL for retry",
+ zap.String("media_id", mediaID),
+ zap.String("field", field),
+ zap.String("candidate", candidate),
+ zap.Error(err))
+ return candidate, ""
+ }
+ if current != "" && isHTTPish(current) {
+ return candidate, current
+ }
+ return candidate, ""
+}
+
+func (s *ScraperService) removeCachedScrapedArtwork(urls ...string) {
+ if s == nil || s.images == nil {
+ return
+ }
+ seen := map[string]struct{}{}
+ for _, raw := range urls {
+ raw = strings.TrimSpace(raw)
+ if raw == "" {
+ continue
+ }
+ if _, ok := seen[raw]; ok {
+ continue
+ }
+ seen[raw] = struct{}{}
+ if err := s.images.RemoveCached(raw); err != nil {
+ s.log.Debug("remove old scraped artwork cache failed", zap.String("url", raw), zap.Error(err))
+ }
+ }
+}
diff --git a/internal/service/scraper_artwork_test.go b/internal/service/scraper_artwork_test.go
new file mode 100644
index 0000000..de65d48
--- /dev/null
+++ b/internal/service/scraper_artwork_test.go
@@ -0,0 +1,213 @@
+package service
+
+import (
+ "errors"
+ "io"
+ "net/http"
+ "os"
+ "path/filepath"
+ "strings"
+ "testing"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestScrapeDelayUsesSettings(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+ if err := repos.DB.AutoMigrate(&model.Setting{}); err != nil {
+ t.Fatal(err)
+ }
+
+ if got := scraper.scrapeDelay(t.Context()); got < 250*time.Millisecond || got > 500*time.Millisecond {
+ t.Fatalf("default scrapeDelay = %s, want 250-500ms", got)
+ }
+
+ if err := repos.Setting.Set(t.Context(), "scrape.delay_min_ms", "0"); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Setting.Set(t.Context(), "scrape.delay_max_ms", "0"); err != nil {
+ t.Fatal(err)
+ }
+ if got := scraper.scrapeDelay(t.Context()); got != 0 {
+ t.Fatalf("disabled scrapeDelay = %s, want 0", got)
+ }
+
+ if err := repos.Setting.Set(t.Context(), "scrape.delay_min_ms", "800"); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Setting.Set(t.Context(), "scrape.delay_max_ms", "200"); err != nil {
+ t.Fatal(err)
+ }
+ if got := scraper.scrapeDelay(t.Context()); got != 800*time.Millisecond {
+ t.Fatalf("normalized scrapeDelay = %s, want 800ms", got)
+ }
+}
+
+func TestApplyProviderMatchInvalidatesMediaCache(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ cache := NewRuntimeCacheService(&config.Config{}, zap.NewNop())
+ cache.SetJSON(t.Context(), "media:list:stale", map[string]string{"poster": ""}, time.Minute)
+ cache.SetJSON(t.Context(), "stats:snapshot:base", map[string]int{"media": 1}, time.Minute)
+ scraper.SetRuntimeCache(cache)
+
+ libPath := t.TempDir()
+ lib := model.Library{Name: "Movies", Path: libPath, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{LibraryID: lib.ID, Title: "Raw", Path: "/media/movies/raw.mkv", ScrapeStatus: "pending"}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ match := &Match{Title: "Matched", PosterURL: "https://image.tmdb.org/t/p/w500/poster.jpg", BackdropURL: "https://image.tmdb.org/t/p/w1280/backdrop.jpg"}
+ if err := scraper.applyProviderMatch(t.Context(), &media, &lib, match); err != nil {
+ t.Fatal(err)
+ }
+
+ var stale map[string]string
+ if cache.GetJSON(t.Context(), "media:list:stale", &stale) {
+ t.Fatal("scraper should invalidate media list cache after applying artwork")
+ }
+ var stats map[string]int
+ if cache.GetJSON(t.Context(), "stats:snapshot:base", &stats) {
+ t.Fatal("scraper should invalidate stats cache after applying artwork")
+ }
+ var stored model.Media
+ if err := repos.DB.First(&stored, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if stored.PosterURL == "" || stored.BackdropURL == "" || stored.ScrapeStatus != "matched" {
+ t.Fatalf("match not saved: poster=%q backdrop=%q status=%q", stored.PosterURL, stored.BackdropURL, stored.ScrapeStatus)
+ }
+}
+
+func TestApplyProviderMatchKeepsExistingArtworkWhenNewPrefetchFails(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ images := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: filepath.Join(t.TempDir(), "cache")}}, zap.NewNop())
+ images.client = &http.Client{Transport: imageRoundTripFunc(func(req *http.Request) (*http.Response, error) {
+ return &http.Response{
+ StatusCode: http.StatusBadGateway,
+ Status: "502 Bad Gateway",
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader("bad gateway")),
+ Request: req,
+ }, nil
+ })}
+ scraper.SetImageProxy(images)
+
+ libPath := t.TempDir()
+ lib := model.Library{Name: "Movies", Path: libPath, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ oldPoster := "https://image.tmdb.org/t/p/w500/old-poster.jpg"
+ oldBackdrop := "https://image.tmdb.org/t/p/w1280/old-backdrop.jpg"
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "Raw",
+ Path: filepath.Join(libPath, "raw.mkv"),
+ PosterURL: oldPoster,
+ BackdropURL: oldBackdrop,
+ ScrapeStatus: "matched",
+ OriginalName: "Raw",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ match := &Match{
+ Title: "Matched",
+ PosterURL: "https://image.tmdb.org/t/p/w500/new-broken-poster.jpg",
+ BackdropURL: "",
+ }
+ if err := scraper.applyProviderMatch(t.Context(), &media, &lib, match); err != nil {
+ t.Fatal(err)
+ }
+
+ var stored model.Media
+ if err := repos.DB.First(&stored, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if stored.PosterURL != oldPoster {
+ t.Fatalf("poster should keep existing URL when new prefetch fails: got %q want %q", stored.PosterURL, oldPoster)
+ }
+ if stored.BackdropURL != oldBackdrop {
+ t.Fatalf("blank match backdrop should not clear existing backdrop: got %q want %q", stored.BackdropURL, oldBackdrop)
+ }
+}
+
+func TestApplyProviderMatchReplacesArtworkAndRemovesOldCache(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ images := NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: filepath.Join(t.TempDir(), "cache")}}, zap.NewNop())
+ images.client = &http.Client{Transport: imageRoundTripFunc(func(req *http.Request) (*http.Response, error) {
+ body := "image:" + req.URL.Path
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Status: "200 OK",
+ Header: http.Header{"Content-Type": []string{"image/jpeg"}},
+ Body: io.NopCloser(strings.NewReader(body)),
+ Request: req,
+ }, nil
+ })}
+ scraper.SetImageProxy(images)
+
+ libPath := t.TempDir()
+ lib := model.Library{Name: "Movies", Path: libPath, Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ oldPoster := "https://image.tmdb.org/t/p/w500/old-cache-poster.jpg"
+ newPoster := "https://image.tmdb.org/t/p/w500/new-cache-poster.jpg"
+ if err := images.PrefetchRemote(t.Context(), oldPoster); err != nil {
+ t.Fatal(err)
+ }
+ _, oldCachePath, _, err := images.remoteImageCachePaths(oldPoster)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err := os.Stat(oldCachePath); err != nil {
+ t.Fatalf("old poster cache should exist before replace: %v", err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "Raw",
+ Path: filepath.Join(libPath, "raw.mkv"),
+ PosterURL: oldPoster,
+ ScrapeStatus: "matched",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ match := &Match{Title: "Matched", PosterURL: newPoster}
+ if err := scraper.applyProviderMatch(t.Context(), &media, &lib, match); err != nil {
+ t.Fatal(err)
+ }
+
+ var stored model.Media
+ if err := repos.DB.First(&stored, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if stored.PosterURL != newPoster {
+ t.Fatalf("poster = %q, want %q", stored.PosterURL, newPoster)
+ }
+ if _, err := os.Stat(oldCachePath); !errors.Is(err, os.ErrNotExist) {
+ t.Fatalf("old poster cache should be removed, stat err=%v", err)
+ }
+ _, newCachePath, _, err := images.remoteImageCachePaths(newPoster)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err := os.Stat(newCachePath); err != nil {
+ t.Fatalf("new poster cache should exist: %v", err)
+ }
+}
diff --git a/internal/service/scraper_episode_details_test.go b/internal/service/scraper_episode_details_test.go
new file mode 100644
index 0000000..a108179
--- /dev/null
+++ b/internal/service/scraper_episode_details_test.go
@@ -0,0 +1,210 @@
+package service
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "path/filepath"
+ "strings"
+ "sync"
+ "testing"
+
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestEnrichLibraryDefersEpisodeDetailsUntilMainMetadataFinishes(t *testing.T) {
+ var mu sync.Mutex
+ paths := []string{}
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ mu.Lock()
+ paths = append(paths, r.URL.Path)
+ mu.Unlock()
+
+ w.Header().Set("Content-Type", "application/json")
+ switch {
+ case strings.HasPrefix(r.URL.Path, "/search/tv"):
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 12345,
+ "name": "间谍过家家",
+ "original_name": "SPY×FAMILY",
+ "overview": "测试简介",
+ "poster_path": "/poster.jpg",
+ "backdrop_path": "/backdrop.jpg",
+ "first_air_date": "2022-04-09",
+ "vote_average": 8.6,
+ }},
+ })
+ case r.URL.Path == "/tv/12345/season/2/episode/1":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "name": "任务代号: 猫",
+ "overview": "第一集剧情",
+ "still_path": "/still-1.jpg",
+ "vote_average": 8.9,
+ "runtime": 24,
+ })
+ case r.URL.Path == "/tv/12345/season/2/episode/2":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "name": "接近目标",
+ "overview": "第二集剧情",
+ "still_path": "/still-2.jpg",
+ "vote_average": 9.0,
+ "runtime": 25,
+ })
+ case r.URL.Path == "/tv/12345":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": 12345,
+ "name": "间谍过家家",
+ "overview": "测试简介",
+ "poster_path": "/poster.jpg",
+ "backdrop_path": "/backdrop.jpg",
+ "first_air_date": "2022-04-09",
+ "vote_average": 8.6,
+ "origin_country": []string{"JP"},
+ "spoken_languages": []map[string]any{{
+ "iso_639_1": "ja",
+ }},
+ "genres": []map[string]any{{
+ "name": "Animation",
+ }},
+ })
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ cfg.Secrets.TMDbImageProxy = upstream.URL + "/images"
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ rows := []model.Media{
+ {
+ Base: model.Base{ID: "episode-1"},
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ ScrapeStatus: "pending",
+ },
+ {
+ Base: model.Base{ID: "episode-2"},
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E02.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 2,
+ ScrapeStatus: "pending",
+ },
+ }
+ if err := repos.DB.Create(&rows).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ result, err := scraper.EnrichLibraryDetailed(t.Context(), lib.ID, true)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if result.Processed != 2 || result.Matched != 2 {
+ t.Fatalf("result=%+v, want two matched episodes", result)
+ }
+
+ mu.Lock()
+ gotPaths := append([]string(nil), paths...)
+ mu.Unlock()
+ firstEpisodeDetail := firstIndexFunc(gotPaths, func(path string) bool {
+ return strings.Contains(path, "/season/")
+ })
+ lastMainMetadata := lastIndexFunc(gotPaths, func(path string) bool {
+ return strings.HasPrefix(path, "/search/tv") || path == "/tv/12345"
+ })
+ if firstEpisodeDetail < 0 {
+ t.Fatalf("no deferred episode detail requests recorded: %v", gotPaths)
+ }
+ if firstEpisodeDetail <= lastMainMetadata {
+ t.Fatalf("episode detail ran before main metadata finished: paths=%v", gotPaths)
+ }
+
+ var stored []model.Media
+ if err := repos.DB.Where("library_id = ?", lib.ID).Order("episode_num ASC").Find(&stored).Error; err != nil {
+ t.Fatal(err)
+ }
+ if len(stored) != 2 || stored[0].Overview != "第一集剧情" || stored[1].Overview != "第二集剧情" {
+ t.Fatalf("deferred episode metadata not saved: %+v", stored)
+ }
+ if stored[0].EpisodeTitle != "任务代号: 猫" || stored[1].EpisodeTitle != "接近目标" {
+ t.Fatalf("deferred episode titles not saved: %+v", stored)
+ }
+ if stored[0].OriginalName != "SPY×FAMILY" || stored[1].OriginalName != "SPY×FAMILY" {
+ t.Fatalf("series original_name should stay shared: %+v", stored)
+ }
+}
+
+func TestEnrichLibrarySkipsDeferredEpisodeStillWhenDisabled(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ episodeArtwork := false
+ result, err := scraper.EnrichLibraryDetailedWithOptions(t.Context(), lib.ID, ScrapeOptions{
+ RetryNoMatch: true,
+ EpisodeArtwork: &episodeArtwork,
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if result.Processed != 1 || result.Matched != 1 {
+ t.Fatalf("result=%+v, want one matched episode", result)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.Overview != "单集剧情" || got.DurationSec != 24*60 {
+ t.Fatalf("deferred episode text metadata should still be saved: overview=%q duration=%d", got.Overview, got.DurationSec)
+ }
+ if strings.HasSuffix(got.BackdropURL, "/images/w500/still.jpg") {
+ t.Fatalf("deferred episode still should not be saved when disabled: backdrop=%q", got.BackdropURL)
+ }
+ if !strings.HasSuffix(got.BackdropURL, "/images/w1280/backdrop.jpg") {
+ t.Fatalf("series backdrop should remain available when episode still is disabled: got %q", got.BackdropURL)
+ }
+}
diff --git a/internal/service/scraper_library.go b/internal/service/scraper_library.go
new file mode 100644
index 0000000..662c28a
--- /dev/null
+++ b/internal/service/scraper_library.go
@@ -0,0 +1,235 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// lookup runs the provider chain after local NFO has been considered:
+// TMDb -> Douban -> Bangumi -> TheTVDB. Douban and Bangumi do not require API
+// keys; providers that are unavailable or return an error are skipped.
+func (s *ScraperService) lookup(ctx context.Context, lib *model.Library, media *model.Media, query string, year int) *Match {
+ kind := ""
+ if lib != nil {
+ kind = lib.Type
+ }
+ if mediaIsEpisodic(media, lib) {
+ kind = "tv"
+ }
+ if s.tmdb != nil && s.tmdb.Enabled() {
+ // anime / tv 先用 TMDb /search/tv(剧名通常是 TV 类目)。
+ if kind == "anime" || kind == "tv" || kind == "variety" || kind == "show" || kind == "shows" {
+ if m, err := s.tmdb.SearchTV(ctx, query, year); err == nil && m != nil {
+ return m
+ } else if err != nil {
+ s.log.Debug("tmdb tv search failed", zap.String("query", query), zap.Error(err))
+ }
+ }
+ if m, err := s.tmdb.SearchMovie(ctx, query, year); err == nil && m != nil {
+ return m
+ } else if err != nil {
+ s.log.Debug("tmdb movie search failed", zap.String("query", query), zap.Error(err))
+ }
+ }
+ if s.douban != nil && s.douban.Enabled() {
+ if m, err := s.douban.SearchMatch(ctx, query); err == nil && m != nil {
+ return m
+ } else if err != nil {
+ s.log.Debug("douban search failed", zap.String("query", query), zap.Error(err))
+ }
+ }
+ if s.bangumi != nil && s.bangumi.Enabled() {
+ if m, err := s.bangumi.Search(ctx, query); err == nil && m != nil {
+ return m
+ } else if err != nil {
+ s.log.Debug("bangumi search failed", zap.String("query", query), zap.Error(err))
+ }
+ }
+ if (kind == "anime" || kind == "tv" || kind == "variety" || kind == "show" || kind == "shows") && s.thetvdb != nil && s.thetvdb.Enabled() {
+ if m, err := s.thetvdb.SearchSeries(ctx, query); err == nil && m != nil {
+ return m
+ } else if err != nil {
+ s.log.Debug("thetvdb search failed", zap.String("query", query), zap.Error(err))
+ }
+ }
+ return nil
+}
+
+// EnrichLibrary runs the provider chain for every pending media in a library.
+// When retryNoMatch is true it also retries rows previously marked no_match,
+// which is the expected behaviour for a manual "重新刮削" action. Scanner-driven
+// automatic enrichment keeps the default false path to avoid repeated scraping.
+//
+// Pending status includes both the canonical "pending" string and the
+// empty / NULL values, because MediaRepository.Upsert can wipe the GORM
+// default when re-running a scan over an already-existing row.
+func (s *ScraperService) EnrichLibrary(ctx context.Context, libraryID string, retryNoMatch ...bool) (int, error) {
+ result, err := s.EnrichLibraryDetailed(ctx, libraryID, retryNoMatch...)
+ return result.Matched, err
+}
+
+type EnrichLibraryResult struct {
+ LibraryID string
+ Matched int
+ Processed int
+ Failed int
+ Candidates int
+}
+
+func (s *ScraperService) EnrichLibraryDetailed(ctx context.Context, libraryID string, retryNoMatch ...bool) (EnrichLibraryResult, error) {
+ options := ScrapeOptions{}
+ if len(retryNoMatch) > 0 {
+ options.RetryNoMatch = retryNoMatch[0]
+ }
+ return s.EnrichLibraryDetailedWithOptions(ctx, libraryID, options)
+}
+
+func (s *ScraperService) EnrichLibraryDetailedWithOptions(ctx context.Context, libraryID string, options ScrapeOptions) (EnrichLibraryResult, error) {
+ result := EnrichLibraryResult{LibraryID: libraryID}
+ rows, err := s.scrapeCandidateRows(ctx, libraryID, options)
+ if err != nil {
+ return result, err
+ }
+ result.Candidates = len(rows)
+ runOptions := options
+ runOptions.DeferEpisodeDetails = true
+ for i := range rows {
+ select {
+ case <-ctx.Done():
+ return result, ctx.Err()
+ default:
+ }
+ if err := s.EnrichOneWithOptions(ctx, &rows[i], runOptions); err != nil {
+ s.log.Warn("enrich failed", zap.String("media", rows[i].ID), zap.Error(err))
+ s.notifyScrapeFailed(rows[i], err)
+ result.Failed++
+ continue
+ }
+ result.Processed++
+ if s.mediaIsMatched(ctx, rows[i].ID) {
+ result.Matched++
+ }
+ if i < len(rows)-1 {
+ if delay := s.scrapeDelay(ctx); delay > 0 {
+ select {
+ case <-ctx.Done():
+ return result, ctx.Err()
+ case <-time.After(delay):
+ }
+ }
+ }
+ }
+ if err := s.enrichDeferredEpisodeDetails(ctx, rows, options); err != nil {
+ return result, err
+ }
+ s.hub.Publish("scrape", map[string]any{
+ "library_id": libraryID,
+ "finished": true,
+ "matched": result.Matched,
+ "processed": result.Processed,
+ "failed": result.Failed,
+ "candidates": result.Candidates,
+ })
+ return result, nil
+}
+
+func (s *ScraperService) scrapeCandidateRows(ctx context.Context, libraryID string, options ScrapeOptions) ([]model.Media, error) {
+ var rows []model.Media
+ libraryIDs := []string{}
+ if strings.TrimSpace(libraryID) != "" {
+ var err error
+ libraryIDs, err = MergedLibraryIDsForLibrary(ctx, s.repo, libraryID)
+ if err != nil {
+ return nil, err
+ }
+ }
+ statusFilter := "scrape_status IS NULL OR scrape_status = '' OR scrape_status = ?"
+ statusArgs := []any{"pending"}
+ if options.RetryNoMatch {
+ statusFilter += " OR scrape_status = ?"
+ statusArgs = append(statusArgs, "no_match")
+ }
+ if options.IncludeMatched {
+ statusFilter += " OR scrape_status = ?"
+ statusArgs = append(statusArgs, "matched")
+ }
+ q := s.repo.DB.WithContext(ctx).Where(statusFilter, statusArgs...)
+ if len(libraryIDs) > 0 {
+ q = q.Where("library_id IN ?", libraryIDs)
+ }
+ if err := q.
+ Order("CASE WHEN COALESCE(season_num, 0) > 0 OR COALESCE(episode_num, 0) > 0 THEN 1 ELSE 0 END").
+ Order("id ASC").
+ Find(&rows).Error; err != nil {
+ return nil, err
+ }
+ return rows, nil
+}
+
+func (s *ScraperService) notifyScrapeFailed(m model.Media, err error) {
+ if s == nil || s.notify == nil || err == nil {
+ return
+ }
+ body := strings.TrimSpace(m.Title)
+ if body == "" {
+ body = m.Path
+ }
+ body = "媒体:" + body + "\n错误:" + err.Error()
+ go func() {
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ s.notify.Broadcast(ctx, "MediaStationGo 刮削失败", body, EventScrapeFailed)
+ }()
+}
+
+func (s *ScraperService) scrapeDelay(ctx context.Context) time.Duration {
+ minMS := s.scrapeDelaySetting(ctx, "scrape.delay_min_ms", defaultScrapeDelayMinMS)
+ maxMS := s.scrapeDelaySetting(ctx, "scrape.delay_max_ms", defaultScrapeDelayMaxMS)
+ if minMS < 0 {
+ minMS = 0
+ }
+ if maxMS < 0 {
+ maxMS = 0
+ }
+ if minMS > maxScrapeDelayMS {
+ minMS = maxScrapeDelayMS
+ }
+ if maxMS > maxScrapeDelayMS {
+ maxMS = maxScrapeDelayMS
+ }
+ if maxMS < minMS {
+ maxMS = minMS
+ }
+ if maxMS == 0 {
+ return 0
+ }
+ if maxMS == minMS {
+ return time.Duration(minMS) * time.Millisecond
+ }
+ return time.Duration(minMS+secureRandomIntn(maxMS-minMS+1)) * time.Millisecond
+}
+
+func (s *ScraperService) scrapeDelaySetting(ctx context.Context, key string, fallback int) int {
+ if s == nil || s.repo == nil || s.repo.Setting == nil {
+ return fallback
+ }
+ value, err := s.repo.Setting.Get(ctx, key)
+ if err != nil || strings.TrimSpace(value) == "" {
+ return fallback
+ }
+ return parseIntSettingDefault(strings.TrimSpace(value), fallback)
+}
+
+func (s *ScraperService) mediaIsMatched(ctx context.Context, mediaID string) bool {
+ var status string
+ err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).
+ Select("scrape_status").
+ Where("id = ?", mediaID).
+ Scan(&status).Error
+ return err == nil && status == "matched"
+}
diff --git a/internal/service/scraper_query.go b/internal/service/scraper_query.go
new file mode 100644
index 0000000..995daf5
--- /dev/null
+++ b/internal/service/scraper_query.go
@@ -0,0 +1,396 @@
+package service
+
+import (
+ "path/filepath"
+ "regexp"
+ "strconv"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// yearPattern extracts a 4-digit year (1900-2099).
+var yearPattern = regexp.MustCompile(`(?:^|[^\d])(19\d{2}|20\d{2})(?:[^\d]|$)`)
+
+// noiseTokens are stripped before search.
+var noiseTokens = []string{
+ // 视频规格
+ "1080p", "2160p", "4k", "720p", "480p", "uhd", "ds4k", "fhd",
+ "bd", "bdrip", "brrip", "dvd", "dvdrip", "hdtv", "pdtv", "webdl",
+ "hdrip", "bluray", "blu-ray", "webrip", "web-dl", "web",
+ "x264", "x265", "h264", "h265", "hevc", "avc", "10bit", "8bit", "hi10p", "hi10",
+ "hdr", "hdr10", "sdr", "dts", "ddp", "ddp5", "dd5", "dd2", "eac3", "truehd",
+ "dovi", "atmos", "aac", "ac3", "flac",
+ "remux", "extended", "uncut", "remastered", "repack", "proper", "internal",
+ "limited", "imax", "directors-cut", "directors_cut",
+ "hkfree", "yify", "rarbg", "ettv", "fgt", "tgx", "ctrlhd", "ntb", "flux",
+
+ // 流媒体平台 / 字幕组 / 国家版本(动漫常见)
+ "netflix", "nf", "amzn", "hulu", "disney", "max", "hbo",
+ "linetv", "ourtv", "iqiyi", "youku", "bilibili", "qiyi", "krj",
+ "crunchyroll", "funimation", "anidb", "horriblesubs", "subsplease",
+ "erai-raws", "judas", "asw", "smcat", "leopard-raws", "ohys-raws", "colortv",
+
+ // 中文字幕标记
+ "zm", "zw", "ch", "chs", "cht", "cn", "tc", "sc",
+ "中字", "繁字", "简中", "繁中", "国语", "粤语", "日语",
+
+ // 季数前缀残留 — ParseEpisode 已抽取过
+ "season", "264", "265",
+}
+
+var noiseTokenSet = func() map[string]struct{} {
+ set := make(map[string]struct{}, len(noiseTokens)+1)
+ for _, token := range noiseTokens {
+ set[token] = struct{}{}
+ }
+ set["dl"] = struct{}{}
+ return set
+}()
+
+var releaseBoundaryTokenSet = map[string]struct{}{
+ "1080p": {}, "2160p": {}, "4k": {}, "720p": {}, "480p": {}, "uhd": {}, "fhd": {},
+ "bd": {}, "bdrip": {}, "brrip": {}, "dvd": {}, "dvdrip": {}, "hdtv": {}, "pdtv": {},
+ "webdl": {}, "hdrip": {}, "bluray": {}, "webrip": {}, "web": {}, "remux": {},
+ "x264": {}, "x265": {}, "h264": {}, "h265": {}, "hevc": {}, "avc": {},
+}
+
+// bracketedTag matches "[anything]", "(anything)" or "{anything}" segments.
+var bracketedTag = regexp.MustCompile(`[\[\(\{][^\]\)\}]*[\]\)\}]`)
+var multiWordNoise = []*regexp.Regexp{
+ regexp.MustCompile(`(?i)\bweb[\s._-]*dl\b`),
+ regexp.MustCompile(`(?i)\bblu[\s._-]*ray\b`),
+ regexp.MustCompile(`(?i)\bdirectors[\s._-]*cut\b`),
+ regexp.MustCompile(`(?i)\berai[\s._-]*raws\b`),
+ regexp.MustCompile(`(?i)\bohys[\s._-]*raws\b`),
+}
+
+// CleanQuery converts a filename like "Inception.2010.1080p.BluRay.x264.mkv"
+// into a TMDb-friendly title plus an optional year hint.
+func CleanQuery(raw string) (title string, year int) {
+ base := pathBaseSlash(raw)
+ if base == "" {
+ base = strings.TrimSpace(raw)
+ }
+ name := strings.TrimSuffix(base, filepath.Ext(base))
+ lower := strings.ToLower(name)
+
+ if m := yearPattern.FindStringSubmatch(lower); len(m) >= 2 {
+ if v, err := strconv.Atoi(m[1]); err == nil {
+ year = v
+ lower = strings.ReplaceAll(lower, m[1], " ")
+ }
+ }
+
+ lower = bracketedTag.ReplaceAllString(lower, " ")
+
+ lower = patSEnE.ReplaceAllString(lower, " ")
+ lower = patNxE.ReplaceAllString(lower, " ")
+ lower = patEP.ReplaceAllString(lower, " ")
+ lower = patCN.ReplaceAllString(lower, " ")
+ // 去掉中文季/部标记(如「第二季」「第2部」),避免残留在标题里既污染
+ // 搜索查询又导致整理后的目录名重复季信息。
+ lower = patSeasonOnly.ReplaceAllString(lower, " ")
+ lower = patCNSeason.ReplaceAllString(lower, " ")
+
+ for _, pat := range multiWordNoise {
+ lower = pat.ReplaceAllString(lower, " ")
+ }
+ for _, sep := range []string{".", "_", "-", "[", "]", "(", ")", "×"} {
+ lower = strings.ReplaceAll(lower, sep, " ")
+ }
+ // 拆分后丢掉过短(≤1)且全为 ASCII 数字 / 字母的"碎片",避免
+ // 「2」「0」「v」之类残留干扰 TMDb 搜索。中文字符不算碎片。
+ out := make([]string, 0, 8)
+ seenReleaseBoundary := false
+ for _, w := range strings.Fields(lower) {
+ if _, ok := noiseTokenSet[w]; ok {
+ if _, boundary := releaseBoundaryTokenSet[w]; boundary {
+ seenReleaseBoundary = true
+ }
+ continue
+ }
+ if seenReleaseBoundary && isASCIIWord(w) {
+ continue
+ }
+ if len(w) <= 1 {
+ r := []rune(w)
+ if len(r) == 1 && r[0] < 128 {
+ continue
+ }
+ }
+ out = append(out, w)
+ }
+ title = strings.TrimSpace(strings.Join(out, " "))
+ return title, year
+}
+
+func isASCIIWord(s string) bool {
+ if s == "" {
+ return false
+ }
+ for _, r := range s {
+ if r >= 128 {
+ return false
+ }
+ if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') {
+ continue
+ }
+ return false
+ }
+ return true
+}
+
+func scrapeQueryCandidates(m *model.Media, lib *model.Library) []string {
+ seen := map[string]struct{}{}
+ var out []string
+ add := func(raw string) {
+ cleaned, _ := CleanQuery(raw)
+ if cleaned == "" {
+ cleaned = strings.TrimSpace(raw)
+ }
+ for _, candidate := range titleCandidates(cleaned) {
+ key := strings.ToLower(candidate)
+ if _, ok := seen[key]; ok || candidate == "" {
+ continue
+ }
+ seen[key] = struct{}{}
+ out = append(out, candidate)
+ }
+ }
+ episodic := mediaIsEpisodic(m, lib)
+ if lib != nil && episodic {
+ add(seriesFolderTitle(m.Path, lib.Path))
+ }
+ if lib != nil {
+ add(mediaFolderTitle(m.Path, lib.Path))
+ }
+ add(m.Title)
+ add(m.Path)
+ if len(out) == 0 {
+ base := pathBaseSlash(m.Path)
+ out = append(out, strings.TrimSuffix(base, filepath.Ext(base)))
+ }
+ return out
+}
+
+func mediaFolderTitle(mediaPath, libraryRoot string) string {
+ dir := parentSlashPath(mediaPath)
+ root := comparableLibraryRoot(libraryRoot)
+ for depth := 0; depth < 5 && dir != ""; depth++ {
+ if root != "" && sameSlashPath(dir, root) {
+ return libraryRootTitle(libraryRoot)
+ }
+ base := pathBaseSlash(dir)
+ if base == "" || base == "." {
+ return ""
+ }
+ if isTechnicalMediaFolder(base) || strictSeasonFolderMatched(base) {
+ dir = parentSlashPath(dir)
+ continue
+ }
+ if isGenericMediaCategoryFolder(base) {
+ return ""
+ }
+ title, _ := CleanQuery(base)
+ if title == "" {
+ title = strings.TrimSpace(base)
+ }
+ return strings.TrimSpace(title)
+ }
+ return ""
+}
+
+func isTechnicalMediaFolder(name string) bool {
+ key := strings.ToLower(strings.TrimSpace(name))
+ compact := strings.NewReplacer(" ", "", "_", "", ".", "", "-", "").Replace(key)
+ switch compact {
+ case "bdmv", "stream", "certificate", "videots", "audiots",
+ "subs", "subtitles", "subtitle", "sample", "samples",
+ "extra", "extras", "featurette", "featurettes":
+ return true
+ default:
+ return numberedTechnicalFolder(compact, "disc") ||
+ numberedTechnicalFolder(compact, "disk") ||
+ numberedTechnicalFolder(compact, "cd") ||
+ numberedTechnicalFolder(compact, "dvd") ||
+ numberedTechnicalFolder(compact, "part")
+ }
+}
+
+func numberedTechnicalFolder(value, prefix string) bool {
+ if !strings.HasPrefix(value, prefix) || len(value) == len(prefix) {
+ return false
+ }
+ for _, r := range value[len(prefix):] {
+ if r < '0' || r > '9' {
+ return false
+ }
+ }
+ return true
+}
+
+func cleanSlashPath(value string) string {
+ value = strings.TrimSpace(strings.ReplaceAll(value, "\\", "/"))
+ return strings.TrimRight(value, "/")
+}
+
+func comparableLibraryRoot(libraryRoot string) string {
+ if info, ok := ParseCloudLibraryMount(libraryRoot); ok {
+ if strings.TrimSpace(info.DisplayDir) == "" {
+ return "cloud://" + info.Provider
+ }
+ return "cloud://" + info.Provider + "/" + info.DisplayDir
+ }
+ return cleanSlashPath(libraryRoot)
+}
+
+func sameSlashPath(a, b string) bool {
+ return strings.EqualFold(cleanSlashPath(a), cleanSlashPath(b))
+}
+
+func parentSlashPath(value string) string {
+ value = cleanSlashPath(value)
+ if value == "" {
+ return ""
+ }
+ idx := strings.LastIndex(value, "/")
+ if idx < 0 {
+ return ""
+ }
+ return strings.TrimRight(value[:idx], "/")
+}
+
+func seriesFolderTitle(mediaPath, libraryRoot string) string {
+ dir := parentSlashPath(mediaPath)
+ if strictSeasonFolderMatched(pathBaseSlash(dir)) {
+ dir = parentSlashPath(dir)
+ }
+ if root := comparableLibraryRoot(libraryRoot); root != "" && sameSlashPath(dir, root) {
+ return libraryRootTitle(libraryRoot)
+ }
+ base := pathBaseSlash(dir)
+ if base == "" || base == "." {
+ return ""
+ }
+ if isGenericMediaCategoryFolder(base) || isTechnicalMediaFolder(base) || strictSeasonFolderMatched(base) {
+ return ""
+ }
+ return base
+}
+
+func libraryRootTitle(libraryRoot string) string {
+ base := ""
+ if info, ok := ParseCloudLibraryMount(libraryRoot); ok {
+ base = pathBaseSlash(info.DisplayDir)
+ } else {
+ base = pathBaseSlash(libraryRoot)
+ }
+ if base == "" || base == "." || isGenericMediaCategoryFolder(base) || isTechnicalMediaFolder(base) || strictSeasonFolderMatched(base) {
+ return ""
+ }
+ return base
+}
+
+func isGenericMediaCategoryFolder(name string) bool {
+ key := strings.ToLower(strings.TrimSpace(name))
+ key = strings.Trim(key, `\/`)
+ switch key {
+ case "",
+ "电影", "movies", "movie",
+ "电视剧", "剧集", "tv", "shows", "series",
+ "动漫", "动画", "anime", "bangumi",
+ "国产剧", "国剧", "大陆剧", "国产电视剧",
+ "欧美剧", "欧美电视剧",
+ "日韩剧", "日剧", "韩剧",
+ "华语电影", "国产电影", "大陆电影",
+ "外语电影", "欧美电影", "日韩电影",
+ "动画电影", "动漫电影",
+ "国漫", "国产动漫", "日番", "日漫", "日本动漫", "日本动画",
+ "综艺", "真人秀",
+ "纪录片", "纪录",
+ "儿童", "少儿",
+ "成人", "番号", "9kg",
+ "未分类", "uncategorized":
+ return true
+ default:
+ return false
+ }
+}
+
+func strictSeasonFolder(name string) int {
+ if season, ok := seasonFromDir(name); ok {
+ return season
+ }
+ return 0
+}
+
+func strictSeasonFolderMatched(name string) bool {
+ _, ok := seasonFromDir(name)
+ return ok
+}
+
+func titleCandidates(title string) []string {
+ title = strings.Join(strings.Fields(strings.TrimSpace(title)), " ")
+ if title == "" {
+ return nil
+ }
+ out := make([]string, 0, 2)
+ if cjk := cjkTitleOnly(title); cjk != "" {
+ out = append(out, cjk)
+ if cjk != title {
+ return out
+ }
+ }
+ out = append(out, title)
+ return out
+}
+
+func cjkTitleOnly(title string) string {
+ parts := make([]string, 0, 4)
+ for _, field := range strings.Fields(title) {
+ if containsCJK(field) {
+ parts = append(parts, field)
+ }
+ }
+ return strings.Join(parts, " ")
+}
+
+func containsCJK(s string) bool {
+ for _, r := range s {
+ switch {
+ case r >= '\u3400' && r <= '\u4dbf':
+ return true
+ case r >= '\u4e00' && r <= '\u9fff':
+ return true
+ case r >= '\uf900' && r <= '\ufaff':
+ return true
+ }
+ }
+ return false
+}
+
+func mediaIsEpisodic(m *model.Media, lib *model.Library) bool {
+ if m != nil && (m.SeasonNum > 0 || m.EpisodeNum > 0) {
+ return true
+ }
+ if m != nil {
+ season, episode := ParseEpisode(m.Path)
+ if season > 0 || episode > 0 {
+ return true
+ }
+ }
+ return librarySupportsSeasons(lib)
+}
+
+func librarySupportsSeasons(lib *model.Library) bool {
+ if lib == nil {
+ return false
+ }
+ switch strings.ToLower(strings.TrimSpace(lib.Type)) {
+ case "tv", "anime", "variety", "show", "shows":
+ return true
+ default:
+ return false
+ }
+}
diff --git a/internal/service/scraper_query_test.go b/internal/service/scraper_query_test.go
new file mode 100644
index 0000000..5ded697
--- /dev/null
+++ b/internal/service/scraper_query_test.go
@@ -0,0 +1,427 @@
+package service
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "path/filepath"
+ "strings"
+ "testing"
+
+ "github.com/glebarez/sqlite"
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestCleanQuery(t *testing.T) {
+ cases := []struct {
+ in string
+ wantTitle string
+ wantYear int
+ }{
+ {"Inception.2010.1080p.BluRay.x264.mkv", "inception", 2010},
+ {"The_Matrix_(1999).1080p.WEB-DL.H265.mp4", "the matrix", 1999},
+ {"interstellar.2014.4k.hdr.dts.atmos.mkv", "interstellar", 2014},
+ {"My Movie 2022 [HDR] (1080p) [TGx].mp4", "my movie", 2022},
+ {"NoYearOrTags.mkv", "noyearortags", 0},
+ {"亏成首富从游戏开始 The Richest in Game - S01E11 - 4K.mp4", "亏成首富从游戏开始 the richest in game", 0},
+ {"紫川.2024.S02E24.第24集.2160p.WEB-DL.H.265-ColorTV.mkv", "紫川", 2024},
+ {"紫川 (2024) {tmdb-247590}", "紫川", 2024},
+ }
+ for _, tc := range cases {
+ t.Run(tc.in, func(t *testing.T) {
+ gotTitle, gotYear := CleanQuery(tc.in)
+ if gotTitle != tc.wantTitle || gotYear != tc.wantYear {
+ t.Errorf("CleanQuery(%q) = (%q, %d), want (%q, %d)",
+ tc.in, gotTitle, gotYear, tc.wantTitle, tc.wantYear)
+ }
+ })
+ }
+}
+
+func TestExternalIDHintsFromText(t *testing.T) {
+ hints := externalIDHintsFromText("国漫/折腰 (2025) {tmdb 296753}/Season 1/折腰.S01E01.mkv")
+ if hints.TMDbID != 296753 {
+ t.Fatalf("tmdb hint = %d, want 296753", hints.TMDbID)
+ }
+ hints = externalIDHintsFromText("Movie (2026) {tmdb-1630433} [douban=3622222] {bgm 456789} {tvdb:12345}")
+ if hints.TMDbID != 1630433 || hints.DoubanID != "3622222" || hints.BangumiID != 456789 || hints.TheTVDBID != "12345" {
+ t.Fatalf("external hints not parsed: %+v", hints)
+ }
+}
+
+func TestPathHintMetadataDoesNotMarkMediaMatched(t *testing.T) {
+ meta, hints := pathHintMetadata("cloud://openlist/国漫/折腰 (2025) {tmdb 296753}/Season 1/折腰.S01E01.mkv", true)
+ if meta == nil || hints.TMDbID != 296753 || meta.TMDbID != 296753 || meta.Title != "折腰" || meta.Year != 2025 {
+ t.Fatalf("path hint metadata = %+v hints=%+v", meta, hints)
+ }
+ media := &model.Media{Title: "折腰", ScrapeStatus: "pending"}
+ applyLocalMetadata(media, meta)
+ if media.ScrapeStatus != "pending" {
+ t.Fatalf("path hints alone must not mark media matched, got %q", media.ScrapeStatus)
+ }
+}
+
+func TestEnrichOneCloudPathHintOverridesStaleTMDbID(t *testing.T) {
+ var requested []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ requested = append(requested, r.URL.Path)
+ w.Header().Set("Content-Type", "application/json")
+ switch r.URL.Path {
+ case "/tv/296753":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": 296753,
+ "name": "折腰",
+ "overview": "正确的剧集条目",
+ "poster_path": "/zheyao.jpg",
+ "first_air_date": "2025-05-13",
+ "origin_country": []string{"CN"},
+ })
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ cfg.Secrets.TMDbImageProxy = upstream.URL + "/images"
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "OpenList · 国产剧", Path: "cloud://openlist/国产剧", Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "折腰",
+ Path: "cloud://openlist/国产剧/折腰 (2025) {tmdb-296753}/Season 1/折腰.S01E01.mkv",
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ TMDbID: 220269,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ if err := scraper.EnrichOne(t.Context(), &media); err != nil {
+ t.Fatal(err)
+ }
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.ScrapeStatus != "matched" || got.TMDbID != 296753 || got.Title != "折腰" || got.PosterURL == "" {
+ t.Fatalf("path hint was not authoritative: status=%q tmdb=%d title=%q poster=%q", got.ScrapeStatus, got.TMDbID, got.Title, got.PosterURL)
+ }
+ for _, path := range requested {
+ if path == "/tv/220269" || path == "/movie/220269" {
+ t.Fatalf("scraper queried stale tmdb id; requests=%v", requested)
+ }
+ }
+}
+
+func TestEnrichOneUsesLocalPathExternalIDHints(t *testing.T) {
+ var requested []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ requested = append(requested, r.URL.Path)
+ w.Header().Set("Content-Type", "application/json")
+ switch r.URL.Path {
+ case "/movie/27205":
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": 27205,
+ "title": "Inception",
+ "overview": "A thief enters dreams.",
+ "poster_path": "/inception.jpg",
+ "release_date": "2010-07-16",
+ "vote_average": 8.4,
+ "original_title": "Inception",
+ })
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ root := t.TempDir()
+ mediaPath := filepath.Join(root, "错误标题 (2010) {tmdb-27205}", "bad-file-name.mkv")
+ lib := model.Library{Name: "电影", Path: root, Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "bad local title",
+ Path: mediaPath,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ if err := scraper.EnrichOne(t.Context(), &media); err != nil {
+ t.Fatal(err)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.ScrapeStatus != "matched" || got.TMDbID != 27205 || got.Title != "Inception" {
+ t.Fatalf("local path tmdb hint was not used: status=%q tmdb=%d title=%q requests=%v", got.ScrapeStatus, got.TMDbID, got.Title, requested)
+ }
+ if len(requested) == 0 || requested[0] != "/movie/27205" {
+ t.Fatalf("scraper should query by hinted tmdb id first, requests=%v", requested)
+ }
+}
+
+func TestScrapeQueryCandidatesPreferSeriesFolderAndCJKTitle(t *testing.T) {
+ lib := &model.Library{
+ Path: `F:\downloads\国产剧`,
+ Type: "movie",
+ }
+ media := &model.Media{
+ Title: "亏成首富从游戏开始 the ri est in game",
+ Path: `F:\downloads\国产剧\亏成首富从游戏开始 The Richest in Game\Season 01\亏成首富从游戏开始 The Richest in Game - S01E11 - 4K.mp4`,
+ SeasonNum: 1,
+ EpisodeNum: 11,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] != "亏成首富从游戏开始" {
+ t.Fatalf("first query candidate = %q, want Chinese series title", got[0])
+ }
+ for _, candidate := range got {
+ if strings.Contains(candidate, "ri est") {
+ t.Fatalf("query candidate kept substring-stripped title: %#v", got)
+ }
+ }
+}
+
+func TestScrapeQueryCandidatesUseCloudSeriesFolder(t *testing.T) {
+ lib := &model.Library{
+ Path: "cloud://openlist/国产剧",
+ Type: "movie",
+ }
+ media := &model.Media{
+ Title: "折腰 S01E01",
+ Path: "cloud://openlist/国产剧/折腰 (2025)/Season 1/折腰.S01E01.mkv",
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] != "折腰" {
+ t.Fatalf("first query candidate = %q, want cloud series folder title; all candidates=%#v", got[0], got)
+ }
+}
+
+func TestScrapeQueryCandidatesUseSeriesLibraryRootWhenMountedAtShowFolder(t *testing.T) {
+ lib := &model.Library{
+ Path: `/downloads/国产剧/折腰 (2025)`,
+ Type: "tv",
+ }
+ media := &model.Media{
+ Title: "第 1 集",
+ Path: `/downloads/国产剧/折腰 (2025)/Season 01/第01集.mkv`,
+ SeasonNum: 1,
+ EpisodeNum: 1,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] != "折腰" {
+ t.Fatalf("first query candidate = %q, want library root show title; all candidates=%#v", got[0], got)
+ }
+}
+
+func TestMediaIsEpisodicUsesEpisodePatternInPath(t *testing.T) {
+ lib := &model.Library{
+ Path: `/media/movies`,
+ Type: "movie",
+ }
+ media := &model.Media{
+ Title: "折腰 S01E01",
+ Path: `/media/movies/折腰/Season 01/折腰.S01E01.mkv`,
+ }
+
+ if !mediaIsEpisodic(media, lib) {
+ t.Fatal("media with an SxxEyy path should be treated as episodic even in a movie library")
+ }
+}
+
+func TestScrapeQueryCandidatesSkipCategoryFolderAsSeriesTitle(t *testing.T) {
+ lib := &model.Library{
+ Path: `/downloads`,
+ Type: "tv",
+ }
+ media := &model.Media{
+ Title: "Ashes To Crown",
+ Path: `/downloads/国产剧/Ashes.to.Crown.S01E06.1080p.WEB-DL.mkv`,
+ SeasonNum: 1,
+ EpisodeNum: 6,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] == "国产剧" {
+ t.Fatalf("first query candidate = %q, category folders must not be used as title candidates: %#v", got[0], got)
+ }
+ if !strings.EqualFold(got[0], "Ashes To Crown") {
+ t.Fatalf("first query candidate = %q, want release title; all candidates=%#v", got[0], got)
+ }
+}
+
+func TestScrapeQueryCandidatesUseMovieFolderForGenericFilename(t *testing.T) {
+ lib := &model.Library{
+ Path: `/media/movies`,
+ Type: "movie",
+ }
+ media := &model.Media{
+ Title: "00000",
+ Path: `/media/movies/Inception (2010)/BDMV/STREAM/00000.m2ts`,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] != "inception" {
+ t.Fatalf("first query candidate = %q, want movie folder title; all candidates=%#v", got[0], got)
+ }
+ for _, candidate := range got {
+ switch strings.ToLower(candidate) {
+ case "bdmv", "stream":
+ t.Fatalf("query candidates kept technical filename/folder: %#v", got)
+ }
+ }
+}
+
+func TestScrapeQueryCandidatesUseMovieLibraryRootWhenMountedAtMovieFolder(t *testing.T) {
+ lib := &model.Library{
+ Path: `/media/movies/Inception (2010)`,
+ Type: "movie",
+ }
+ media := &model.Media{
+ Title: "00000",
+ Path: `/media/movies/Inception (2010)/BDMV/STREAM/00000.m2ts`,
+ }
+
+ got := scrapeQueryCandidates(media, lib)
+ if len(got) == 0 {
+ t.Fatal("scrapeQueryCandidates returned no candidates")
+ }
+ if got[0] != "inception" {
+ t.Fatalf("first query candidate = %q, want movie library root title; all candidates=%#v", got[0], got)
+ }
+}
+
+func TestEnrichOneUsesMovieFolderWhenFilenameIsGeneric(t *testing.T) {
+ var queries []string
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ queries = append(queries, r.URL.Query().Get("query"))
+ w.Header().Set("Content-Type", "application/json")
+ if r.URL.Path != "/search/movie" {
+ http.NotFound(w, r)
+ return
+ }
+ if r.URL.Query().Get("query") != "inception" {
+ _ = json.NewEncoder(w).Encode(map[string]any{"results": []any{}})
+ return
+ }
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "results": []map[string]any{{
+ "id": 27205,
+ "title": "Inception",
+ "overview": "A thief enters dreams.",
+ "poster_path": "/inception.jpg",
+ "release_date": "2010-07-16",
+ "vote_average": 8.4,
+ "original_title": "Inception",
+ }},
+ })
+ }))
+ defer upstream.Close()
+
+ db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
+ t.Fatal(err)
+ }
+ repos := repository.New(db)
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test-key"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ log := zap.NewNop()
+ scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
+
+ lib := model.Library{Name: "Movies", Path: `/media/movies`, Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "00000",
+ Path: `/media/movies/Inception (2010)/BDMV/STREAM/00000.m2ts`,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ if err := scraper.EnrichOne(t.Context(), &media); err != nil {
+ t.Fatal(err)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.ScrapeStatus != "matched" || got.TMDbID != 27205 || got.Title != "Inception" {
+ t.Fatalf("generic filename scrape did not use folder title: status=%q tmdb=%d title=%q queries=%v", got.ScrapeStatus, got.TMDbID, got.Title, queries)
+ }
+ if len(queries) == 0 || queries[0] != "inception" {
+ t.Fatalf("first tmdb query = %q, want folder title; all queries=%v", firstQuery(queries), queries)
+ }
+}
diff --git a/internal/service/scraper_test.go b/internal/service/scraper_test.go
index 20bea78..4ba802b 100644
--- a/internal/service/scraper_test.go
+++ b/internal/service/scraper_test.go
@@ -8,7 +8,6 @@ import (
"path/filepath"
"strings"
"testing"
- "time"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
@@ -19,327 +18,6 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
-func TestCleanQuery(t *testing.T) {
- cases := []struct {
- in string
- wantTitle string
- wantYear int
- }{
- {"Inception.2010.1080p.BluRay.x264.mkv", "inception", 2010},
- {"The_Matrix_(1999).1080p.WEB-DL.H265.mp4", "the matrix", 1999},
- {"interstellar.2014.4k.hdr.dts.atmos.mkv", "interstellar", 2014},
- {"My Movie 2022 [HDR] (1080p) [TGx].mp4", "my movie", 2022},
- {"NoYearOrTags.mkv", "noyearortags", 0},
- {"亏成首富从游戏开始 The Richest in Game - S01E11 - 4K.mp4", "亏成首富从游戏开始 the richest in game", 0},
- {"紫川.2024.S02E24.第24集.2160p.WEB-DL.H.265-ColorTV.mkv", "紫川", 2024},
- {"紫川 (2024) {tmdb-247590}", "紫川", 2024},
- }
- for _, tc := range cases {
- t.Run(tc.in, func(t *testing.T) {
- gotTitle, gotYear := CleanQuery(tc.in)
- if gotTitle != tc.wantTitle || gotYear != tc.wantYear {
- t.Errorf("CleanQuery(%q) = (%q, %d), want (%q, %d)",
- tc.in, gotTitle, gotYear, tc.wantTitle, tc.wantYear)
- }
- })
- }
-}
-
-func TestExternalIDHintsFromText(t *testing.T) {
- hints := externalIDHintsFromText("国漫/折腰 (2025) {tmdb 296753}/Season 1/折腰.S01E01.mkv")
- if hints.TMDbID != 296753 {
- t.Fatalf("tmdb hint = %d, want 296753", hints.TMDbID)
- }
- hints = externalIDHintsFromText("Movie (2026) {tmdb-1630433} [douban=3622222] {bgm 456789} {tvdb:12345}")
- if hints.TMDbID != 1630433 || hints.DoubanID != "3622222" || hints.BangumiID != 456789 || hints.TheTVDBID != "12345" {
- t.Fatalf("external hints not parsed: %+v", hints)
- }
-}
-
-func TestPathHintMetadataDoesNotMarkMediaMatched(t *testing.T) {
- meta, hints := pathHintMetadata("cloud://openlist/国漫/折腰 (2025) {tmdb 296753}/Season 1/折腰.S01E01.mkv", true)
- if meta == nil || hints.TMDbID != 296753 || meta.TMDbID != 296753 || meta.Title != "折腰" || meta.Year != 2025 {
- t.Fatalf("path hint metadata = %+v hints=%+v", meta, hints)
- }
- media := &model.Media{Title: "折腰", ScrapeStatus: "pending"}
- applyLocalMetadata(media, meta)
- if media.ScrapeStatus != "pending" {
- t.Fatalf("path hints alone must not mark media matched, got %q", media.ScrapeStatus)
- }
-}
-
-func TestEnrichOneCloudPathHintOverridesStaleTMDbID(t *testing.T) {
- var requested []string
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- requested = append(requested, r.URL.Path)
- w.Header().Set("Content-Type", "application/json")
- switch r.URL.Path {
- case "/tv/296753":
- _ = json.NewEncoder(w).Encode(map[string]any{
- "id": 296753,
- "name": "折腰",
- "overview": "正确的剧集条目",
- "poster_path": "/zheyao.jpg",
- "first_air_date": "2025-05-13",
- "origin_country": []string{"CN"},
- })
- default:
- http.NotFound(w, r)
- }
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- cfg := &config.Config{}
- cfg.Secrets.TMDbAPIKey = "test-key"
- cfg.Secrets.TMDbAPIProxy = upstream.URL
- cfg.Secrets.TMDbImageProxy = upstream.URL + "/images"
- log := zap.NewNop()
- scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
-
- lib := model.Library{Name: "OpenList · 国产剧", Path: "cloud://openlist/国产剧", Type: "tv", Enabled: true}
- if err := repos.DB.Create(&lib).Error; err != nil {
- t.Fatal(err)
- }
- media := model.Media{
- LibraryID: lib.ID,
- Title: "折腰",
- Path: "cloud://openlist/国产剧/折腰 (2025) {tmdb-296753}/Season 1/折腰.S01E01.mkv",
- SeasonNum: 1,
- EpisodeNum: 1,
- TMDbID: 220269,
- ScrapeStatus: "pending",
- }
- if err := repos.DB.Create(&media).Error; err != nil {
- t.Fatal(err)
- }
-
- if err := scraper.EnrichOne(t.Context(), &media); err != nil {
- t.Fatal(err)
- }
- var got model.Media
- if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
- t.Fatal(err)
- }
- if got.ScrapeStatus != "matched" || got.TMDbID != 296753 || got.Title != "折腰" || got.PosterURL == "" {
- t.Fatalf("path hint was not authoritative: status=%q tmdb=%d title=%q poster=%q", got.ScrapeStatus, got.TMDbID, got.Title, got.PosterURL)
- }
- for _, path := range requested {
- if path == "/tv/220269" || path == "/movie/220269" {
- t.Fatalf("scraper queried stale tmdb id; requests=%v", requested)
- }
- }
-}
-
-func TestManualRequestMatchFallsBackToCandidatePayload(t *testing.T) {
- scraper := &ScraperService{}
- match, err := scraper.manualRequestMatch(t.Context(), ManualScrapeRequest{
- Source: "douban",
- Title: "手动选择的电影",
- DoubanID: "1234567",
- Year: 2026,
- })
- if err != nil {
- t.Fatal(err)
- }
- if match.Title != "手动选择的电影" || match.DoubanID != "1234567" || match.Year != 2026 {
- t.Fatalf("fallback match = %#v", match)
- }
-}
-
-func TestManualSearchReturnsTMDbCandidatePage(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("Content-Type", "application/json")
- if r.URL.Path != "/search/movie" {
- http.NotFound(w, r)
- return
- }
- _ = json.NewEncoder(w).Encode(map[string]any{
- "results": []map[string]any{
- {
- "id": 101,
- "title": "错误的同名电影",
- "poster_path": "/wrong.jpg",
- "release_date": "2021-01-01",
- "vote_average": 5.1,
- "genre_ids": []int{18},
- "backdrop_path": "/wrong-backdrop.jpg",
- },
- {
- "id": 202,
- "title": "正确的同名电影",
- "poster_path": "/right.jpg",
- "release_date": "2021-08-01",
- "vote_average": 8.2,
- "genre_ids": []int{28},
- "backdrop_path": "/right-backdrop.jpg",
- },
- },
- })
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- cfg := &config.Config{}
- cfg.Secrets.TMDbAPIKey = "test-key"
- cfg.Secrets.TMDbAPIProxy = upstream.URL
- log := zap.NewNop()
- scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
-
- lib := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true}
- if err := repos.DB.Create(&lib).Error; err != nil {
- t.Fatal(err)
- }
- media := model.Media{LibraryID: lib.ID, Title: "同名电影", Path: "/media/movie/同名电影.mkv"}
- if err := repos.DB.Create(&media).Error; err != nil {
- t.Fatal(err)
- }
-
- results, err := scraper.ManualSearch(t.Context(), &media, "同名电影", "tmdb", "movie")
- if err != nil {
- t.Fatal(err)
- }
- if len(results) != 2 || results[0].TMDbID != 101 || results[1].TMDbID != 202 {
- t.Fatalf("manual TMDb candidates = %#v", results)
- }
-}
-
-func TestManualSearchIncludesAdultProvider(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- switch r.URL.Path {
- case "/search":
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- _, _ = w.Write([]byte(`SSIS-001 手动候选`))
- case "/v/ssis001":
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- _, _ = w.Write([]byte(`SSIS-001 手动成人标题
`))
- default:
- http.NotFound(w, r)
- }
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}, &model.APIConfig{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- apiConfig := NewAPIConfigService(zap.NewNop(), repos, NewCryptoService("", zap.NewNop()))
- baseURL := upstream.URL
- if _, err := apiConfig.Update(t.Context(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
- t.Fatal(err)
- }
- log := zap.NewNop()
- scraper := NewScraperService(&config.Config{}, log, repos, nil, nil, nil, nil, NewHub(log), NewAdultProvider(log, apiConfig))
-
- lib := model.Library{Name: "成人", Path: "/media/adult", Type: "movie", Enabled: true}
- if err := repos.DB.Create(&lib).Error; err != nil {
- t.Fatal(err)
- }
- media := model.Media{LibraryID: lib.ID, Title: "SSIS-001", OriginalName: "SSIS-001", Path: "/media/adult/SSIS-001.mkv"}
- if err := repos.DB.Create(&media).Error; err != nil {
- t.Fatal(err)
- }
-
- results, err := scraper.ManualSearch(t.Context(), &media, "SSIS-001", "adult", "adult")
- if err != nil {
- t.Fatal(err)
- }
- if len(results) != 1 || results[0].Source != "adult" || results[0].MediaType != "adult" || !results[0].NSFW || results[0].OriginalName != "SSIS-001" {
- t.Fatalf("manual adult candidates = %#v", results)
- }
-}
-
-func TestApplyManualMatchSavesSelectedCloudMatchWhenDetailsSlow(t *testing.T) {
- oldTimeout := tmdbDetailsTimeout
- tmdbDetailsTimeout = 20 * time.Millisecond
- defer func() { tmdbDetailsTimeout = oldTimeout }()
-
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/movie/77" {
- http.NotFound(w, r)
- return
- }
- select {
- case <-r.Context().Done():
- return
- case <-time.After(time.Second):
- _ = json.NewEncoder(w).Encode(map[string]any{
- "id": 77,
- "title": "Slow Details",
- })
- }
- }))
- defer upstream.Close()
-
- db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- cfg := &config.Config{}
- cfg.Secrets.TMDbAPIKey = "test-key"
- cfg.Secrets.TMDbAPIProxy = upstream.URL
- log := zap.NewNop()
- scraper := NewScraperService(cfg, log, repos, NewTMDbProvider(cfg, log, nil), nil, nil, nil, NewHub(log))
-
- lib := model.Library{Name: "OpenList · Movies", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
- if err := repos.DB.Create(&lib).Error; err != nil {
- t.Fatal(err)
- }
- media := model.Media{
- LibraryID: lib.ID,
- Title: "bad cloud title",
- Path: "cloud://openlist/Movies/Bad.Title.2026.mkv",
- ScrapeStatus: "pending",
- }
- if err := repos.DB.Create(&media).Error; err != nil {
- t.Fatal(err)
- }
-
- start := time.Now()
- if _, err := scraper.ApplyManualMatch(t.Context(), media.ID, ManualScrapeRequest{
- Source: "manual",
- MediaType: "movie",
- Title: "Correct Cloud Movie",
- TMDbID: 77,
- Year: 2026,
- }); err != nil {
- t.Fatal(err)
- }
- if elapsed := time.Since(start); elapsed > 500*time.Millisecond {
- t.Fatalf("manual apply waited for optional details: %s", elapsed)
- }
-
- var got model.Media
- if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
- t.Fatal(err)
- }
- if got.Title != "Correct Cloud Movie" || got.ScrapeStatus != "matched" || got.TMDbID != 77 {
- t.Fatalf("manual cloud match was not saved: title=%q status=%q tmdb=%d", got.Title, got.ScrapeStatus, got.TMDbID)
- }
-}
-
func TestEnrichOneUsesExistingTMDbIDForCloudMedia(t *testing.T) {
scraper, repos, closeServer := newTestScraper(t)
defer closeServer()
@@ -373,56 +51,6 @@ func TestEnrichOneUsesExistingTMDbIDForCloudMedia(t *testing.T) {
}
}
-func TestScrapeQueryCandidatesPreferSeriesFolderAndCJKTitle(t *testing.T) {
- lib := &model.Library{
- Path: `F:\downloads\国产剧`,
- Type: "movie",
- }
- media := &model.Media{
- Title: "亏成首富从游戏开始 the ri est in game",
- Path: `F:\downloads\国产剧\亏成首富从游戏开始 The Richest in Game\Season 01\亏成首富从游戏开始 The Richest in Game - S01E11 - 4K.mp4`,
- SeasonNum: 1,
- EpisodeNum: 11,
- }
-
- got := scrapeQueryCandidates(media, lib)
- if len(got) == 0 {
- t.Fatal("scrapeQueryCandidates returned no candidates")
- }
- if got[0] != "亏成首富从游戏开始" {
- t.Fatalf("first query candidate = %q, want Chinese series title", got[0])
- }
- for _, candidate := range got {
- if strings.Contains(candidate, "ri est") {
- t.Fatalf("query candidate kept substring-stripped title: %#v", got)
- }
- }
-}
-
-func TestScrapeQueryCandidatesSkipCategoryFolderAsSeriesTitle(t *testing.T) {
- lib := &model.Library{
- Path: `/downloads`,
- Type: "tv",
- }
- media := &model.Media{
- Title: "Ashes To Crown",
- Path: `/downloads/国产剧/Ashes.to.Crown.S01E06.1080p.WEB-DL.mkv`,
- SeasonNum: 1,
- EpisodeNum: 6,
- }
-
- got := scrapeQueryCandidates(media, lib)
- if len(got) == 0 {
- t.Fatal("scrapeQueryCandidates returned no candidates")
- }
- if got[0] == "国产剧" {
- t.Fatalf("first query candidate = %q, category folders must not be used as title candidates: %#v", got[0], got)
- }
- if !strings.EqualFold(got[0], "Ashes To Crown") {
- t.Fatalf("first query candidate = %q, want release title; all candidates=%#v", got[0], got)
- }
-}
-
func TestEnrichOneWritesTMDbIDColumn(t *testing.T) {
scraper, repos, closeServer := newTestScraper(t)
defer closeServer()
@@ -460,6 +88,39 @@ func TestEnrichOneWritesTMDbIDColumn(t *testing.T) {
}
}
+func TestEnrichOneTreatsEpisodicMediaInMovieLibraryAsTV(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ lib := model.Library{Name: "混合库", Path: t.TempDir(), Type: "movie", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "间谍过家家 S02E01",
+ Path: filepath.Join(lib.Path, "间谍过家家", "Season 02", "间谍过家家 - S02E01.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ if err := scraper.EnrichOne(t.Context(), &media); err != nil {
+ t.Fatal(err)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.ScrapeStatus != "matched" || got.TMDbID != 12345 {
+ t.Fatalf("episodic media in movie library should use tv scrape: status=%q tmdb=%d", got.ScrapeStatus, got.TMDbID)
+ }
+}
+
func TestEnrichOneWritesTMDbEpisodeMetadata(t *testing.T) {
scraper, repos, closeServer := newTestScraper(t)
defer closeServer()
@@ -469,12 +130,16 @@ func TestEnrichOneWritesTMDbEpisodeMetadata(t *testing.T) {
t.Fatal(err)
}
mediaPath := filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv")
+ existingPoster := "https://image.tmdb.org/t/p/w500/existing-poster.jpg"
+ existingBackdrop := "https://image.tmdb.org/t/p/w1280/existing-backdrop.jpg"
media := model.Media{
LibraryID: lib.ID,
Title: "间谍过家家",
Path: mediaPath,
SeasonNum: 2,
EpisodeNum: 1,
+ PosterURL: existingPoster,
+ BackdropURL: existingBackdrop,
ScrapeStatus: "pending",
}
if err := repos.DB.Create(&media).Error; err != nil {
@@ -499,6 +164,9 @@ func TestEnrichOneWritesTMDbEpisodeMetadata(t *testing.T) {
if got.Rating < 9.09 || got.Rating > 9.11 {
t.Fatalf("episode rating = %v, want 9.1", got.Rating)
}
+ if got.EpisodeTitle != "任务代号: 猫" {
+ t.Fatalf("episode_title should store per-episode name, got %q", got.EpisodeTitle)
+ }
// original_name 必须保持「整剧原名」,绝不能被单集名(任务代号: 猫)覆盖,
// 否则同剧每集 original_name 不同会导致合集被拆成多集无法合并。
if got.OriginalName != "SPY×FAMILY" {
@@ -506,6 +174,108 @@ func TestEnrichOneWritesTMDbEpisodeMetadata(t *testing.T) {
}
}
+func TestEnrichOneSkipsTMDbEpisodeStillWhenDisabled(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ mediaPath := filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv")
+ existingPoster := "https://image.tmdb.org/t/p/w500/existing-poster.jpg"
+ existingBackdrop := "https://image.tmdb.org/t/p/w1280/existing-backdrop.jpg"
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: mediaPath,
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ PosterURL: existingPoster,
+ BackdropURL: existingBackdrop,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ episodeArtwork := false
+ if err := scraper.EnrichOneWithOptions(t.Context(), &media, ScrapeOptions{EpisodeArtwork: &episodeArtwork}); err != nil {
+ t.Fatal(err)
+ }
+
+ var got model.Media
+ if err := repos.DB.First(&got, "id = ?", media.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.Overview != "单集剧情" || got.DurationSec != 24*60 {
+ t.Fatalf("episode metadata should still be saved: overview=%q duration=%d", got.Overview, got.DurationSec)
+ }
+ if got.Rating < 9.09 || got.Rating > 9.11 {
+ t.Fatalf("episode rating = %v, want 9.1", got.Rating)
+ }
+ if strings.HasSuffix(got.BackdropURL, "/images/w500/still.jpg") {
+ t.Fatalf("episode still should not be saved when disabled: backdrop=%q", got.BackdropURL)
+ }
+ if !strings.HasSuffix(got.PosterURL, "/images/w500/poster.jpg") {
+ t.Fatalf("series poster should still be saved when episode artwork is disabled: got %q", got.PosterURL)
+ }
+ if !strings.HasSuffix(got.BackdropURL, "/images/w1280/backdrop.jpg") {
+ t.Fatalf("series backdrop should still be saved when episode artwork is disabled: got %q", got.BackdropURL)
+ }
+ if got.PosterURL == existingPoster || got.BackdropURL == existingBackdrop {
+ t.Fatalf("main artwork should be refreshed while episode still is skipped: poster=%q backdrop=%q", got.PosterURL, got.BackdropURL)
+ }
+}
+
+func TestApplyManualMatchSkipsTMDbEpisodeStillWhenDisabled(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "待匹配",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ ScrapeStatus: "pending",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ episodeArtwork := false
+ got, err := scraper.ApplyManualMatch(t.Context(), media.ID, ManualScrapeRequest{
+ Source: "tmdb",
+ MediaType: "tv",
+ Title: "间谍过家家",
+ TMDbID: 12345,
+ EpisodeArtwork: &episodeArtwork,
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if got == nil {
+ t.Fatal("manual match returned nil media")
+ }
+ if got.Overview != "单集剧情" || got.DurationSec != 24*60 {
+ t.Fatalf("episode metadata should still be saved: overview=%q duration=%d", got.Overview, got.DurationSec)
+ }
+ if strings.HasSuffix(got.BackdropURL, "/images/w500/still.jpg") {
+ t.Fatalf("manual episode still should not be saved when disabled: backdrop=%q", got.BackdropURL)
+ }
+ if !strings.HasSuffix(got.PosterURL, "/images/w500/poster.jpg") {
+ t.Fatalf("series poster should still be saved when manual episode artwork is disabled: got %q", got.PosterURL)
+ }
+ if !strings.HasSuffix(got.BackdropURL, "/images/w1280/backdrop.jpg") {
+ t.Fatalf("series backdrop should still be saved when manual episode artwork is disabled: got %q", got.BackdropURL)
+ }
+}
+
func TestEnrichOneRejectsWrongYearMatchFromSeriesFolder(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
@@ -634,6 +404,9 @@ func TestEnrichOnePrefersLocalMetadataWithoutProvider(t *testing.T) {
if got.SeasonNum != 2 || got.EpisodeNum != 12 || got.Overview != "本地剧情" {
t.Fatalf("unexpected local episode data: s=%d e=%d overview=%q", got.SeasonNum, got.EpisodeNum, got.Overview)
}
+ if got.EpisodeTitle != "企鹅公园" {
+ t.Fatalf("episode_title = %q, want local episode title", got.EpisodeTitle)
+ }
}
func TestManualEnrichLibraryRetriesNoMatchAndCountsRealMatches(t *testing.T) {
@@ -664,35 +437,129 @@ func TestManualEnrichLibraryRetriesNoMatchAndCountsRealMatches(t *testing.T) {
}
}
-func TestScrapeDelayUsesSettings(t *testing.T) {
+func TestManualEnrichLibraryCanRefreshAlreadyMatchedRows(t *testing.T) {
scraper, repos, closeServer := newTestScraper(t)
defer closeServer()
- if err := repos.DB.AutoMigrate(&model.Setting{}); err != nil {
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E02.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 2,
+ ScrapeStatus: "matched",
+ }
+ if err := repos.DB.Create(&media).Error; err != nil {
t.Fatal(err)
}
- if got := scraper.scrapeDelay(t.Context()); got < 250*time.Millisecond || got > 500*time.Millisecond {
- t.Fatalf("default scrapeDelay = %s, want 250-500ms", got)
+ defaultResult, err := scraper.EnrichLibraryDetailedWithOptions(t.Context(), lib.ID, ScrapeOptions{RetryNoMatch: true})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if defaultResult.Processed != 0 || defaultResult.Candidates != 0 {
+ t.Fatalf("default manual scrape result=%+v, want matched rows skipped without IncludeMatched", defaultResult)
}
- if err := repos.Setting.Set(t.Context(), "scrape.delay_min_ms", "0"); err != nil {
+ refreshResult, err := scraper.EnrichLibraryDetailedWithOptions(t.Context(), lib.ID, ScrapeOptions{
+ RetryNoMatch: true,
+ IncludeMatched: true,
+ })
+ if err != nil {
t.Fatal(err)
}
- if err := repos.Setting.Set(t.Context(), "scrape.delay_max_ms", "0"); err != nil {
+ if refreshResult.Processed != 1 || refreshResult.Matched != 1 || refreshResult.Candidates != 1 {
+ t.Fatalf("refresh result=%+v, want already matched row reprocessed", refreshResult)
+ }
+}
+
+func TestScrapeCandidateRowsPrioritizeLibraryArtworkBeforeEpisodes(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ lib := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ if err := repos.DB.Create(&lib).Error; err != nil {
t.Fatal(err)
}
- if got := scraper.scrapeDelay(t.Context()); got != 0 {
- t.Fatalf("disabled scrapeDelay = %s, want 0", got)
+ rows := []model.Media{
+ {
+ Base: model.Base{ID: "001-episode"},
+ LibraryID: lib.ID,
+ Title: "间谍过家家 第 1 集",
+ Path: filepath.Join(lib.Path, "间谍过家家 - S02E01.mkv"),
+ SeasonNum: 2,
+ EpisodeNum: 1,
+ ScrapeStatus: "pending",
+ },
+ {
+ Base: model.Base{ID: "999-series"},
+ LibraryID: lib.ID,
+ Title: "间谍过家家",
+ Path: filepath.Join(lib.Path, "间谍过家家.mkv"),
+ ScrapeStatus: "pending",
+ },
+ }
+ if err := repos.DB.Create(&rows).Error; err != nil {
+ t.Fatal(err)
}
- if err := repos.Setting.Set(t.Context(), "scrape.delay_min_ms", "800"); err != nil {
+ got, err := scraper.scrapeCandidateRows(t.Context(), lib.ID, ScrapeOptions{})
+ if err != nil {
t.Fatal(err)
}
- if err := repos.Setting.Set(t.Context(), "scrape.delay_max_ms", "200"); err != nil {
+ if len(got) != 2 {
+ t.Fatalf("candidate rows = %d, want 2", len(got))
+ }
+ if got[0].ID != "999-series" || got[1].ID != "001-episode" {
+ t.Fatalf("scrape order = [%s, %s], want series-level row before episode row", got[0].ID, got[1].ID)
+ }
+}
+
+func TestEnrichLibraryIncludesMergedCloudLibraryMedia(t *testing.T) {
+ scraper, repos, closeServer := newTestScraper(t)
+ defer closeServer()
+
+ local := model.Library{Name: "番剧", Path: t.TempDir(), Type: "tv", Enabled: true}
+ cloud := model.Library{
+ Name: "OpenList · 番剧",
+ Path: BuildCloudLibraryPath("openlist", "/番剧", "/番剧"),
+ Type: "tv",
+ Enabled: true,
+ }
+ if err := repos.DB.Create(&local).Error; err != nil {
t.Fatal(err)
}
- if got := scraper.scrapeDelay(t.Context()); got != 800*time.Millisecond {
- t.Fatalf("normalized scrapeDelay = %s, want 800ms", got)
+ if err := repos.DB.Create(&cloud).Error; err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.DB.Create(&model.Media{
+ LibraryID: cloud.ID,
+ Title: "间谍过家家",
+ Path: "cloud://openlist/番剧/间谍过家家 - S02E02.mkv",
+ SeasonNum: 2,
+ EpisodeNum: 2,
+ ScrapeStatus: "pending",
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ result, err := scraper.EnrichLibraryDetailed(t.Context(), local.ID, true)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if result.Matched != 1 || result.Processed != 1 || result.Candidates != 1 || result.Failed != 0 {
+ t.Fatalf("result=%+v, want merged cloud media to be scraped once", result)
+ }
+ var got model.Media
+ if err := repos.DB.First(&got, "library_id = ?", cloud.ID).Error; err != nil {
+ t.Fatal(err)
+ }
+ if got.ScrapeStatus != "matched" || got.TMDbID != 12345 {
+ t.Fatalf("merged cloud media was not enriched: status=%q tmdb=%d", got.ScrapeStatus, got.TMDbID)
}
}
diff --git a/internal/service/scraper_test_helpers_test.go b/internal/service/scraper_test_helpers_test.go
new file mode 100644
index 0000000..df044ac
--- /dev/null
+++ b/internal/service/scraper_test_helpers_test.go
@@ -0,0 +1,26 @@
+package service
+
+func firstIndexFunc(values []string, match func(string) bool) int {
+ for i, value := range values {
+ if match(value) {
+ return i
+ }
+ }
+ return -1
+}
+
+func firstQuery(values []string) string {
+ if len(values) == 0 {
+ return ""
+ }
+ return values[0]
+}
+
+func lastIndexFunc(values []string, match func(string) bool) int {
+ for i := len(values) - 1; i >= 0; i-- {
+ if match(values[i]) {
+ return i
+ }
+ }
+ return -1
+}
diff --git a/internal/service/scraper_tmdb_details.go b/internal/service/scraper_tmdb_details.go
new file mode 100644
index 0000000..60724e6
--- /dev/null
+++ b/internal/service/scraper_tmdb_details.go
@@ -0,0 +1,176 @@
+package service
+
+import (
+ "context"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *ScraperService) fetchAndSaveTMDbExtendedMetadata(ctx context.Context, mediaID string, tmdbID int, mediaType string) {
+ detailCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), tmdbDetailsTimeout)
+ details, err := s.tmdb.GetDetails(detailCtx, tmdbID, mediaType)
+ cancel()
+ if err != nil {
+ s.log.Warn("failed to get details from tmdb",
+ zap.Int("tmdb_id", tmdbID),
+ zap.String("type", mediaType),
+ zap.Error(err))
+ return
+ }
+ if details == nil {
+ return
+ }
+ updates := map[string]any{}
+ if len(details.Languages) > 0 {
+ updates["languages"] = strings.Join(details.Languages, ",")
+ }
+ if len(details.Countries) > 0 {
+ updates["countries"] = strings.Join(details.Countries, ",")
+ }
+ if len(details.Genres) > 0 {
+ updates["genres"] = strings.Join(details.Genres, ",")
+ }
+ if len(updates) > 0 {
+ if err := s.repo.DB.Model(&model.Media{}).Where("id = ?", mediaID).
+ Updates(updates).Error; err != nil {
+ s.log.Warn("failed to save tmdb extended metadata",
+ zap.String("media_id", mediaID),
+ zap.Int("tmdb_id", tmdbID),
+ zap.Error(err))
+ }
+ }
+ s.log.Debug("enrich: saved extended metadata",
+ zap.String("media_id", mediaID),
+ zap.Strings("languages", details.Languages),
+ zap.Strings("countries", details.Countries),
+ zap.Strings("genres", details.Genres))
+}
+
+func (s *ScraperService) fetchAndSaveTMDbEpisodeDetails(ctx context.Context, m *model.Media, tmdbID int, matchYear int, options ScrapeOptions) bool {
+ if s == nil || s.tmdb == nil || !s.tmdb.Enabled() || m == nil || tmdbID <= 0 || m.EpisodeNum <= 0 {
+ return false
+ }
+ episodeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), tmdbDetailsTimeout)
+ episode, err := s.tmdb.GetTVEpisodeDetails(episodeCtx, tmdbID, m.SeasonNum, m.EpisodeNum)
+ cancel()
+ if err != nil {
+ s.log.Debug("failed to get tmdb episode details",
+ zap.String("media_id", m.ID),
+ zap.Int("tmdb_id", tmdbID),
+ zap.Int("season", m.SeasonNum),
+ zap.Int("episode", m.EpisodeNum),
+ zap.Error(err))
+ return false
+ }
+ if episode == nil {
+ return false
+ }
+ updates := tmdbEpisodeMetadataUpdates(m, episode, matchYear, options)
+ if len(updates) == 0 {
+ return false
+ }
+ if err := s.repo.DB.Model(&model.Media{}).Where("id = ?", m.ID).
+ Updates(updates).Error; err != nil {
+ s.log.Warn("failed to save tmdb episode metadata",
+ zap.String("media_id", m.ID),
+ zap.Int("tmdb_id", tmdbID),
+ zap.Int("season", m.SeasonNum),
+ zap.Int("episode", m.EpisodeNum),
+ zap.Error(err))
+ return false
+ }
+ return true
+}
+
+func tmdbEpisodeMetadataUpdates(m *model.Media, episode *TMDbEpisodeDetails, matchYear int, options ScrapeOptions) map[string]any {
+ updates := map[string]any{}
+ if episode == nil {
+ return updates
+ }
+ // Keep original_name at series level. Per-episode names can split one show
+ // into multiple cards because original_name participates in grouping.
+ if strings.TrimSpace(episode.Name) != "" {
+ updates["episode_title"] = strings.TrimSpace(episode.Name)
+ }
+ if strings.TrimSpace(episode.Overview) != "" {
+ updates["overview"] = strings.TrimSpace(episode.Overview)
+ }
+ if strings.TrimSpace(episode.StillURL) != "" && options.episodeArtworkEnabled() {
+ updates["backdrop_url"] = strings.TrimSpace(episode.StillURL)
+ }
+ if episode.Rating > 0 {
+ updates["rating"] = episode.Rating
+ }
+ if episode.AirYear > 0 && matchYear <= 0 {
+ updates["year"] = episode.AirYear
+ }
+ if m != nil && episode.Runtime > 0 && m.DurationSec <= 0 {
+ updates["duration_sec"] = episode.Runtime * 60
+ }
+ return updates
+}
+
+func (s *ScraperService) enrichDeferredEpisodeDetails(ctx context.Context, rows []model.Media, options ScrapeOptions) error {
+ if s == nil || s.tmdb == nil || !s.tmdb.Enabled() {
+ return nil
+ }
+ for i := range rows {
+ if rows[i].EpisodeNum <= 0 {
+ continue
+ }
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ default:
+ }
+ media, err := s.repo.Media.FindByID(ctx, rows[i].ID)
+ if err != nil || media == nil {
+ s.log.Debug("deferred episode metadata media missing", zap.String("media_id", rows[i].ID), zap.Error(err))
+ continue
+ }
+ if media.TMDbID <= 0 || media.EpisodeNum <= 0 {
+ continue
+ }
+ lib, _ := s.repo.Library.FindByID(ctx, media.LibraryID)
+ if !mediaIsEpisodic(media, lib) {
+ continue
+ }
+ if s.fetchAndSaveTMDbEpisodeDetails(ctx, media, media.TMDbID, media.Year, options) {
+ s.writeMediaNFOAfterScrape(ctx, media, lib)
+ s.invalidateMediaCache(ctx)
+ }
+ if i < len(rows)-1 {
+ if delay := s.scrapeDelay(ctx); delay > 0 {
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ case <-time.After(delay):
+ }
+ }
+ }
+ }
+ return nil
+}
+
+func (s *ScraperService) writeMediaNFOAfterScrape(ctx context.Context, m *model.Media, lib *model.Library) {
+ if s == nil || m == nil {
+ return
+ }
+ cloudMedia := isCloudMediaPath(m.Path) || (lib != nil && isCloudMediaPath(lib.Path))
+ if cloudMedia {
+ return
+ }
+ refreshed, err := s.repo.Media.FindByID(ctx, m.ID)
+ if err != nil || refreshed == nil {
+ return
+ }
+ if path, err := WriteMediaNFO(refreshed); err != nil {
+ s.log.Warn("write nfo after scrape failed", zap.String("media_id", m.ID), zap.Error(err))
+ } else {
+ s.log.Debug("write nfo after scrape", zap.String("media_id", m.ID), zap.String("path", path))
+ }
+}
diff --git a/internal/service/service.go b/internal/service/service.go
index e22c50d..c9fd2ec 100644
--- a/internal/service/service.go
+++ b/internal/service/service.go
@@ -75,6 +75,7 @@ type Container struct {
Site *SiteService
Device *DeviceService
Cache *RuntimeCacheService
+ Sessions *SessionTrackerService
stopCtx context.Context
stopCancel context.CancelFunc
@@ -82,179 +83,7 @@ type Container struct {
// New 构建服务容器。
func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Container {
- ApplyRuntimeSettings(context.Background(), cfg, repos, log)
-
- hub := NewHub(log)
- go hub.Run()
- tasks := NewTaskTrackerService(log, hub)
-
- // 初始化 SSE Hub
- sseHub := NewSSEHub(log)
- go sseHub.Run()
-
- probe := NewFFprobeService(cfg, log)
- runtimeCache := NewRuntimeCacheService(cfg, log)
- if searchBackend := repository.NewOpenSearchMediaBackend(cfg.Search); searchBackend != nil && repos != nil && repos.Media != nil {
- repos.Media.SetSearchBackend(searchBackend)
- if log != nil {
- log.Info("opensearch media search enabled", zap.String("index", cfg.Search.Index), zap.String("url", cfg.Search.OpenSearchURL))
- }
- }
- crypto := NewCryptoService(cfg.Secrets.JWTSecret, log)
- apiConfig := NewAPIConfigService(log, repos, crypto)
- tmdb := NewTMDbProvider(cfg, log, apiConfig)
- bangumi := NewBangumiProvider(cfg, log)
- thetvdb := NewTheTVDBProvider(cfg, log)
- douban := NewDoubanProvider(cfg, log)
- fanart := NewFanartProvider(cfg, log)
- adult := NewAdultProvider(log, apiConfig)
- scraper := NewScraperService(cfg, log, repos, tmdb, bangumi, thetvdb, fanart, hub, adult)
- scraper.SetDouban(douban)
- organizer := NewOrganizerService(cfg, log, repos)
- organizer.SetProbe(probe)
- organizer.SetScraper(scraper)
- discover := NewDiscoverService(log, tmdb)
- transcoder := NewTranscoderService(cfg, log, repos, hub)
- scanner := NewScannerService(cfg, log, repos, hub, probe, scraper)
- scanner.SetRuntimeCache(runtimeCache)
- organizePipeline := NewOrganizePipelineService(log, repos, organizer, scanner, tasks)
- watcher := NewWatcherService(log, repos, scanner)
- nfo := NewNFOService(log, repos)
- ai := NewAIService(cfg, log, apiConfig)
- duplicate := NewDuplicateService(log, repos, hub)
- filemanager := NewFileManagerService(cfg, log, repos)
- dlna := NewDLNAService(log)
- storage := NewStorageService(log, repos)
- emby := NewEmbyService(cfg, log, repos)
- backup := NewBackupService(cfg, log, repos.DB)
- notifier := NewNotifierService(log, repos)
- notifyChannels := NewNotifyChannelService(log, repos)
- scanner.SetNotifyChannels(notifyChannels)
- scraper.SetNotifyChannels(notifyChannels)
- playProfiles := NewPlayProfileService(log, repos)
- permissions := NewPermissionService(log, repos)
- storageCfg := NewStorageConfigService(log, repos, crypto)
- strmSvc := NewSTRMService(log, repos, cfg)
- scanner.SetStorageConfig(storageCfg)
- emby.SetRuntimeCache(runtimeCache)
- emby.SetCloudProbe(storageCfg, probe)
- downloadClients := NewDownloadClientService(log, repos)
- assistant := NewAssistantService(log, repos, ai)
- scheduler := NewSchedulerService(log, repos, scanner, transcoder, organizer, storageCfg, hub, cfg.Cache.CacheDir)
- scheduler.SetTaskTracker(tasks)
- scheduler.SetOrganizePipeline(organizePipeline)
-
- // 初始化认证相关服务
- tokenSvc := NewTokenService(cfg, log, repos)
- authSvc := NewAuthService(cfg, log, repos, tokenSvc, permissions)
- deviceSvc := NewDeviceService(log, repos)
- telegramBot := NewTelegramBotService(log, repos, crypto, authSvc)
- telegramBot.SetDeviceService(deviceSvc)
- telegramBot.SetBackupService(backup)
- // Allow the device-enforcement service to DM users (warnings / deletions)
- // through their Telegram binding before any destructive action.
- deviceSvc.SetNotifier(telegramBot.NotifyUserByID)
- apiConfigSvc := NewApiConfigService(cfg, log, repos, crypto)
- downloadMgr := NewDownloadManager(log, repos, crypto)
- notifySvc := NewNotifyService(log, repos, crypto)
-
- // 构建 FlareSolverr URL(如果启用)
- flareSolverrURL := ""
- if cfg.FlareSolverr.Enabled && cfg.FlareSolverr.URL != "" {
- flareSolverrURL = cfg.FlareSolverr.URL
- }
- siteSvc := NewSiteService(log, repos, flareSolverrURL)
- downloads := NewDownloadService(log, repos, hub, organizer, siteSvc)
- downloads.SetScanner(scanner)
- downloads.SetTaskTracker(tasks)
- downloads.SetOrganizePipeline(organizePipeline)
- downloads.SetNotifyChannels(notifyChannels)
- subscription := NewSubscriptionService(cfg, log, repos, downloads, siteSvc, hub)
- subscription.SetScraper(scraper)
- subscription.SetNotifyChannels(notifyChannels)
-
- // 让图片代理把媒体库根目录视为可读的本地图片位置:海报/封面等
- // sidecar 资源就存放在这些(用户自定义、任意)目录下,否则会被
- // 路径白名单挡掉、退化成占位图导致前端图片不显示。
- imageProxy := NewImageProxy(cfg, log)
- imageProxy.SetLibraryRootsProvider(func() []string {
- libs, err := repos.Library.List(context.Background())
- if err != nil {
- return nil
- }
- roots := make([]string, 0, len(libs))
- for _, l := range libs {
- if strings.TrimSpace(l.Path) != "" {
- roots = append(roots, l.Path)
- }
- }
- return roots
- })
- scanner.SetImageProxy(imageProxy)
-
- ctx, cancel := context.WithCancel(context.Background())
-
- return &Container{
- Cfg: cfg,
- Log: log,
- Repo: repos,
- WSHub: hub,
- SSEHub: sseHub,
- Tasks: tasks,
- Auth: authSvc,
- Media: NewMediaService(cfg, log, repos).SetRuntimeCache(runtimeCache),
- Scan: scanner,
- Stream: NewStreamService(cfg, log, repos, transcoder),
- Transcoder: transcoder,
- FFprobe: probe,
- TMDb: tmdb,
- Bangumi: bangumi,
- TheTVDB: thetvdb,
- Fanart: fanart,
- Scraper: scraper,
- Discover: discover,
- Playback: NewPlaybackService(log, repos),
- ImageProxy: imageProxy,
- Watcher: watcher,
- Downloads: downloads,
- Subscription: subscription,
- Subtitle: NewSubtitleService(log, repos),
- Stats: NewStatsService(log, repos).SetRuntimeCache(runtimeCache),
- Profile: NewProfileService(log, repos),
- Audit: NewAuditService(log, repos),
- NFO: nfo,
- AI: ai,
- APIConfig: apiConfig,
- Crypto: crypto,
- Duplicate: duplicate,
- FileManager: filemanager,
- DLNA: dlna,
- Scheduler: scheduler,
- Storage: storage,
- Emby: emby,
- Backup: backup,
- Notifier: notifier,
- NotifyChannels: notifyChannels,
- TelegramBot: telegramBot,
- PlayProfiles: playProfiles,
- Permissions: permissions,
- StorageCfg: storageCfg,
- STRM: strmSvc,
- DownloadClients: downloadClients,
- Assistant: assistant,
- Organizer: organizer,
- OrganizePipeline: organizePipeline,
- Douban: douban,
- Token: tokenSvc,
- ApiConfig: apiConfigSvc,
- DownloadMgr: downloadMgr,
- Notify: notifySvc,
- Site: siteSvc,
- Device: deviceSvc,
- Cache: runtimeCache,
- stopCtx: ctx,
- stopCancel: cancel,
- }
+ return newServiceContainer(cfg, log, repos)
}
// Boot 启动后台工作进程(watcher, downloads poller, subscription scheduler)。
@@ -480,6 +309,3 @@ func (c *Container) Close() {
c.Scheduler.Stop()
}
}
-
-// unused guard
-var _ = time.Now
diff --git a/internal/service/service_builder.go b/internal/service/service_builder.go
new file mode 100644
index 0000000..888e64e
--- /dev/null
+++ b/internal/service/service_builder.go
@@ -0,0 +1,196 @@
+package service
+
+import (
+ "context"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+type serviceContainerBuilder struct {
+ cfg *config.Config
+ log *zap.Logger
+ repos *repository.Container
+ c *Container
+}
+
+func newServiceContainer(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Container {
+ ApplyRuntimeSettings(context.Background(), cfg, repos, log)
+
+ builder := &serviceContainerBuilder{
+ cfg: cfg,
+ log: log,
+ repos: repos,
+ c: &Container{
+ Cfg: cfg,
+ Log: log,
+ Repo: repos,
+ },
+ }
+ builder.startRealtimeServices()
+ builder.initProviderServices()
+ builder.initContentServices()
+ builder.initAccessAndStorageServices()
+ builder.initIdentityServices()
+ builder.initSiteDownloadServices()
+ builder.initImageProxy()
+ builder.attachRuntimeContext()
+ return builder.c
+}
+
+func (b *serviceContainerBuilder) startRealtimeServices() {
+ b.c.WSHub = NewHub(b.log)
+ go b.c.WSHub.Run()
+ b.c.Tasks = NewTaskTrackerService(b.log, b.c.WSHub)
+
+ b.c.SSEHub = NewSSEHub(b.log)
+ go b.c.SSEHub.Run()
+}
+
+func (b *serviceContainerBuilder) initProviderServices() {
+ b.c.FFprobe = NewFFprobeService(b.cfg, b.log)
+ b.c.Cache = NewRuntimeCacheService(b.cfg, b.log)
+ b.configureMediaSearchBackend()
+
+ b.c.Crypto = NewCryptoService(b.cfg.Secrets.JWTSecret, b.log)
+ b.c.APIConfig = NewAPIConfigService(b.log, b.repos, b.c.Crypto)
+ b.c.TMDb = NewTMDbProvider(b.cfg, b.log, b.c.APIConfig)
+ b.c.Bangumi = NewBangumiProvider(b.cfg, b.log)
+ b.c.TheTVDB = NewTheTVDBProvider(b.cfg, b.log)
+ b.c.Douban = NewDoubanProvider(b.cfg, b.log)
+ b.c.Fanart = NewFanartProvider(b.cfg, b.log)
+
+ adult := NewAdultProvider(b.log, b.c.APIConfig)
+ b.c.Scraper = NewScraperService(
+ b.cfg, b.log, b.repos,
+ b.c.TMDb, b.c.Bangumi, b.c.TheTVDB, b.c.Fanart,
+ b.c.WSHub, adult,
+ )
+ b.c.Scraper.SetRuntimeCache(b.c.Cache)
+ b.c.Scraper.SetDouban(b.c.Douban)
+}
+
+func (b *serviceContainerBuilder) configureMediaSearchBackend() {
+ searchBackend := repository.NewOpenSearchMediaBackend(b.cfg.Search)
+ if searchBackend == nil || b.repos == nil || b.repos.Media == nil {
+ return
+ }
+ b.repos.Media.SetSearchBackend(searchBackend)
+ if b.log != nil {
+ b.log.Info("opensearch media search enabled", zap.String("index", b.cfg.Search.Index), zap.String("url", b.cfg.Search.OpenSearchURL))
+ }
+}
+
+func (b *serviceContainerBuilder) initContentServices() {
+ b.c.Organizer = NewOrganizerService(b.cfg, b.log, b.repos)
+ b.c.Organizer.SetProbe(b.c.FFprobe)
+ b.c.Organizer.SetScraper(b.c.Scraper)
+ b.c.Discover = NewDiscoverService(b.log, b.c.TMDb)
+ b.c.Transcoder = NewTranscoderService(b.cfg, b.log, b.repos, b.c.WSHub)
+ b.c.Scan = NewScannerService(b.cfg, b.log, b.repos, b.c.WSHub, b.c.FFprobe, b.c.Scraper)
+ b.c.Scan.SetRuntimeCache(b.c.Cache)
+ b.c.OrganizePipeline = NewOrganizePipelineService(b.log, b.repos, b.c.Organizer, b.c.Scan, b.c.Tasks)
+ b.c.Watcher = NewWatcherService(b.log, b.repos, b.c.Scan)
+ b.c.NFO = NewNFOService(b.log, b.repos)
+ b.c.AI = NewAIService(b.cfg, b.log, b.c.APIConfig)
+ b.c.Duplicate = NewDuplicateService(b.log, b.repos, b.c.WSHub)
+ b.c.FileManager = NewFileManagerService(b.cfg, b.log, b.repos)
+ b.c.DLNA = NewDLNAService(b.log)
+ b.c.Storage = NewStorageService(b.log, b.repos)
+ b.c.Emby = NewEmbyService(b.cfg, b.log, b.repos)
+ b.c.Backup = NewBackupService(b.cfg, b.log, b.repos.DB)
+ b.c.Notifier = NewNotifierService(b.log, b.repos)
+ b.c.NotifyChannels = NewNotifyChannelService(b.log, b.repos)
+ b.c.Scan.SetNotifyChannels(b.c.NotifyChannels)
+ b.c.Scraper.SetNotifyChannels(b.c.NotifyChannels)
+ b.c.Media = NewMediaService(b.cfg, b.log, b.repos).SetRuntimeCache(b.c.Cache)
+ b.c.Stream = NewStreamService(b.cfg, b.log, b.repos, b.c.Transcoder)
+ b.c.Playback = NewPlaybackService(b.log, b.repos)
+ b.c.Subtitle = NewSubtitleService(b.log, b.repos)
+ b.c.Stats = NewStatsService(b.log, b.repos).SetRuntimeCache(b.c.Cache)
+ b.c.Profile = NewProfileService(b.log, b.repos)
+ b.c.Audit = NewAuditService(b.log, b.repos)
+}
+
+func (b *serviceContainerBuilder) initAccessAndStorageServices() {
+ b.c.PlayProfiles = NewPlayProfileService(b.log, b.repos)
+ b.c.Permissions = NewPermissionService(b.log, b.repos)
+ b.c.StorageCfg = NewStorageConfigService(b.log, b.repos, b.c.Crypto)
+ b.c.STRM = NewSTRMService(b.log, b.repos, b.cfg)
+ b.c.Scan.SetStorageConfig(b.c.StorageCfg)
+ b.c.Subtitle.SetStorageConfig(b.c.StorageCfg)
+ b.c.Emby.SetRuntimeCache(b.c.Cache)
+ b.c.Emby.SetCloudProbe(b.c.StorageCfg, b.c.FFprobe)
+ b.c.DownloadClients = NewDownloadClientService(b.log, b.repos)
+ b.c.Assistant = NewAssistantService(b.log, b.repos, b.c.AI)
+ b.c.Scheduler = NewSchedulerService(
+ b.log, b.repos, b.c.Scan, b.c.Transcoder,
+ b.c.Organizer, b.c.StorageCfg, b.c.WSHub, b.cfg.Cache.CacheDir,
+ )
+ b.c.Scheduler.SetTaskTracker(b.c.Tasks)
+ b.c.Scheduler.SetOrganizePipeline(b.c.OrganizePipeline)
+}
+
+func (b *serviceContainerBuilder) initIdentityServices() {
+ b.c.Token = NewTokenService(b.cfg, b.log, b.repos)
+ b.c.Auth = NewAuthService(b.cfg, b.log, b.repos, b.c.Token, b.c.Permissions)
+ b.c.Sessions = NewSessionTrackerService(b.log)
+ b.c.Device = NewDeviceService(b.log, b.repos)
+ b.c.Device.SetSessionTracker(b.c.Sessions)
+ b.c.TelegramBot = NewTelegramBotService(b.log, b.repos, b.c.Crypto, b.c.Auth)
+ b.c.TelegramBot.SetDeviceService(b.c.Device)
+ b.c.TelegramBot.SetBackupService(b.c.Backup)
+ // Device enforcement notifies users through their Telegram binding before destructive actions.
+ b.c.Device.SetNotifier(b.c.TelegramBot.NotifyUserByID)
+ b.c.ApiConfig = NewApiConfigService(b.cfg, b.log, b.repos, b.c.Crypto)
+ b.c.DownloadMgr = NewDownloadManager(b.log, b.repos, b.c.Crypto)
+ b.c.Notify = NewNotifyService(b.log, b.repos, b.c.Crypto)
+}
+
+func (b *serviceContainerBuilder) initSiteDownloadServices() {
+ b.c.Site = NewSiteService(b.log, b.repos, b.flareSolverrURL())
+ b.c.Downloads = NewDownloadService(b.log, b.repos, b.c.WSHub, b.c.Organizer, b.c.Site)
+ b.c.Downloads.SetScanner(b.c.Scan)
+ b.c.Downloads.SetTaskTracker(b.c.Tasks)
+ b.c.Downloads.SetOrganizePipeline(b.c.OrganizePipeline)
+ b.c.Downloads.SetNotifyChannels(b.c.NotifyChannels)
+ b.c.Subscription = NewSubscriptionService(b.cfg, b.log, b.repos, b.c.Downloads, b.c.Site, b.c.WSHub)
+ b.c.Subscription.SetScraper(b.c.Scraper)
+ b.c.Subscription.SetNotifyChannels(b.c.NotifyChannels)
+}
+
+func (b *serviceContainerBuilder) initImageProxy() {
+ b.c.ImageProxy = NewImageProxy(b.cfg, b.log)
+ b.c.ImageProxy.SetLibraryRootsProvider(b.libraryRoots)
+ b.c.Scan.SetImageProxy(b.c.ImageProxy)
+ b.c.Scraper.SetImageProxy(b.c.ImageProxy)
+ b.c.Discover.SetImageProxy(b.c.ImageProxy)
+}
+
+func (b *serviceContainerBuilder) libraryRoots() []string {
+ libs, err := b.repos.Library.List(context.Background())
+ if err != nil {
+ return nil
+ }
+ roots := make([]string, 0, len(libs))
+ for _, l := range libs {
+ if strings.TrimSpace(l.Path) != "" {
+ roots = append(roots, l.Path)
+ }
+ }
+ return roots
+}
+
+func (b *serviceContainerBuilder) flareSolverrURL() string {
+ if b.cfg.FlareSolverr.Enabled && b.cfg.FlareSolverr.URL != "" {
+ return b.cfg.FlareSolverr.URL
+ }
+ return ""
+}
+
+func (b *serviceContainerBuilder) attachRuntimeContext() {
+ b.c.stopCtx, b.c.stopCancel = context.WithCancel(context.Background())
+}
diff --git a/internal/service/session_tracker.go b/internal/service/session_tracker.go
new file mode 100644
index 0000000..52f53f7
--- /dev/null
+++ b/internal/service/session_tracker.go
@@ -0,0 +1,313 @@
+package service
+
+import (
+ "context"
+ "sort"
+ "strings"
+ "sync"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+const (
+ realtimeSessionTTL = 30 * time.Minute
+ realtimeSessionOnlineTTL = 5 * time.Minute
+)
+
+func RealtimeDeletionGuardWindow() time.Duration {
+ return realtimeSessionTTL
+}
+
+// RealtimeSession is an in-memory Emby-compatible session view. It mirrors the
+// information reported by Emby clients through AuthenticateByName and
+// /Sessions/Playing/* without requiring Playback Reporting persistence.
+type RealtimeSession struct {
+ ID string `json:"id"`
+ UserID string `json:"user_id"`
+ UserName string `json:"user_name,omitempty"`
+ DeviceID string `json:"device_id"`
+ DeviceName string `json:"device_name,omitempty"`
+ Client string `json:"client,omitempty"`
+ RemoteEndPoint string `json:"remote_end_point,omitempty"`
+ LastActivityAt time.Time `json:"last_activity_at"`
+ ItemID string `json:"item_id,omitempty"`
+ PositionTicks int64 `json:"position_ticks,omitempty"`
+ RuntimeTicks int64 `json:"runtime_ticks,omitempty"`
+ IsPlaying bool `json:"is_playing"`
+ IsPaused bool `json:"is_paused"`
+ LastPlaybackAt *time.Time `json:"last_playback_at,omitempty"`
+}
+
+type realtimeSessionInput struct {
+ UserID string
+ UserName string
+ DeviceID string
+ DeviceName string
+ Client string
+ RemoteEndPoint string
+ ItemID string
+ PositionTicks int64
+ RuntimeTicks int64
+ IsPlaying bool
+ IsPaused bool
+}
+
+type userRealtimeActivity struct {
+ LastActivityAt *time.Time
+ ActiveDeviceCount int
+ Online bool
+}
+
+// SessionTrackerService keeps recent client state in memory. It is intentionally
+// not durable: process restart clears transient online status, while normal
+// login/playback requests repopulate it immediately.
+type SessionTrackerService struct {
+ log *zap.Logger
+
+ mu sync.RWMutex
+ sessions map[string]RealtimeSession
+ now func() time.Time
+}
+
+func NewSessionTrackerService(log *zap.Logger) *SessionTrackerService {
+ return &SessionTrackerService{
+ log: log,
+ sessions: make(map[string]RealtimeSession),
+ now: time.Now,
+ }
+}
+
+func (s *SessionTrackerService) RecordLogin(ctx context.Context, userID, userName, deviceID, deviceName, client, remoteEndPoint string) {
+ if s == nil {
+ return
+ }
+ s.upsert(ctx, realtimeSessionInput{
+ UserID: userID,
+ UserName: userName,
+ DeviceID: deviceID,
+ DeviceName: deviceName,
+ Client: client,
+ RemoteEndPoint: remoteEndPoint,
+ })
+}
+
+func (s *SessionTrackerService) RecordPlayback(ctx context.Context, userID, userName, deviceID, deviceName, client, remoteEndPoint, itemID string, positionTicks, runtimeTicks int64, stopped bool) {
+ if s == nil {
+ return
+ }
+ s.upsert(ctx, realtimeSessionInput{
+ UserID: userID,
+ UserName: userName,
+ DeviceID: deviceID,
+ DeviceName: deviceName,
+ Client: client,
+ RemoteEndPoint: remoteEndPoint,
+ ItemID: itemID,
+ PositionTicks: positionTicks,
+ RuntimeTicks: runtimeTicks,
+ IsPlaying: !stopped,
+ })
+}
+
+func (s *SessionTrackerService) Logout(ctx context.Context, userID, deviceID, remoteEndPoint string) {
+ if s == nil {
+ return
+ }
+ userID = strings.TrimSpace(userID)
+ if userID == "" {
+ return
+ }
+ deviceID = strings.TrimSpace(deviceID)
+ remoteEndPoint = strings.TrimSpace(remoteEndPoint)
+ s.mu.Lock()
+ defer s.mu.Unlock()
+ for key, sess := range s.sessions {
+ if sess.UserID != userID {
+ continue
+ }
+ if deviceID != "" && sess.DeviceID != deviceID {
+ continue
+ }
+ if deviceID == "" && remoteEndPoint != "" && sess.RemoteEndPoint != remoteEndPoint {
+ continue
+ }
+ delete(s.sessions, key)
+ }
+}
+
+func (s *SessionTrackerService) List(ctx context.Context) []RealtimeSession {
+ if s == nil {
+ return nil
+ }
+ now := s.now()
+ s.mu.Lock()
+ defer s.mu.Unlock()
+ s.pruneLocked(now)
+ out := make([]RealtimeSession, 0, len(s.sessions))
+ for _, sess := range s.sessions {
+ out = append(out, sess)
+ }
+ sort.SliceStable(out, func(i, j int) bool {
+ return out[i].LastActivityAt.After(out[j].LastActivityAt)
+ })
+ return out
+}
+
+func (s *SessionTrackerService) ListByUser(ctx context.Context, userID string) []RealtimeSession {
+ userID = strings.TrimSpace(userID)
+ if userID == "" {
+ return nil
+ }
+ all := s.List(ctx)
+ out := make([]RealtimeSession, 0, len(all))
+ for _, sess := range all {
+ if sess.UserID == userID {
+ out = append(out, sess)
+ }
+ }
+ return out
+}
+
+func (s *SessionTrackerService) ApplyToUsers(ctx context.Context, users []model.User) {
+ if s == nil || len(users) == 0 {
+ return
+ }
+ activity := s.activityByUser(ctx)
+ for i := range users {
+ a, ok := activity[users[i].ID]
+ if !ok {
+ continue
+ }
+ if a.LastActivityAt != nil && (users[i].LastLoginAt == nil || a.LastActivityAt.After(*users[i].LastLoginAt)) {
+ t := *a.LastActivityAt
+ users[i].LastLoginAt = &t
+ }
+ users[i].RealtimeOnline = a.Online
+ users[i].RealtimeDeviceCount = a.ActiveDeviceCount
+ }
+}
+
+func (s *SessionTrackerService) UserRecentlyActive(ctx context.Context, userID string, within time.Duration) bool {
+ if s == nil || within <= 0 {
+ return false
+ }
+ activity := s.activityByUser(ctx)[strings.TrimSpace(userID)]
+ return activity.LastActivityAt != nil && activity.LastActivityAt.After(s.now().Add(-within))
+}
+
+func (s *SessionTrackerService) activityByUser(ctx context.Context) map[string]userRealtimeActivity {
+ sessions := s.List(ctx)
+ now := s.now()
+ out := make(map[string]userRealtimeActivity)
+ seenDevices := make(map[string]map[string]struct{})
+ for _, sess := range sessions {
+ if strings.TrimSpace(sess.UserID) == "" {
+ continue
+ }
+ a := out[sess.UserID]
+ if a.LastActivityAt == nil || sess.LastActivityAt.After(*a.LastActivityAt) {
+ t := sess.LastActivityAt
+ a.LastActivityAt = &t
+ }
+ if sess.LastActivityAt.After(now.Add(-realtimeSessionOnlineTTL)) {
+ a.Online = true
+ }
+ if seenDevices[sess.UserID] == nil {
+ seenDevices[sess.UserID] = map[string]struct{}{}
+ }
+ seenDevices[sess.UserID][sessionDeviceKey(sess)] = struct{}{}
+ a.ActiveDeviceCount = len(seenDevices[sess.UserID])
+ out[sess.UserID] = a
+ }
+ return out
+}
+
+func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionInput) {
+ userID := strings.TrimSpace(in.UserID)
+ if userID == "" {
+ return
+ }
+ now := s.now()
+ in.DeviceID = strings.TrimSpace(in.DeviceID)
+ in.DeviceName = strings.TrimSpace(in.DeviceName)
+ in.Client = strings.TrimSpace(in.Client)
+ in.RemoteEndPoint = strings.TrimSpace(in.RemoteEndPoint)
+ if in.DeviceID == "" {
+ in.DeviceID = fallbackSessionDeviceID(in.DeviceName, in.Client, in.RemoteEndPoint)
+ }
+ key := userID + "\x00" + in.DeviceID
+ s.mu.Lock()
+ defer s.mu.Unlock()
+ s.pruneLocked(now)
+ existing := s.sessions[key]
+ if strings.TrimSpace(in.UserName) == "" {
+ in.UserName = existing.UserName
+ }
+ if in.DeviceName == "" {
+ in.DeviceName = existing.DeviceName
+ }
+ if in.Client == "" {
+ in.Client = existing.Client
+ }
+ if in.RemoteEndPoint == "" {
+ in.RemoteEndPoint = existing.RemoteEndPoint
+ }
+ lastPlaybackAt := existing.LastPlaybackAt
+ if in.ItemID != "" || in.IsPlaying {
+ t := now
+ lastPlaybackAt = &t
+ }
+ s.sessions[key] = RealtimeSession{
+ ID: key,
+ UserID: userID,
+ UserName: strings.TrimSpace(in.UserName),
+ DeviceID: in.DeviceID,
+ DeviceName: in.DeviceName,
+ Client: in.Client,
+ RemoteEndPoint: in.RemoteEndPoint,
+ LastActivityAt: now,
+ ItemID: firstNonEmptyString(in.ItemID, existing.ItemID),
+ PositionTicks: in.PositionTicks,
+ RuntimeTicks: in.RuntimeTicks,
+ IsPlaying: in.IsPlaying,
+ IsPaused: in.IsPaused,
+ LastPlaybackAt: lastPlaybackAt,
+ }
+}
+
+func (s *SessionTrackerService) pruneLocked(now time.Time) {
+ expiresBefore := now.Add(-realtimeSessionTTL)
+ for key, sess := range s.sessions {
+ if sess.LastActivityAt.Before(expiresBefore) {
+ delete(s.sessions, key)
+ }
+ }
+}
+
+func fallbackSessionDeviceID(deviceName, client, remoteEndPoint string) string {
+ parts := []string{strings.TrimSpace(deviceName), strings.TrimSpace(client), strings.TrimSpace(remoteEndPoint)}
+ joined := strings.Trim(strings.Join(parts, "|"), "|")
+ if joined == "" {
+ joined = "unknown"
+ }
+ return "rt-" + fingerprint(client, joined)
+}
+
+func sessionDeviceKey(sess RealtimeSession) string {
+ if strings.TrimSpace(sess.DeviceID) != "" {
+ return strings.TrimSpace(sess.DeviceID)
+ }
+ return fallbackSessionDeviceID(sess.DeviceName, sess.Client, sess.RemoteEndPoint)
+}
+
+func firstNonEmptyString(values ...string) string {
+ for _, value := range values {
+ if strings.TrimSpace(value) != "" {
+ return strings.TrimSpace(value)
+ }
+ }
+ return ""
+}
diff --git a/internal/service/session_tracker_test.go b/internal/service/session_tracker_test.go
new file mode 100644
index 0000000..bcc2125
--- /dev/null
+++ b/internal/service/session_tracker_test.go
@@ -0,0 +1,92 @@
+package service
+
+import (
+ "testing"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestSessionTrackerAppliesRealtimeActivityToUsers(t *testing.T) {
+ tracker := NewSessionTrackerService(zap.NewNop())
+ now := time.Date(2026, 6, 21, 10, 0, 0, 0, time.UTC)
+ tracker.now = func() time.Time { return now }
+ old := now.Add(-8 * time.Hour)
+ users := []model.User{{Base: model.Base{ID: "u1"}, Username: "admin", LastLoginAt: &old}}
+
+ tracker.RecordLogin(t.Context(), "u1", "admin", "web-1", "Web", "Browser", "127.0.0.1")
+ tracker.ApplyToUsers(t.Context(), users)
+
+ if users[0].LastLoginAt == nil || !users[0].LastLoginAt.Equal(now) {
+ t.Fatalf("last_login_at = %v, want realtime %v", users[0].LastLoginAt, now)
+ }
+ if !users[0].RealtimeOnline || users[0].RealtimeDeviceCount != 1 {
+ t.Fatalf("realtime flags online=%v devices=%d", users[0].RealtimeOnline, users[0].RealtimeDeviceCount)
+ }
+}
+
+func TestDeviceListMergesRealtimeSessions(t *testing.T) {
+ repos := newSessionTrackerTestRepos(t)
+ user := model.User{Base: model.Base{ID: "u1"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
+ if err := repos.User.Create(t.Context(), &user); err != nil {
+ t.Fatal(err)
+ }
+ tracker := NewSessionTrackerService(zap.NewNop())
+ now := time.Date(2026, 6, 21, 11, 0, 0, 0, time.UTC)
+ tracker.now = func() time.Time { return now }
+ tracker.RecordPlayback(t.Context(), user.ID, user.Username, "dev-1", "Apple TV", "Yamby", "10.0.0.8", "media-1", 123, 456, false)
+
+ device := NewDeviceService(zap.NewNop(), repos)
+ device.SetSessionTracker(tracker)
+ rows, err := device.ListDevices(t.Context(), user.ID)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 1 {
+ t.Fatalf("devices = %d, want 1", len(rows))
+ }
+ if !rows[0].Realtime || !rows[0].Online || !rows[0].Playing {
+ t.Fatalf("realtime device flags = realtime:%v online:%v playing:%v", rows[0].Realtime, rows[0].Online, rows[0].Playing)
+ }
+ if rows[0].DeviceName != "Apple TV" || rows[0].Client != "Yamby" || !rows[0].LastSeenAt.Equal(now) {
+ t.Fatalf("device row = %#v", rows[0])
+ }
+}
+
+func TestRealtimeRecentLoginProtectsCleanupCandidate(t *testing.T) {
+ repos := newSessionTrackerTestRepos(t)
+ now := time.Date(2026, 6, 21, 12, 0, 0, 0, time.UTC)
+ old := now.Add(-30 * 24 * time.Hour)
+ user := model.User{Base: model.Base{ID: "u1"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &old}
+ if err := repos.User.Create(t.Context(), &user); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Setting.Set(t.Context(), SettingAccountCleanupEnabled, "true"); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Setting.Set(t.Context(), SettingAccountCleanupRules, `[{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7}]`); err != nil {
+ t.Fatal(err)
+ }
+ tracker := NewSessionTrackerService(zap.NewNop())
+ tracker.now = func() time.Time { return now }
+ tracker.RecordLogin(t.Context(), user.ID, user.Username, "web", "Web", "Browser", "127.0.0.1")
+ device := NewDeviceService(zap.NewNop(), repos)
+ device.SetSessionTracker(tracker)
+
+ candidates, err := device.PreviewAccountCleanup(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(candidates) != 0 {
+ t.Fatalf("recent realtime login should protect user, got candidates %#v", candidates)
+ }
+}
+
+func newSessionTrackerTestRepos(t *testing.T) *repository.Container {
+ t.Helper()
+ db := newServiceTestDB(t, &model.User{}, &model.Setting{}, &model.UserDevice{}, &model.SignIn{}, &model.PlaybackHistory{})
+ return repository.New(db)
+}
diff --git a/internal/service/site_adapter_test.go b/internal/service/site_adapter_test.go
index c167805..bd4ae7b 100644
--- a/internal/service/site_adapter_test.go
+++ b/internal/service/site_adapter_test.go
@@ -11,9 +11,6 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
- "gorm.io/gorm"
-
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
@@ -201,13 +198,7 @@ func TestMTeamPublishedAPIRateLimits(t *testing.T) {
}
func TestPersistentSiteAPIRateLimiterPersistsSlidingWindow(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Setting{})
repos := repository.New(db)
now := time.Date(2026, 6, 20, 12, 0, 0, 0, time.UTC)
limiter := newPersistentSiteAPIRateLimiter(repos)
@@ -220,7 +211,7 @@ func TestPersistentSiteAPIRateLimiterPersistsSlidingWindow(t *testing.T) {
if err := limiter.Allow(t.Context(), "mteam:test", limit); err != nil {
t.Fatalf("second allow: %v", err)
}
- err = limiter.Allow(t.Context(), "mteam:test", limit)
+ err := limiter.Allow(t.Context(), "mteam:test", limit)
var limited *siteAPIRateLimitError
if !errors.As(err, &limited) {
t.Fatalf("third allow error = %v, want siteAPIRateLimitError", err)
diff --git a/internal/service/site_test.go b/internal/service/site_test.go
index e7f6ef1..1cbe131 100644
--- a/internal/service/site_test.go
+++ b/internal/service/site_test.go
@@ -7,22 +7,14 @@ import (
"strings"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestSiteUpdateKeepsSecretsWhenPatchIsBlank(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Site{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Site{})
svc := NewSiteService(zap.NewNop(), &repository.Container{DB: db}, "")
site := &model.Site{
Name: "M-Team",
@@ -63,13 +55,7 @@ func TestYemaPTTestConnectionDoesNotFallbackAfterAuthFailure(t *testing.T) {
}))
defer server.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Site{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Site{})
repos := repository.New(db)
svc := NewSiteService(zap.NewNop(), repos, "")
site := &model.Site{
diff --git a/internal/service/stats_test.go b/internal/service/stats_test.go
index 309b590..7844395 100644
--- a/internal/service/stats_test.go
+++ b/internal/service/stats_test.go
@@ -3,22 +3,14 @@ package service
import (
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
func TestStatsComputeFiltersDisabledLibraries(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.User{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.User{})
repos := repository.New(db)
enabled := &model.Library{Name: "电影", Path: "/media/movies", Type: "movie", Enabled: true}
disabled := &model.Library{Name: "停用库", Path: "/media/disabled", Type: "movie", Enabled: false}
diff --git a/internal/service/storage_cloud_resolve.go b/internal/service/storage_cloud_resolve.go
new file mode 100644
index 0000000..942e363
--- /dev/null
+++ b/internal/service/storage_cloud_resolve.go
@@ -0,0 +1,315 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+type cloudResolveCacheEntry struct {
+ link *cloud.DirectLink
+ expiresAt time.Time
+ hits int
+ lastHit time.Time
+}
+
+type cloudResolveCall struct {
+ done chan struct{}
+ link *cloud.DirectLink
+ err error
+}
+
+const (
+ cloudResolveHotHitThreshold = 3
+ cloudResolveBackgroundRefreshMax = 30 * time.Second
+)
+
+// CloudResolve resolves a cloud file reference to a direct link.
+//
+// clientUA is the User-Agent of the playback client that will follow the 302
+// redirect. Some provider CDN links are bound to the UA used to request them,
+// so we resolve with the client's own UA. When clientUA is empty the provider's
+// default UA is used.
+func (s *StorageConfigService) CloudResolve(ctx context.Context, typ, fileRef, clientUA string) (*cloud.DirectLink, error) {
+ if s == nil {
+ return nil, errors.New("storage config service unavailable")
+ }
+ cacheKey := s.resolveCacheKey(typ, fileRef, clientUA)
+ if link, ok, refresh := s.cachedResolve(cacheKey, typ); ok {
+ if refresh {
+ s.refreshResolveInBackground(cacheKey, typ, fileRef, clientUA)
+ }
+ return link, nil
+ }
+ if call, owner := s.beginResolve(cacheKey); !owner {
+ select {
+ case <-call.done:
+ if call.err != nil {
+ return nil, call.err
+ }
+ return cloneDirectLink(call.link), nil
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ }
+ } else {
+ defer s.finishResolve(cacheKey, call)
+ p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
+ if err != nil {
+ call.err = err
+ return nil, err
+ }
+ link, err := p.Resolve(ctx, fileRef)
+ if err != nil {
+ call.err = err
+ return nil, err
+ }
+ call.link = cloneDirectLink(link)
+ s.storeResolvedLink(cacheKey, typ, link)
+ return cloneDirectLink(link), nil
+ }
+}
+
+func (s *StorageConfigService) resolveCacheKey(typ, fileRef, clientUA string) string {
+ return strings.TrimSpace(typ) + "\x00" + strings.TrimSpace(fileRef) + "\x00" + strings.TrimSpace(clientUA)
+}
+
+func (s *StorageConfigService) cachedResolve(key, typ string) (*cloud.DirectLink, bool, bool) {
+ s.resolveMu.Lock()
+ defer s.resolveMu.Unlock()
+ if s.resolveCache == nil {
+ s.resolveCache = make(map[string]cloudResolveCacheEntry)
+ return nil, false, false
+ }
+ entry, ok := s.resolveCache[key]
+ now := time.Now()
+ if !ok || now.After(entry.expiresAt) {
+ if ok {
+ delete(s.resolveCache, key)
+ }
+ return nil, false, false
+ }
+ entry.hits++
+ entry.lastHit = now
+ s.resolveCache[key] = entry
+ refreshWindow := cloudResolveHotRefreshWindow(cloudResolveCacheTTL(typ))
+ shouldRefresh := entry.hits >= cloudResolveHotHitThreshold &&
+ refreshWindow > 0 &&
+ now.Add(refreshWindow).After(entry.expiresAt)
+ return cloneDirectLink(entry.link), true, shouldRefresh
+}
+
+func (s *StorageConfigService) beginResolve(key string) (*cloudResolveCall, bool) {
+ s.resolveMu.Lock()
+ defer s.resolveMu.Unlock()
+ if s.resolveFlight == nil {
+ s.resolveFlight = make(map[string]*cloudResolveCall)
+ }
+ if call := s.resolveFlight[key]; call != nil {
+ return call, false
+ }
+ call := &cloudResolveCall{done: make(chan struct{})}
+ s.resolveFlight[key] = call
+ return call, true
+}
+
+func (s *StorageConfigService) finishResolve(key string, call *cloudResolveCall) {
+ s.resolveMu.Lock()
+ if current := s.resolveFlight[key]; current == call {
+ delete(s.resolveFlight, key)
+ }
+ s.resolveMu.Unlock()
+ close(call.done)
+}
+
+func (s *StorageConfigService) refreshResolveInBackground(key, typ, fileRef, clientUA string) {
+ if s == nil {
+ return
+ }
+ go func() {
+ call, owner := s.beginResolve(key)
+ if !owner {
+ return
+ }
+ defer s.finishResolve(key, call)
+ ctx, cancel := context.WithTimeout(context.Background(), cloudResolveBackgroundRefreshMax)
+ defer cancel()
+ p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
+ if err != nil {
+ call.err = err
+ if s.log != nil {
+ s.log.Debug("refresh cloud direct link failed", zap.String("provider", typ), zap.Error(err))
+ }
+ return
+ }
+ link, err := p.Resolve(ctx, fileRef)
+ if err != nil {
+ call.err = err
+ if s.log != nil {
+ s.log.Debug("refresh cloud direct link failed", zap.String("provider", typ), zap.Error(err))
+ }
+ return
+ }
+ call.link = cloneDirectLink(link)
+ s.storeResolvedLink(key, typ, link)
+ }()
+}
+
+func (s *StorageConfigService) storeResolvedLink(key, typ string, link *cloud.DirectLink) {
+ if link == nil || strings.TrimSpace(link.URL) == "" {
+ return
+ }
+ ttl := cloudResolveCacheTTL(typ)
+ if ttl <= 0 {
+ return
+ }
+ s.resolveMu.Lock()
+ defer s.resolveMu.Unlock()
+ if s.resolveCache == nil {
+ s.resolveCache = make(map[string]cloudResolveCacheEntry)
+ }
+ now := time.Now()
+ hits := 0
+ if existing, ok := s.resolveCache[key]; ok {
+ hits = existing.hits
+ }
+ s.resolveCache[key] = cloudResolveCacheEntry{link: cloneDirectLink(link), expiresAt: now.Add(ttl), hits: hits, lastHit: now}
+}
+
+func cloudResolveHotRefreshWindow(ttl time.Duration) time.Duration {
+ if ttl <= 0 {
+ return 0
+ }
+ window := ttl / 4
+ if window < 15*time.Second {
+ window = 15 * time.Second
+ }
+ if window > 2*time.Minute {
+ window = 2 * time.Minute
+ }
+ return window
+}
+
+func cloudResolveCacheTTL(typ string) time.Duration {
+ switch typ {
+ case cloud.Type115, cloud.TypeCloudDrive2, cloud.TypeOpenList:
+ return 2 * time.Minute
+ default:
+ return 5 * time.Minute
+ }
+}
+
+func cloneDirectLink(link *cloud.DirectLink) *cloud.DirectLink {
+ if link == nil {
+ return nil
+ }
+ out := &cloud.DirectLink{
+ URL: link.URL,
+ Headers: make(map[string]string, len(link.Headers)),
+ Proxy: link.Proxy,
+ }
+ for k, v := range link.Headers {
+ out.Headers[k] = v
+ }
+ return out
+}
+
+func (s *StorageConfigService) clearResolveCacheForType(typ string) {
+ typ = strings.TrimSpace(typ)
+ if typ == "" {
+ return
+ }
+ prefix := typ + "\x00"
+ s.resolveMu.Lock()
+ defer s.resolveMu.Unlock()
+ for key := range s.resolveCache {
+ if strings.HasPrefix(key, prefix) {
+ delete(s.resolveCache, key)
+ }
+ }
+ for key, call := range s.resolveFlight {
+ if strings.HasPrefix(key, prefix) && call != nil {
+ call.err = fmt.Errorf("%s storage config changed", typ)
+ }
+ }
+}
+
+func (s *StorageConfigService) CloudResolveUncached(ctx context.Context, typ, fileRef, clientUA string) (*cloud.DirectLink, error) {
+ p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
+ if err != nil {
+ return nil, err
+ }
+ return p.Resolve(ctx, fileRef)
+}
+
+// CloudReadText resolves a small cloud file and returns its text payload. It is
+// used for cloud-hosted .strm files: the scanner reads the STRM target once and
+// stores the real playback URL, while the media bytes still stay in the cloud.
+func (s *StorageConfigService) CloudReadText(ctx context.Context, typ, fileRef string, limit int64) (string, error) {
+ if limit <= 0 {
+ limit = 64 << 10
+ }
+ link, err := s.CloudResolve(ctx, typ, fileRef, "")
+ if err != nil {
+ return "", err
+ }
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, link.URL, nil)
+ if err != nil {
+ return "", err
+ }
+ for k, v := range link.Headers {
+ req.Header.Set(k, v)
+ }
+ resp, err := s.client.Do(req)
+ if err != nil {
+ return "", err
+ }
+ defer resp.Body.Close()
+ if resp.StatusCode < 200 || resp.StatusCode >= 300 {
+ return "", fmt.Errorf("%s: read strm returned http %d", typ, resp.StatusCode)
+ }
+ body, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
+ if err != nil {
+ return "", err
+ }
+ if int64(len(body)) > limit {
+ return "", fmt.Errorf("%s: strm file is too large", typ)
+ }
+ return strings.TrimSpace(strings.TrimPrefix(string(body), "\ufeff")), nil
+}
+
+// cloudProviderWithUA builds a provider, overriding the request UA when a
+// non-empty clientUA is supplied.
+func (s *StorageConfigService) cloudProviderWithUA(ctx context.Context, typ, clientUA string) (cloud.Provider, error) {
+ if !cloud.IsCloudType(typ) {
+ return nil, fmt.Errorf("not a cloud provider: %q", typ)
+ }
+ view, err := s.Get(ctx, typ)
+ if err != nil {
+ return nil, err
+ }
+ if view == nil {
+ return nil, fmt.Errorf("%s storage not configured", typ)
+ }
+ if !view.Enabled {
+ return nil, fmt.Errorf("%s storage disabled", typ)
+ }
+ cfg := view.Config
+ if strings.TrimSpace(clientUA) != "" {
+ // Copy so we never mutate the cached view config.
+ cp := make(map[string]any, len(cfg)+1)
+ for k, v := range cfg {
+ cp[k] = v
+ }
+ cp["ua"] = clientUA
+ cfg = cp
+ }
+ return cloud.New(typ, cfg, s.clientForConfig(cfg))
+}
diff --git a/internal/service/storage_config.go b/internal/service/storage_config.go
index e7e694c..984be87 100644
--- a/internal/service/storage_config.go
+++ b/internal/service/storage_config.go
@@ -1,4 +1,4 @@
-// Package service — Alist / S3 / WebDAV configuration management.
+// Package service — external storage configuration management.
//
// StorageConfigService stores connection settings encrypted at rest
// (via CryptoService). It also exposes a Test() probe so the React UI
@@ -10,9 +10,7 @@ import (
"encoding/json"
"errors"
"fmt"
- "io"
"net/http"
- "net/url"
"strconv"
"strings"
"sync"
@@ -36,24 +34,6 @@ type StorageConfigService struct {
resolveFlight map[string]*cloudResolveCall
}
-type cloudResolveCacheEntry struct {
- link *cloud.DirectLink
- expiresAt time.Time
- hits int
- lastHit time.Time
-}
-
-type cloudResolveCall struct {
- done chan struct{}
- link *cloud.DirectLink
- err error
-}
-
-const (
- cloudResolveHotHitThreshold = 3
- cloudResolveBackgroundRefreshMax = 30 * time.Second
-)
-
// NewStorageConfigService is the constructor.
func NewStorageConfigService(log *zap.Logger, repo *repository.Container, crypto *CryptoService) *StorageConfigService {
return &StorageConfigService{
@@ -108,6 +88,9 @@ func (s *StorageConfigService) List(ctx context.Context) ([]StorageView, error)
}
out := make([]StorageView, 0, len(rows))
for _, r := range rows {
+ if !IsAdminStorageConfigurable(r.Type) {
+ continue
+ }
plain := s.crypto.Decrypt(r.Config)
var cfg map[string]any
_ = json.Unmarshal([]byte(plain), &cfg)
@@ -334,7 +317,7 @@ func (s *StorageConfigService) Test(ctx context.Context, in StorageInput) error
}
defer resp.Body.Close()
return nil
- case cloud.TypeQuark, cloud.Type115, cloud.TypeCloudDrive2:
+ case cloud.Type115, cloud.TypeCloudDrive2:
p, err := cloud.New(in.Type, cfg, client)
if err != nil {
return err
@@ -373,287 +356,40 @@ func (s *StorageConfigService) CloudList(ctx context.Context, typ, dirID string)
return p.List(ctx, dirID)
}
-// CloudResolve resolves a cloud file reference to a direct link.
-//
-// clientUA is the User-Agent of the playback client that will follow the 302
-// redirect. 115/夸克 CDN links are bound to the UA used to request them, so we
-// resolve with the client's own UA — that way the pure 302 the host issues
-// points at a link the client can fetch directly (true offload). When clientUA
-// is empty the provider's default UA is used.
-func (s *StorageConfigService) CloudResolve(ctx context.Context, typ, fileRef, clientUA string) (*cloud.DirectLink, error) {
- if s == nil {
- return nil, errors.New("storage config service unavailable")
- }
- cacheKey := s.resolveCacheKey(typ, fileRef, clientUA)
- if link, ok, refresh := s.cachedResolve(cacheKey, typ); ok {
- if refresh {
- s.refreshResolveInBackground(cacheKey, typ, fileRef, clientUA)
- }
- return link, nil
- }
- if call, owner := s.beginResolve(cacheKey); !owner {
- select {
- case <-call.done:
- if call.err != nil {
- return nil, call.err
- }
- return cloneDirectLink(call.link), nil
- case <-ctx.Done():
- return nil, ctx.Err()
- }
- } else {
- defer s.finishResolve(cacheKey, call)
- p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
- if err != nil {
- call.err = err
- return nil, err
- }
- link, err := p.Resolve(ctx, fileRef)
- if err != nil {
- call.err = err
- return nil, err
- }
- call.link = cloneDirectLink(link)
- s.storeResolvedLink(cacheKey, typ, link)
- return cloneDirectLink(link), nil
- }
-}
-
-func (s *StorageConfigService) resolveCacheKey(typ, fileRef, clientUA string) string {
- return strings.TrimSpace(typ) + "\x00" + strings.TrimSpace(fileRef) + "\x00" + strings.TrimSpace(clientUA)
-}
-
-func (s *StorageConfigService) cachedResolve(key, typ string) (*cloud.DirectLink, bool, bool) {
- s.resolveMu.Lock()
- defer s.resolveMu.Unlock()
- if s.resolveCache == nil {
- s.resolveCache = make(map[string]cloudResolveCacheEntry)
- return nil, false, false
- }
- entry, ok := s.resolveCache[key]
- now := time.Now()
- if !ok || now.After(entry.expiresAt) {
- if ok {
- delete(s.resolveCache, key)
- }
- return nil, false, false
- }
- entry.hits++
- entry.lastHit = now
- s.resolveCache[key] = entry
- refreshWindow := cloudResolveHotRefreshWindow(cloudResolveCacheTTL(typ))
- shouldRefresh := entry.hits >= cloudResolveHotHitThreshold &&
- refreshWindow > 0 &&
- now.Add(refreshWindow).After(entry.expiresAt)
- return cloneDirectLink(entry.link), true, shouldRefresh
-}
-
-func (s *StorageConfigService) beginResolve(key string) (*cloudResolveCall, bool) {
- s.resolveMu.Lock()
- defer s.resolveMu.Unlock()
- if s.resolveFlight == nil {
- s.resolveFlight = make(map[string]*cloudResolveCall)
- }
- if call := s.resolveFlight[key]; call != nil {
- return call, false
- }
- call := &cloudResolveCall{done: make(chan struct{})}
- s.resolveFlight[key] = call
- return call, true
-}
-
-func (s *StorageConfigService) finishResolve(key string, call *cloudResolveCall) {
- s.resolveMu.Lock()
- if current := s.resolveFlight[key]; current == call {
- delete(s.resolveFlight, key)
- }
- s.resolveMu.Unlock()
- close(call.done)
-}
-
-func (s *StorageConfigService) refreshResolveInBackground(key, typ, fileRef, clientUA string) {
- if s == nil {
- return
- }
- go func() {
- call, owner := s.beginResolve(key)
- if !owner {
- return
- }
- defer s.finishResolve(key, call)
- ctx, cancel := context.WithTimeout(context.Background(), cloudResolveBackgroundRefreshMax)
- defer cancel()
- p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
- if err != nil {
- call.err = err
- if s.log != nil {
- s.log.Debug("refresh cloud direct link failed", zap.String("provider", typ), zap.Error(err))
- }
- return
- }
- link, err := p.Resolve(ctx, fileRef)
- if err != nil {
- call.err = err
- if s.log != nil {
- s.log.Debug("refresh cloud direct link failed", zap.String("provider", typ), zap.Error(err))
- }
- return
- }
- call.link = cloneDirectLink(link)
- s.storeResolvedLink(key, typ, link)
- }()
-}
-
-func (s *StorageConfigService) storeResolvedLink(key, typ string, link *cloud.DirectLink) {
- if link == nil || strings.TrimSpace(link.URL) == "" {
- return
- }
- ttl := cloudResolveCacheTTL(typ)
- if ttl <= 0 {
- return
- }
- s.resolveMu.Lock()
- defer s.resolveMu.Unlock()
- if s.resolveCache == nil {
- s.resolveCache = make(map[string]cloudResolveCacheEntry)
- }
- now := time.Now()
- hits := 0
- if existing, ok := s.resolveCache[key]; ok {
- hits = existing.hits
- }
- s.resolveCache[key] = cloudResolveCacheEntry{link: cloneDirectLink(link), expiresAt: now.Add(ttl), hits: hits, lastHit: now}
-}
-
-func cloudResolveHotRefreshWindow(ttl time.Duration) time.Duration {
- if ttl <= 0 {
- return 0
- }
- window := ttl / 4
- if window < 15*time.Second {
- window = 15 * time.Second
- }
- if window > 2*time.Minute {
- window = 2 * time.Minute
- }
- return window
-}
-
-func cloudResolveCacheTTL(typ string) time.Duration {
- switch typ {
- case cloud.TypeQuark, cloud.Type115, cloud.TypeCloudDrive2, cloud.TypeOpenList:
- return 2 * time.Minute
- default:
- return 5 * time.Minute
- }
-}
-
-func cloneDirectLink(link *cloud.DirectLink) *cloud.DirectLink {
- if link == nil {
- return nil
- }
- out := &cloud.DirectLink{
- URL: link.URL,
- Headers: make(map[string]string, len(link.Headers)),
- Proxy: link.Proxy,
- }
- for k, v := range link.Headers {
- out.Headers[k] = v
- }
- return out
-}
-
-func (s *StorageConfigService) clearResolveCacheForType(typ string) {
- typ = strings.TrimSpace(typ)
- if typ == "" {
- return
- }
- prefix := typ + "\x00"
- s.resolveMu.Lock()
- defer s.resolveMu.Unlock()
- for key := range s.resolveCache {
- if strings.HasPrefix(key, prefix) {
- delete(s.resolveCache, key)
- }
- }
- for key, call := range s.resolveFlight {
- if strings.HasPrefix(key, prefix) && call != nil {
- call.err = fmt.Errorf("%s storage config changed", typ)
- }
- }
-}
-
-func (s *StorageConfigService) CloudResolveUncached(ctx context.Context, typ, fileRef, clientUA string) (*cloud.DirectLink, error) {
- p, err := s.cloudProviderWithUA(ctx, typ, clientUA)
+func (s *StorageConfigService) CloudMkdir(ctx context.Context, typ, parentDir, name string) (*cloud.FileEntry, error) {
+ p, err := s.CloudProvider(ctx, typ)
if err != nil {
return nil, err
}
- return p.Resolve(ctx, fileRef)
+ mutable, ok := p.(cloud.MutableProvider)
+ if !ok {
+ return nil, fmt.Errorf("%s does not support folder creation", typ)
+ }
+ return mutable.Mkdir(ctx, parentDir, name)
}
-// CloudReadText resolves a small cloud file and returns its text payload. It is
-// used for cloud-hosted .strm files: the scanner reads the STRM target once and
-// stores the real playback URL, while the media bytes still stay in the cloud.
-func (s *StorageConfigService) CloudReadText(ctx context.Context, typ, fileRef string, limit int64) (string, error) {
- if limit <= 0 {
- limit = 64 << 10
- }
- link, err := s.CloudResolve(ctx, typ, fileRef, "")
- if err != nil {
- return "", err
- }
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, link.URL, nil)
- if err != nil {
- return "", err
- }
- for k, v := range link.Headers {
- req.Header.Set(k, v)
- }
- resp, err := s.client.Do(req)
- if err != nil {
- return "", err
- }
- defer resp.Body.Close()
- if resp.StatusCode < 200 || resp.StatusCode >= 300 {
- return "", fmt.Errorf("%s: read strm returned http %d", typ, resp.StatusCode)
- }
- body, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
- if err != nil {
- return "", err
- }
- if int64(len(body)) > limit {
- return "", fmt.Errorf("%s: strm file is too large", typ)
- }
- return strings.TrimSpace(strings.TrimPrefix(string(body), "\ufeff")), nil
-}
-
-// cloudProviderWithUA builds a provider, overriding the request UA when a
-// non-empty clientUA is supplied.
-func (s *StorageConfigService) cloudProviderWithUA(ctx context.Context, typ, clientUA string) (cloud.Provider, error) {
- if !cloud.IsCloudType(typ) {
- return nil, fmt.Errorf("not a cloud provider: %q", typ)
- }
- view, err := s.Get(ctx, typ)
+func (s *StorageConfigService) CloudRename(ctx context.Context, typ, ref, name string) (*cloud.FileEntry, error) {
+ p, err := s.CloudProvider(ctx, typ)
if err != nil {
return nil, err
}
- if view == nil {
- return nil, fmt.Errorf("%s storage not configured", typ)
+ mutable, ok := p.(cloud.MutableProvider)
+ if !ok {
+ return nil, fmt.Errorf("%s does not support rename", typ)
}
- if !view.Enabled {
- return nil, fmt.Errorf("%s storage disabled", typ)
+ return mutable.Rename(ctx, ref, name)
+}
+
+func (s *StorageConfigService) CloudMove(ctx context.Context, typ, ref, targetDir, name string) (*cloud.FileEntry, error) {
+ p, err := s.CloudProvider(ctx, typ)
+ if err != nil {
+ return nil, err
}
- cfg := view.Config
- if strings.TrimSpace(clientUA) != "" {
- // Copy so we never mutate the cached view config.
- cp := make(map[string]any, len(cfg)+1)
- for k, v := range cfg {
- cp[k] = v
- }
- cp["ua"] = clientUA
- cfg = cp
+ movable, ok := p.(cloud.MovableProvider)
+ if !ok {
+ return nil, fmt.Errorf("%s does not support move", typ)
}
- return cloud.New(typ, cfg, s.clientForConfig(cfg))
+ return movable.Move(ctx, ref, targetDir, name)
}
func (s *StorageConfigService) clientForConfig(cfg map[string]any) *http.Client {
@@ -704,8 +440,6 @@ func storageTimeoutFromConfig(cfg map[string]any, fallback time.Duration) time.D
// cloudLibraryName maps a provider type to a friendly Chinese library name.
func cloudLibraryName(typ string) string {
switch typ {
- case cloud.TypeQuark:
- return "夸克网盘"
case cloud.Type115:
return "115 网盘"
case cloud.TypeCloudDrive2:
@@ -766,7 +500,7 @@ func (s *StorageConfigService) CloudImport(ctx context.Context, typ, fileRef, na
Path: cloudMediaPath(typ, fileRef),
SizeBytes: size,
Container: container,
- STRMURL: BuildPublicAPIURL(ctx, s.repo, nil, "/api/cloud/play/"+typ, url.Values{"ref": []string{fileRef}}),
+ STRMURL: BuildRelativeCloudPlayURL(typ, fileRef),
ScrapeStatus: "pending",
}
if err := s.repo.Media.Upsert(ctx, m); err != nil {
@@ -777,7 +511,7 @@ func (s *StorageConfigService) CloudImport(ctx context.Context, typ, fileRef, na
func validStorageType(t string) bool {
switch t {
- case "alist", "s3", "webdav", cloud.TypeQuark, cloud.Type115, cloud.TypeCloudDrive2, cloud.TypeOpenList:
+ case "alist", "s3", "webdav", cloud.Type115, cloud.TypeCloudDrive2, cloud.TypeOpenList:
return true
}
return false
diff --git a/internal/service/storage_config_cache_test.go b/internal/service/storage_config_cache_test.go
index 2260cb5..37a5e4b 100644
--- a/internal/service/storage_config_cache_test.go
+++ b/internal/service/storage_config_cache_test.go
@@ -12,27 +12,27 @@ import (
func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
var resolves atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/file/download" {
+ if r.URL.Path != "/api/fs/get" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
n := resolves.Add(1)
w.Header().Set("Content-Type", "application/json")
- _, _ = fmt.Fprintf(w, `{"status":200,"code":0,"data":[{"fid":"f1","download_url":"http://cdn.local/%d.mkv"}]}`, n)
+ _, _ = fmt.Fprintf(w, `{"code":200,"data":{"raw_url":"http://cdn.local/%d.mkv"}}`, n)
}))
defer upstream.Close()
_, storage := newStorageUploadTestService(t)
if _, err := storage.Save(t.Context(), StorageInput{
- Type: "quark",
+ Type: "openlist",
Config: map[string]any{
- "cookie": "kps=test",
- "base": upstream.URL,
+ "server": upstream.URL,
+ "token": "token",
},
}); err != nil {
t.Fatal(err)
}
- link, err := storage.CloudResolve(t.Context(), "quark", "f1", "Player/1")
+ link, err := storage.CloudResolve(t.Context(), "openlist", "/Movies/f1.mkv", "Player/1")
if err != nil {
t.Fatal(err)
}
@@ -40,7 +40,7 @@ func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
t.Fatalf("first resolve link=%#v resolves=%d", link, resolves.Load())
}
for i := 0; i < cloudResolveHotHitThreshold-1; i++ {
- link, err = storage.CloudResolve(t.Context(), "quark", "f1", "Player/1")
+ link, err = storage.CloudResolve(t.Context(), "openlist", "/Movies/f1.mkv", "Player/1")
if err != nil {
t.Fatal(err)
}
@@ -49,7 +49,7 @@ func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
}
}
- key := storage.resolveCacheKey("quark", "f1", "Player/1")
+ key := storage.resolveCacheKey("openlist", "/Movies/f1.mkv", "Player/1")
storage.resolveMu.Lock()
entry := storage.resolveCache[key]
entry.hits = cloudResolveHotHitThreshold
@@ -57,7 +57,7 @@ func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
storage.resolveCache[key] = entry
storage.resolveMu.Unlock()
- link, err = storage.CloudResolve(t.Context(), "quark", "f1", "Player/1")
+ link, err = storage.CloudResolve(t.Context(), "openlist", "/Movies/f1.mkv", "Player/1")
if err != nil {
t.Fatal(err)
}
@@ -71,7 +71,7 @@ func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
if resolves.Load() < 2 {
t.Fatalf("background refresh did not run, resolves=%d", resolves.Load())
}
- link, err = storage.CloudResolve(t.Context(), "quark", "f1", "Player/1")
+ link, err = storage.CloudResolve(t.Context(), "openlist", "/Movies/f1.mkv", "Player/1")
if err != nil {
t.Fatal(err)
}
@@ -81,7 +81,7 @@ func TestCloudResolveHotCacheRefreshesInBackground(t *testing.T) {
}
func TestCloudResolveCacheTTLUsesShortTTLForCloudPlaybackLinks(t *testing.T) {
- for _, typ := range []string{"quark", "cloud115", "clouddrive2", "openlist"} {
+ for _, typ := range []string{"cloud115", "clouddrive2", "openlist"} {
if got := cloudResolveCacheTTL(typ); got != 2*time.Minute {
t.Fatalf("%s cloud resolve cache ttl = %v, want 2m", typ, got)
}
diff --git a/internal/service/storage_config_logout_test.go b/internal/service/storage_config_logout_test.go
index 468203f..2b390a5 100644
--- a/internal/service/storage_config_logout_test.go
+++ b/internal/service/storage_config_logout_test.go
@@ -72,3 +72,31 @@ func TestStorageConfigLogoutClearsCredentialsAndCloudLibraries(t *testing.T) {
t.Fatalf("cloud media should be purged after logout, count=%d", cloudMediaCount)
}
}
+
+func TestStorageConfigListHidesDeprecatedQuarkRows(t *testing.T) {
+ repos, storage := newStorageUploadTestService(t)
+ if err := repos.StorageConfig.Upsert(t.Context(), &model.StorageConfig{
+ Type: LegacyQuarkProvider,
+ Config: storage.crypto.Encrypt(`{"cookie":"legacy"}`),
+ Enabled: true,
+ }); err != nil {
+ t.Fatalf("insert legacy quark row: %v", err)
+ }
+ if _, err := storage.Save(t.Context(), StorageInput{
+ Type: "openlist",
+ Config: map[string]any{
+ "server": "http://openlist.test",
+ "token": "token",
+ },
+ }); err != nil {
+ t.Fatalf("save openlist row: %v", err)
+ }
+
+ rows, err := storage.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 1 || rows[0].Type != "openlist" {
+ t.Fatalf("storage list = %#v, want only supported OpenList row", rows)
+ }
+}
diff --git a/internal/service/storage_types.go b/internal/service/storage_types.go
new file mode 100644
index 0000000..72ef2a9
--- /dev/null
+++ b/internal/service/storage_types.go
@@ -0,0 +1,31 @@
+package service
+
+import (
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
+)
+
+const LegacyQuarkProvider = "quark"
+
+func IsAdminStorageConfigurable(typ string) bool {
+ switch strings.TrimSpace(typ) {
+ case cloud.TypeOpenList, "alist", "webdav", cloud.TypeCloudDrive2, cloud.Type115:
+ return true
+ default:
+ return false
+ }
+}
+
+func IsAdminCloudConfigurable(typ string) bool {
+ switch strings.TrimSpace(typ) {
+ case cloud.Type115, cloud.TypeCloudDrive2, cloud.TypeOpenList:
+ return true
+ default:
+ return false
+ }
+}
+
+func IsDeprecatedNativeCloudProvider(typ string) bool {
+ return strings.TrimSpace(typ) == LegacyQuarkProvider
+}
diff --git a/internal/service/storage_types_test.go b/internal/service/storage_types_test.go
new file mode 100644
index 0000000..4a3643e
--- /dev/null
+++ b/internal/service/storage_types_test.go
@@ -0,0 +1,29 @@
+package service
+
+import "testing"
+
+func TestAdminStorageConfigurableTypes(t *testing.T) {
+ for _, typ := range []string{"openlist", "alist", "webdav", "clouddrive2", "cloud115"} {
+ if !IsAdminStorageConfigurable(typ) {
+ t.Fatalf("%s should be configurable", typ)
+ }
+ }
+ for _, typ := range []string{"quark", "s3", "", "unknown"} {
+ if IsAdminStorageConfigurable(typ) {
+ t.Fatalf("%s should not be configurable", typ)
+ }
+ }
+}
+
+func TestAdminCloudConfigurableTypes(t *testing.T) {
+ for _, typ := range []string{"openlist", "clouddrive2", "cloud115"} {
+ if !IsAdminCloudConfigurable(typ) {
+ t.Fatalf("%s should be cloud-configurable", typ)
+ }
+ }
+ for _, typ := range []string{"quark", "alist", "webdav", "s3", ""} {
+ if IsAdminCloudConfigurable(typ) {
+ t.Fatalf("%s should not be cloud-configurable", typ)
+ }
+ }
+}
diff --git a/internal/service/storage_upload.go b/internal/service/storage_upload.go
index 27af88e..6fe634d 100644
--- a/internal/service/storage_upload.go
+++ b/internal/service/storage_upload.go
@@ -25,7 +25,7 @@ const (
CloudUploadOverwriteKey = "cloud.upload_overwrite"
CloudUploadTransferModeKey = "cloud.upload_transfer_mode"
CloudUploadIntervalSecondsKey = "cloud.upload_interval_seconds"
- CloudUploadUnsupportedProvider = "本地文件直传目前支持 Alist / OpenList / WebDAV / CloudDrive2;115/夸克原生上传需要各自的分片上传私有接口,建议先用 CloudDrive2、OpenList 或 Alist 桥接后转存。"
+ CloudUploadUnsupportedProvider = "本地文件直传目前支持 Alist / OpenList / WebDAV / CloudDrive2;115 原生上传需要分片上传私有接口,建议先用 CloudDrive2、OpenList 或 Alist 桥接后转存。"
)
type CloudUploadInput struct {
@@ -165,7 +165,7 @@ func (s *StorageConfigService) uploaderForView(typ string, view *StorageView) (s
return newWebDAVUploader(view.Config), nil
case "s3":
return nil, errors.New("s3 local upload is not implemented yet")
- case "cloud115", "quark":
+ case "cloud115":
return nil, errors.New(CloudUploadUnsupportedProvider)
default:
return nil, fmt.Errorf("unsupported storage type %q", typ)
diff --git a/internal/service/storage_upload_test.go b/internal/service/storage_upload_test.go
index 73a2fc0..cbc433b 100644
--- a/internal/service/storage_upload_test.go
+++ b/internal/service/storage_upload_test.go
@@ -11,9 +11,7 @@ import (
"strings"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -447,13 +445,7 @@ func TestStorageConfigUploadLocalMoveDeletesSourceAfterUpload(t *testing.T) {
func newStorageUploadTestService(t *testing.T) (*repository.Container, *StorageConfigService) {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.StorageConfig{}, &model.Setting{}, &model.Library{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.StorageConfig{}, &model.Setting{}, &model.Library{}, &model.Media{})
repos := repository.New(db)
log := zap.NewNop()
return repos, NewStorageConfigService(log, repos, NewCryptoService("", log))
diff --git a/internal/service/stream_test.go b/internal/service/stream_test.go
index ebf5cbc..b2ba69a 100644
--- a/internal/service/stream_test.go
+++ b/internal/service/stream_test.go
@@ -8,9 +8,7 @@ import (
"strings"
"testing"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -223,13 +221,7 @@ func TestServeFileRejectsCloudMediaWhenSelectedModeDisabled(t *testing.T) {
func newStreamTestRepo(t *testing.T) *repository.Container {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Media{}, &model.Setting{})
return repository.New(db)
}
diff --git a/internal/service/strm_generate.go b/internal/service/strm_generate.go
new file mode 100644
index 0000000..fb3d0bf
--- /dev/null
+++ b/internal/service/strm_generate.go
@@ -0,0 +1,254 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "os"
+ "path/filepath"
+ "strconv"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+type GenerateSTRMOptions struct {
+ LibraryID string `json:"library_id"`
+ OutputDir string `json:"output_dir"`
+ BaseURL string `json:"base_url,omitempty"`
+ Enabled bool `json:"enabled"`
+ Overwrite bool `json:"overwrite"`
+ IncludeLocal bool `json:"include_local"`
+ PlaybackToken string `json:"-"`
+}
+
+type GenerateSTRMResult struct {
+ LibraryID string `json:"library_id"`
+ OutputDir string `json:"output_dir"`
+ Generated int `json:"generated"`
+ Updated int `json:"updated"`
+ Skipped int `json:"skipped"`
+ Cleaned int `json:"cleaned"`
+ Errors []string `json:"errors,omitempty"`
+ Items []GenerateSTRMItem `json:"items,omitempty"`
+}
+
+type GenerateSTRMItem struct {
+ MediaID string `json:"media_id"`
+ Title string `json:"title"`
+ FilePath string `json:"file_path"`
+ URL string `json:"url,omitempty"`
+ Action string `json:"action"`
+ Reason string `json:"reason,omitempty"`
+}
+
+func (s *STRMService) GenerateForLibrary(ctx context.Context, opts GenerateSTRMOptions) (*GenerateSTRMResult, error) {
+ if s == nil || s.repo == nil || s.repo.DB == nil {
+ return nil, errors.New("strm service unavailable")
+ }
+ libraryID := strings.TrimSpace(opts.LibraryID)
+ if libraryID == "" {
+ return nil, errors.New("library_id required")
+ }
+ lib, err := s.repo.Library.FindByID(ctx, libraryID)
+ if err != nil {
+ return nil, err
+ }
+ if lib == nil {
+ return nil, errors.New("library not found")
+ }
+ outputDir := s.resolveSTRMOutputDir(ctx, lib, opts)
+ if outputDir == "" || outputDir == "." {
+ return nil, errors.New("output_dir required")
+ }
+ s.saveSTRMGenerationSettings(ctx, outputDir, opts)
+ if err := os.MkdirAll(outputDir, 0o755); err != nil { // #nosec G301 -- STRM output directories must stay readable by NAS/player users.
+ return nil, err
+ }
+
+ rows, err := s.librarySTRMMedia(ctx, libraryID)
+ if err != nil {
+ return nil, err
+ }
+ res := &GenerateSTRMResult{LibraryID: libraryID, OutputDir: outputDir}
+ expectedFiles := map[string]struct{}{}
+ for _, media := range rows {
+ select {
+ case <-ctx.Done():
+ return res, ctx.Err()
+ default:
+ }
+ item := s.generateOne(ctx, *lib, media, outputDir, opts)
+ res.addItem(item)
+ if item.FilePath != "" && item.Action != "error" {
+ expectedFiles[filepath.Clean(item.FilePath)] = struct{}{}
+ }
+ }
+ if opts.Overwrite {
+ cleaned, err := s.cleanupStaleGeneratedSTRM(ctx, outputDir, expectedFiles)
+ if err != nil {
+ res.Errors = append(res.Errors, err.Error())
+ }
+ res.Cleaned += cleaned
+ }
+ return res, nil
+}
+
+func (s *STRMService) GenerateForAllLibraries(ctx context.Context, opts GenerateSTRMOptions) (*GenerateSTRMResult, error) {
+ if s == nil || s.repo == nil || s.repo.Library == nil {
+ return nil, errors.New("strm service unavailable")
+ }
+ libraries, err := s.repo.Library.List(ctx)
+ if err != nil {
+ return nil, err
+ }
+ baseOutputDir := resolveMappedDestinationPath(strings.TrimSpace(opts.OutputDir))
+ result := &GenerateSTRMResult{LibraryID: "*", OutputDir: baseOutputDir}
+ for _, lib := range libraries {
+ select {
+ case <-ctx.Done():
+ return result, ctx.Err()
+ default:
+ }
+ next := opts
+ next.LibraryID = lib.ID
+ if baseOutputDir != "" && baseOutputDir != "." {
+ next.OutputDir = filepath.Join(baseOutputDir, strmLibraryOutputSubdir(lib))
+ }
+ part, err := s.GenerateForLibrary(ctx, next)
+ if err != nil {
+ result.Errors = append(result.Errors, fmt.Sprintf("%s: %v", lib.Name, err))
+ continue
+ }
+ result.merge(part)
+ }
+ if baseOutputDir != "" && baseOutputDir != "." && s.repo.Setting != nil {
+ _ = s.repo.Setting.Set(ctx, "strm.output_dir", baseOutputDir)
+ _ = s.repo.Setting.Set(ctx, "strm.output_scope", "all")
+ result.OutputDir = baseOutputDir
+ }
+ return result, nil
+}
+
+func (s *STRMService) resolveSTRMOutputDir(ctx context.Context, lib *model.Library, opts GenerateSTRMOptions) string {
+ outputDir := resolveMappedDestinationPath(strings.TrimSpace(opts.OutputDir))
+ if (outputDir == "" || outputDir == ".") && s.repo.Setting != nil {
+ if saved, err := s.repo.Setting.Get(ctx, "strm.output_dir"); err == nil {
+ outputDir = resolveMappedDestinationPath(strings.TrimSpace(saved))
+ }
+ }
+ if outputDir == "" || outputDir == "." {
+ outputDir = s.defaultOutputDir(lib)
+ }
+ return outputDir
+}
+
+func (s *STRMService) saveSTRMGenerationSettings(ctx context.Context, outputDir string, opts GenerateSTRMOptions) {
+ if strings.TrimSpace(opts.BaseURL) != "" && s.repo.Setting != nil {
+ baseURL := strings.TrimRight(strings.TrimSpace(opts.BaseURL), "/")
+ _ = s.repo.Setting.Set(ctx, "app.server_url", baseURL)
+ _ = s.repo.Setting.Set(ctx, "strm.base_url", baseURL)
+ }
+ if s.repo.Setting == nil {
+ return
+ }
+ _ = s.repo.Setting.Set(ctx, "strm.auto_generate_enabled", strconv.FormatBool(opts.Enabled))
+ _ = s.repo.Setting.Set(ctx, "strm.output_dir", outputDir)
+ _ = s.repo.Setting.Set(ctx, "strm.output_scope", "library")
+}
+
+func (s *STRMService) librarySTRMMedia(ctx context.Context, libraryID string) ([]model.Media, error) {
+ var rows []model.Media
+ err := s.repo.DB.WithContext(ctx).
+ Where("library_id = ?", libraryID).
+ Order("title asc, season_num asc, episode_num asc, created_at asc").
+ Find(&rows).Error
+ return rows, err
+}
+
+func (s *STRMService) defaultOutputDir(lib *model.Library) string {
+ subdir := strmLibraryOutputSubdir(*lib)
+ if s != nil && s.cfg != nil && strings.TrimSpace(s.cfg.App.DataDir) != "" {
+ return filepath.Join(s.cfg.App.DataDir, "strm", subdir)
+ }
+ return filepath.Join("data", "strm", subdir)
+}
+
+func (s *STRMService) generateOne(ctx context.Context, lib model.Library, media model.Media, outputDir string, opts GenerateSTRMOptions) GenerateSTRMItem {
+ item := GenerateSTRMItem{MediaID: media.ID, Title: media.Title}
+ playURL := s.strmPlaybackURL(ctx, media, opts.BaseURL, opts.PlaybackToken)
+ if playURL == "" {
+ item.Action = "skipped"
+ item.Reason = "no playable strm target"
+ return item
+ }
+ if strings.TrimSpace(media.STRMURL) == "" && !opts.IncludeLocal {
+ item.Action = "skipped"
+ item.Reason = "local media skipped"
+ return item
+ }
+ rel := s.strmRelativePath(lib, media)
+ if rel == "" {
+ item.Action = "skipped"
+ item.Reason = "cannot build file name"
+ return item
+ }
+ filePath := filepath.Join(outputDir, rel)
+ item.FilePath = filePath
+ item.URL = playURL
+ if _, err := os.Stat(filePath); err == nil && !opts.Overwrite {
+ item.Action = "skipped"
+ item.Reason = "target exists"
+ return item
+ }
+ action := "generated"
+ if _, err := os.Stat(filePath); err == nil {
+ action = "updated"
+ }
+ if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil { // #nosec G301 -- STRM output directories must stay readable by NAS/player users.
+ item.Action = "error"
+ item.Reason = err.Error()
+ return item
+ }
+ if err := os.WriteFile(filePath, []byte(playURL+"\n"), 0o644); err != nil { // #nosec G306 -- STRM files are media sidecars intended to be readable by players.
+ item.Action = "error"
+ item.Reason = err.Error()
+ return item
+ }
+ if err := s.upsertGeneratedRecord(ctx, media, filePath, playURL, lib.Type); err != nil {
+ item.Action = "error"
+ item.Reason = err.Error()
+ return item
+ }
+ item.Action = action
+ return item
+}
+
+func (r *GenerateSTRMResult) addItem(item GenerateSTRMItem) {
+ r.Items = append(r.Items, item)
+ switch item.Action {
+ case "generated":
+ r.Generated++
+ case "updated":
+ r.Updated++
+ case "skipped":
+ r.Skipped++
+ case "error":
+ r.Errors = append(r.Errors, fmt.Sprintf("%s: %s", item.Title, item.Reason))
+ }
+}
+
+func (r *GenerateSTRMResult) merge(part *GenerateSTRMResult) {
+ if part == nil {
+ return
+ }
+ if r.OutputDir == "" || r.OutputDir == "." {
+ r.OutputDir = filepath.Dir(part.OutputDir)
+ }
+ r.Generated += part.Generated
+ r.Updated += part.Updated
+ r.Skipped += part.Skipped
+ r.Cleaned += part.Cleaned
+ r.Errors = append(r.Errors, part.Errors...)
+ r.Items = append(r.Items, part.Items...)
+}
diff --git a/internal/service/strm_generate_cleanup.go b/internal/service/strm_generate_cleanup.go
new file mode 100644
index 0000000..9cca145
--- /dev/null
+++ b/internal/service/strm_generate_cleanup.go
@@ -0,0 +1,116 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "io/fs"
+ "net/url"
+ "os"
+ "path/filepath"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *STRMService) upsertGeneratedRecord(ctx context.Context, media model.Media, filePath, playURL, mediaType string) error {
+ protocol := ""
+ if u, err := url.Parse(playURL); err == nil {
+ protocol = strings.ToLower(u.Scheme)
+ }
+ if protocol == "" {
+ protocol = "http"
+ }
+ record := model.STRMRecord{
+ Title: media.Title,
+ URL: playURL,
+ FilePath: filePath,
+ Protocol: protocol,
+ MediaID: media.ID,
+ MediaType: mediaType,
+ SeasonNum: media.SeasonNum,
+ EpisodeNum: media.EpisodeNum,
+ }
+ var existing model.STRMRecord
+ err := s.repo.DB.WithContext(ctx).Where("media_id = ? AND file_path = ?", media.ID, filePath).First(&existing).Error
+ if err == nil {
+ existing.Title = record.Title
+ existing.URL = record.URL
+ existing.Protocol = record.Protocol
+ existing.MediaType = record.MediaType
+ existing.SeasonNum = record.SeasonNum
+ existing.EpisodeNum = record.EpisodeNum
+ return s.repo.DB.WithContext(ctx).Save(&existing).Error
+ }
+ return s.repo.DB.WithContext(ctx).Create(&record).Error
+}
+
+func (s *STRMService) cleanupStaleGeneratedSTRM(ctx context.Context, outputDir string, expected map[string]struct{}) (int, error) {
+ outputDir = filepath.Clean(strings.TrimSpace(outputDir))
+ if outputDir == "" || outputDir == "." {
+ return 0, nil
+ }
+ cleaned, err := removeStaleSTRMFiles(outputDir, expected)
+ if err != nil {
+ return cleaned, err
+ }
+ recordsCleaned, err := s.removeStaleSTRMRecords(ctx, outputDir, expected)
+ return cleaned + recordsCleaned, err
+}
+
+func removeStaleSTRMFiles(outputDir string, expected map[string]struct{}) (int, error) {
+ cleaned := 0
+ err := filepath.WalkDir(outputDir, func(path string, entry fs.DirEntry, walkErr error) error {
+ if walkErr != nil {
+ return nil
+ }
+ if entry.IsDir() || strings.ToLower(filepath.Ext(path)) != ".strm" {
+ return nil
+ }
+ cleanPath := filepath.Clean(path)
+ if _, ok := expected[cleanPath]; ok {
+ return nil
+ }
+ if err := os.Remove(cleanPath); err != nil && !errors.Is(err, os.ErrNotExist) {
+ return err
+ }
+ cleaned++
+ return nil
+ })
+ if err != nil && !errors.Is(err, os.ErrNotExist) {
+ return cleaned, err
+ }
+ return cleaned, nil
+}
+
+func (s *STRMService) removeStaleSTRMRecords(ctx context.Context, outputDir string, expected map[string]struct{}) (int, error) {
+ if s == nil || s.repo == nil || s.repo.DB == nil {
+ return 0, nil
+ }
+ var records []model.STRMRecord
+ if err := s.repo.DB.WithContext(ctx).Find(&records).Error; err != nil {
+ return 0, err
+ }
+ rootAbs, err := filepath.Abs(outputDir)
+ if err != nil {
+ return 0, nil
+ }
+ cleaned := 0
+ for i := range records {
+ filePath := filepath.Clean(strings.TrimSpace(records[i].FilePath))
+ if filePath == "" {
+ continue
+ }
+ fileAbs, err := filepath.Abs(filePath)
+ if err != nil || !pathWithin(fileAbs, rootAbs) {
+ continue
+ }
+ if _, ok := expected[filePath]; ok {
+ continue
+ }
+ if err := s.repo.DB.WithContext(ctx).Delete(&records[i]).Error; err != nil {
+ return cleaned, err
+ }
+ cleaned++
+ }
+ return cleaned, nil
+}
diff --git a/internal/service/strm_output_dir.go b/internal/service/strm_output_dir.go
new file mode 100644
index 0000000..8297038
--- /dev/null
+++ b/internal/service/strm_output_dir.go
@@ -0,0 +1,122 @@
+package service
+
+import (
+ "path/filepath"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func strmLibraryOutputSubdir(lib model.Library) string {
+ parts := strmLibraryCategoryParts(lib)
+ if len(parts) == 0 {
+ return sanitizeFilename(lib.Name)
+ }
+ clean := make([]string, 0, len(parts))
+ for _, part := range parts {
+ if safe := sanitizeFilename(part); safe != "" {
+ clean = append(clean, safe)
+ }
+ }
+ if len(clean) == 0 {
+ return sanitizeFilename(lib.Name)
+ }
+ return filepath.Join(clean...)
+}
+
+func strmLibraryCategoryParts(lib model.Library) []string {
+ if parts := strmCategoryPartsFromPath(strmLibraryPathParts(lib.Path)); len(parts) > 0 {
+ return parts
+ }
+ if parts := strmCategoryPartsFromPath(strmNameParts(lib.Name)); len(parts) > 0 {
+ return parts
+ }
+ if root := mediaTypeRootDir(lib.Type); root != "" {
+ return []string{root}
+ }
+ return nil
+}
+
+func strmLibraryPathParts(raw string) []string {
+ if info, ok := ParseCloudLibraryMount(raw); ok {
+ return strmSlashParts(info.DisplayDir)
+ }
+ clean := cleanPathForVolumeMapping(raw)
+ clean = strings.Trim(pathAfterWindowsDrivePrefix(clean), "/")
+ return strmSlashParts(clean)
+}
+
+func strmNameParts(name string) []string {
+ name = strings.NewReplacer("·", "/", ">", "/", "|", "/", "|", "/", "\\", "/").Replace(name)
+ return strmSlashParts(name)
+}
+
+func strmSlashParts(raw string) []string {
+ raw = strings.Trim(strings.TrimSpace(strings.ReplaceAll(raw, "\\", "/")), "/")
+ if raw == "" || raw == "." {
+ return nil
+ }
+ fields := strings.Split(raw, "/")
+ parts := make([]string, 0, len(fields))
+ for _, part := range fields {
+ part = strings.TrimSpace(part)
+ if part != "" && part != "." {
+ parts = append(parts, part)
+ }
+ }
+ return parts
+}
+
+func strmCategoryPartsFromPath(parts []string) []string {
+ for i, part := range parts {
+ if root := strmCanonicalRoot(part); root != "" {
+ return append([]string{root}, strmSanitizedTail(parts[i+1:])...)
+ }
+ if root := strmCategoryRoot(part); root != "" {
+ return []string{root, part}
+ }
+ }
+ return nil
+}
+
+func strmSanitizedTail(parts []string) []string {
+ out := make([]string, 0, len(parts))
+ for _, part := range parts {
+ if strings.TrimSpace(part) != "" {
+ out = append(out, part)
+ }
+ }
+ return out
+}
+
+func strmCanonicalRoot(part string) string {
+ key := strings.ToLower(strings.TrimSpace(part))
+ switch key {
+ case "电影", "movie", "movies", "film", "films":
+ return "电影"
+ case "电视剧", "剧集", "tv", "tvs", "series", "show", "shows":
+ return "电视剧"
+ case "动漫", "动画", "anime", "bangumi":
+ return "动漫"
+ case "成人", "adult", "adults", "jav", "nsfw", "9kg":
+ return "成人"
+ default:
+ return ""
+ }
+}
+
+func strmCategoryRoot(part string) string {
+ key := strings.ToLower(strings.TrimSpace(part))
+ switch key {
+ case "动画电影", "动漫电影", "华语电影", "国产电影", "外语电影", "欧美电影", "日韩电影":
+ return "电影"
+ case "国产剧", "欧美剧", "日韩剧", "日剧", "韩剧", "综艺", "真人秀", "纪录片", "纪录", "未分类":
+ return "电视剧"
+ case "国漫", "国产动漫", "日番", "番剧", "日漫", "日本动漫", "日本动画", "儿童", "少儿":
+ return "动漫"
+ case "番号":
+ return "成人"
+ default:
+ return ""
+ }
+}
diff --git a/internal/service/strm_proxy.go b/internal/service/strm_proxy.go
new file mode 100644
index 0000000..8d62677
--- /dev/null
+++ b/internal/service/strm_proxy.go
@@ -0,0 +1,89 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// ProxySTRM proxies a STRM target and preserves Range requests for players.
+func (s *STRMService) ProxySTRM(ctx context.Context, id string, req *http.Request, w http.ResponseWriter) error {
+ record, err := s.repo.STRM.FindByID(ctx, id)
+ if err != nil {
+ return err
+ }
+ if record == nil {
+ return ErrSTRMNotFound
+ }
+ if !model.IsAllowedProtocol(record.Protocol) {
+ return ErrSTRMProtocolInvalid
+ }
+
+ targetURL, err := validateSTRMProxyURL(record.URL)
+ if err != nil {
+ return err
+ }
+ proxyReq, err := http.NewRequestWithContext(ctx, req.Method, targetURL.String(), nil)
+ if err != nil {
+ return fmt.Errorf("create proxy request: %w", err)
+ }
+ copySTRMRequestHeaders(req, proxyReq)
+
+ client := &http.Client{Timeout: 60 * time.Second}
+ resp, err := client.Do(proxyReq) // #nosec G107,G704 -- STRM proxy target is validated by validateSTRMProxyURL before request creation.
+ if err != nil {
+ return fmt.Errorf("proxy request failed: %w", err)
+ }
+ defer resp.Body.Close()
+
+ copySTRMResponseHeaders(resp, w)
+ w.WriteHeader(resp.StatusCode)
+ _, err = io.Copy(w, resp.Body)
+ return err
+}
+
+func validateSTRMProxyURL(raw string) (*url.URL, error) {
+ u, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil || u.Scheme == "" || u.Host == "" {
+ return nil, ErrSTRMURLInvalid
+ }
+ switch strings.ToLower(u.Scheme) {
+ case "http", "https":
+ default:
+ return nil, ErrSTRMProtocolInvalid
+ }
+ if isPrivateHost(u.Hostname()) {
+ return nil, ErrSTRMURLInvalid
+ }
+ return u, nil
+}
+
+func copySTRMRequestHeaders(src *http.Request, dst *http.Request) {
+ for _, header := range []string{
+ "Range", "If-Range", "If-Match", "If-None-Match",
+ "If-Modified-Since", "If-Unmodified-Since",
+ "Accept", "Accept-Encoding", "Accept-Language",
+ } {
+ if v := src.Header.Get(header); v != "" {
+ dst.Header.Set(header, v)
+ }
+ }
+}
+
+func copySTRMResponseHeaders(src *http.Response, dst http.ResponseWriter) {
+ for _, header := range []string{
+ "Content-Type", "Content-Length", "Content-Range",
+ "Accept-Ranges", "Last-Modified", "ETag",
+ "Cache-Control", "Content-Disposition",
+ } {
+ if v := src.Header.Get(header); v != "" {
+ dst.Header().Set(header, v)
+ }
+ }
+}
diff --git a/internal/service/strm_svc.go b/internal/service/strm_svc.go
index 7ae382b..ebf0d05 100644
--- a/internal/service/strm_svc.go
+++ b/internal/service/strm_svc.go
@@ -4,17 +4,8 @@ package service
import (
"context"
"errors"
- "fmt"
- "io"
- "net/http"
- "net/url"
- "os"
- "path/filepath"
- "strconv"
"strings"
- "time"
- "github.com/golang-jwt/jwt/v5"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
@@ -36,322 +27,11 @@ type STRMService struct {
cfg *config.Config
}
-type GenerateSTRMOptions struct {
- LibraryID string `json:"library_id"`
- OutputDir string `json:"output_dir"`
- BaseURL string `json:"base_url,omitempty"`
- Enabled bool `json:"enabled"`
- Overwrite bool `json:"overwrite"`
- IncludeLocal bool `json:"include_local"`
- PlaybackToken string `json:"-"`
-}
-
-type GenerateSTRMResult struct {
- LibraryID string `json:"library_id"`
- OutputDir string `json:"output_dir"`
- Generated int `json:"generated"`
- Updated int `json:"updated"`
- Skipped int `json:"skipped"`
- Errors []string `json:"errors,omitempty"`
- Items []GenerateSTRMItem `json:"items,omitempty"`
-}
-
-type GenerateSTRMItem struct {
- MediaID string `json:"media_id"`
- Title string `json:"title"`
- FilePath string `json:"file_path"`
- URL string `json:"url,omitempty"`
- Action string `json:"action"`
- Reason string `json:"reason,omitempty"`
-}
-
// NewSTRMService 创建 STRM 服务。
func NewSTRMService(log *zap.Logger, repo *repository.Container, cfg *config.Config) *STRMService {
return &STRMService{log: log, repo: repo, cfg: cfg}
}
-func (s *STRMService) GenerateForLibrary(ctx context.Context, opts GenerateSTRMOptions) (*GenerateSTRMResult, error) {
- if s == nil || s.repo == nil || s.repo.DB == nil {
- return nil, errors.New("strm service unavailable")
- }
- libraryID := strings.TrimSpace(opts.LibraryID)
- if libraryID == "" {
- return nil, errors.New("library_id required")
- }
- lib, err := s.repo.Library.FindByID(ctx, libraryID)
- if err != nil {
- return nil, err
- }
- if lib == nil {
- return nil, errors.New("library not found")
- }
- outputDir := resolveMappedDestinationPath(strings.TrimSpace(opts.OutputDir))
- if (outputDir == "" || outputDir == ".") && s.repo.Setting != nil {
- if saved, err := s.repo.Setting.Get(ctx, "strm.output_dir"); err == nil {
- outputDir = resolveMappedDestinationPath(strings.TrimSpace(saved))
- }
- }
- if outputDir == "" || outputDir == "." {
- outputDir = s.defaultOutputDir(lib)
- }
- if outputDir == "" || outputDir == "." {
- return nil, errors.New("output_dir required")
- }
- if strings.TrimSpace(opts.BaseURL) != "" && s.repo.Setting != nil {
- baseURL := strings.TrimRight(strings.TrimSpace(opts.BaseURL), "/")
- _ = s.repo.Setting.Set(ctx, "app.server_url", baseURL)
- _ = s.repo.Setting.Set(ctx, "strm.base_url", baseURL)
- }
- if s.repo.Setting != nil {
- _ = s.repo.Setting.Set(ctx, "strm.auto_generate_enabled", strconv.FormatBool(opts.Enabled))
- _ = s.repo.Setting.Set(ctx, "strm.output_dir", outputDir)
- }
- if err := os.MkdirAll(outputDir, 0o755); err != nil { // #nosec G301 -- STRM output directories must stay readable by NAS/player users.
- return nil, err
- }
-
- var rows []model.Media
- if err := s.repo.DB.WithContext(ctx).
- Where("library_id = ?", libraryID).
- Order("title asc, season_num asc, episode_num asc, created_at asc").
- Find(&rows).Error; err != nil {
- return nil, err
- }
-
- res := &GenerateSTRMResult{LibraryID: libraryID, OutputDir: outputDir}
- for _, media := range rows {
- select {
- case <-ctx.Done():
- return res, ctx.Err()
- default:
- }
- item := s.generateOne(ctx, *lib, media, outputDir, opts)
- res.Items = append(res.Items, item)
- switch item.Action {
- case "generated":
- res.Generated++
- case "updated":
- res.Updated++
- case "skipped":
- res.Skipped++
- case "error":
- res.Errors = append(res.Errors, fmt.Sprintf("%s: %s", item.Title, item.Reason))
- }
- }
- return res, nil
-}
-
-func (s *STRMService) defaultOutputDir(lib *model.Library) string {
- if s != nil && s.cfg != nil && strings.TrimSpace(s.cfg.App.DataDir) != "" {
- return filepath.Join(s.cfg.App.DataDir, "strm", sanitizeFilename(lib.Name))
- }
- return filepath.Join("data", "strm", sanitizeFilename(lib.Name))
-}
-
-func (s *STRMService) generateOne(ctx context.Context, lib model.Library, media model.Media, outputDir string, opts GenerateSTRMOptions) GenerateSTRMItem {
- item := GenerateSTRMItem{MediaID: media.ID, Title: media.Title}
- playURL := s.strmPlaybackURL(ctx, media, opts.BaseURL, opts.PlaybackToken)
- if playURL == "" {
- item.Action = "skipped"
- item.Reason = "no playable strm target"
- return item
- }
- if strings.TrimSpace(media.STRMURL) == "" && !opts.IncludeLocal {
- item.Action = "skipped"
- item.Reason = "local media skipped"
- return item
- }
- rel := s.strmRelativePath(lib, media)
- if rel == "" {
- item.Action = "skipped"
- item.Reason = "cannot build file name"
- return item
- }
- filePath := filepath.Join(outputDir, rel)
- item.FilePath = filePath
- item.URL = playURL
- if _, err := os.Stat(filePath); err == nil && !opts.Overwrite {
- item.Action = "skipped"
- item.Reason = "target exists"
- return item
- }
- action := "generated"
- if _, err := os.Stat(filePath); err == nil {
- action = "updated"
- }
- if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil { // #nosec G301 -- STRM output directories must stay readable by NAS/player users.
- item.Action = "error"
- item.Reason = err.Error()
- return item
- }
- if err := os.WriteFile(filePath, []byte(playURL+"\n"), 0o644); err != nil { // #nosec G306 -- STRM files are media sidecars intended to be readable by players.
- item.Action = "error"
- item.Reason = err.Error()
- return item
- }
- if err := s.upsertGeneratedRecord(ctx, media, filePath, playURL, lib.Type); err != nil {
- item.Action = "error"
- item.Reason = err.Error()
- return item
- }
- item.Action = action
- return item
-}
-
-func (s *STRMService) strmPlaybackURL(ctx context.Context, media model.Media, baseURL, playbackToken string) string {
- if media.ID == "" {
- return ""
- }
- query := url.Values{}
- token := strings.TrimSpace(playbackToken)
- if token == "" {
- token = s.defaultSTRMPlaybackToken(ctx)
- }
- if token != "" {
- query.Set("token", token)
- }
- return buildAbsoluteSTRMAPIURL(firstNonEmpty(baseURL, PublicServerURL(ctx, s.repo, s.cfg)), "/api/stream/"+url.PathEscape(media.ID), query)
-}
-
-func (s *STRMService) defaultSTRMPlaybackToken(ctx context.Context) string {
- if s == nil || s.repo == nil || s.repo.User == nil || s.cfg == nil || strings.TrimSpace(s.cfg.Secrets.JWTSecret) == "" {
- return ""
- }
- admin, err := s.repo.User.FirstAdmin(ctx)
- if err != nil || admin == nil {
- if err != nil && s.log != nil {
- s.log.Warn("generate strm playback token failed", zap.Error(err))
- }
- return ""
- }
- token, err := signSTRMPlaybackToken(admin, s.cfg.Secrets.JWTSecret)
- if err != nil {
- if s.log != nil {
- s.log.Warn("sign strm playback token failed", zap.Error(err))
- }
- return ""
- }
- return token
-}
-
-func signSTRMPlaybackToken(u *model.User, secret string) (string, error) {
- if u == nil || strings.TrimSpace(u.ID) == "" || strings.TrimSpace(secret) == "" {
- return "", ErrSTRMURLInvalid
- }
- claims := Claims{
- UserID: u.ID,
- Role: u.Role,
- Tier: u.Tier,
- RegisteredClaims: jwt.RegisteredClaims{
- IssuedAt: jwt.NewNumericDate(time.Now()),
- ExpiresAt: jwt.NewNumericDate(time.Now().Add(EmbyTokenDuration)),
- Issuer: "mediastationgo",
- Subject: u.ID,
- },
- }
- t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
- return t.SignedString([]byte(secret))
-}
-
-func (s *STRMService) strmRelativePath(lib model.Library, media model.Media) string {
- title := strings.TrimSpace(media.Title)
- if title == "" {
- title = strings.TrimSuffix(filepath.Base(media.Path), filepath.Ext(media.Path))
- }
- if title == "" {
- return ""
- }
- seriesLike := isSeriesLibraryType(lib.Type) || media.SeasonNum > 0 || media.EpisodeNum > 0
- if seriesLike {
- show := inferSeriesNameFromPath(media.Path)
- if show == "" {
- show = title
- }
- season := media.SeasonNum
- if season <= 0 {
- season = 1
- }
- name := title
- if media.EpisodeNum > 0 {
- name = fmt.Sprintf("%s - S%02dE%02d", show, season, media.EpisodeNum)
- }
- return filepath.Join(sanitizeFilename(show), fmt.Sprintf("Season %02d", season), sanitizeFilename(name)+".strm")
- }
- folder := title
- if media.Year > 0 && !strings.Contains(folder, strconv.Itoa(media.Year)) {
- folder = fmt.Sprintf("%s (%d)", folder, media.Year)
- }
- safe := sanitizeFilename(folder)
- return filepath.Join(safe, safe+".strm")
-}
-
-func (s *STRMService) upsertGeneratedRecord(ctx context.Context, media model.Media, filePath, playURL, mediaType string) error {
- protocol := ""
- if u, err := url.Parse(playURL); err == nil {
- protocol = strings.ToLower(u.Scheme)
- }
- if protocol == "" {
- protocol = "http"
- }
- record := model.STRMRecord{
- Title: media.Title,
- URL: playURL,
- FilePath: filePath,
- Protocol: protocol,
- MediaID: media.ID,
- MediaType: mediaType,
- SeasonNum: media.SeasonNum,
- EpisodeNum: media.EpisodeNum,
- }
- var existing model.STRMRecord
- err := s.repo.DB.WithContext(ctx).Where("media_id = ? AND file_path = ?", media.ID, filePath).First(&existing).Error
- if err == nil {
- existing.Title = record.Title
- existing.URL = record.URL
- existing.Protocol = record.Protocol
- existing.MediaType = record.MediaType
- existing.SeasonNum = record.SeasonNum
- existing.EpisodeNum = record.EpisodeNum
- return s.repo.DB.WithContext(ctx).Save(&existing).Error
- }
- return s.repo.DB.WithContext(ctx).Create(&record).Error
-}
-
-func absolutizeSTRMURL(raw, baseURL string) string {
- raw = strings.TrimSpace(raw)
- if raw == "" || strings.HasPrefix(raw, "//") {
- return raw
- }
- u, err := url.Parse(raw)
- if err == nil && u.IsAbs() {
- return raw
- }
- return buildAbsoluteSTRMAPIURL(baseURL, raw, nil)
-}
-
-func buildAbsoluteSTRMAPIURL(baseURL, apiPath string, query url.Values) string {
- apiPath = "/" + strings.TrimLeft(strings.TrimSpace(apiPath), "/")
- if query != nil && len(query) > 0 {
- apiPath += "?" + query.Encode()
- }
- baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
- if baseURL == "" {
- return apiPath
- }
- base, err := url.Parse(baseURL)
- if err != nil || base.Scheme == "" || base.Host == "" {
- return apiPath
- }
- target, err := url.Parse(apiPath)
- if err != nil {
- return apiPath
- }
- base.Path = strings.TrimRight(base.Path, "/") + "/" + strings.TrimLeft(target.Path, "/")
- base.RawQuery = target.RawQuery
- base.Fragment = ""
- return base.String()
-}
-
// Create 创建 STRM 记录。
func (s *STRMService) Create(ctx context.Context, record *model.STRMRecord) (*model.STRMRecord, error) {
if err := s.validateSTRM(record); err != nil {
@@ -467,87 +147,6 @@ func (s *STRMService) GetProtocols() []string {
return model.AllowedSTRMProtocols
}
-// ProxySTRM 代理访问 STRM 资源。
-// 支持 Range 请求(206 Partial Content)。
-func (s *STRMService) ProxySTRM(ctx context.Context, id string, req *http.Request, w http.ResponseWriter) error {
- record, err := s.repo.STRM.FindByID(ctx, id)
- if err != nil {
- return err
- }
- if record == nil {
- return ErrSTRMNotFound
- }
-
- if !model.IsAllowedProtocol(record.Protocol) {
- return ErrSTRMProtocolInvalid
- }
-
- targetURL, err := validateSTRMProxyURL(record.URL)
- if err != nil {
- return err
- }
-
- // 创建代理请求
- proxyReq, err := http.NewRequestWithContext(ctx, req.Method, targetURL.String(), nil)
- if err != nil {
- return fmt.Errorf("create proxy request: %w", err)
- }
-
- // 复制 Range 等关键请求头
- for _, header := range []string{
- "Range", "If-Range", "If-Match", "If-None-Match",
- "If-Modified-Since", "If-Unmodified-Since",
- "Accept", "Accept-Encoding", "Accept-Language",
- } {
- if v := req.Header.Get(header); v != "" {
- proxyReq.Header.Set(header, v)
- }
- }
-
- // 对 alist/webdav 协议可能需要特殊处理认证
- if record.Protocol == "alist" || record.Protocol == "alists" {
- // alist 协议可以直接访问,无需额外认证
- }
-
- client := &http.Client{Timeout: 60 * time.Second}
- resp, err := client.Do(proxyReq) // #nosec G107,G704 -- STRM proxy target is validated by validateSTRMProxyURL before request creation.
- if err != nil {
- return fmt.Errorf("proxy request failed: %w", err)
- }
- defer resp.Body.Close()
-
- // 复制响应头
- for _, header := range []string{
- "Content-Type", "Content-Length", "Content-Range",
- "Accept-Ranges", "Last-Modified", "ETag",
- "Cache-Control", "Content-Disposition",
- } {
- if v := resp.Header.Get(header); v != "" {
- w.Header().Set(header, v)
- }
- }
-
- w.WriteHeader(resp.StatusCode)
- _, err = io.Copy(w, resp.Body)
- return err
-}
-
-func validateSTRMProxyURL(raw string) (*url.URL, error) {
- u, err := url.Parse(strings.TrimSpace(raw))
- if err != nil || u.Scheme == "" || u.Host == "" {
- return nil, ErrSTRMURLInvalid
- }
- switch strings.ToLower(u.Scheme) {
- case "http", "https":
- default:
- return nil, ErrSTRMProtocolInvalid
- }
- if isPrivateHost(u.Hostname()) {
- return nil, ErrSTRMURLInvalid
- }
- return u, nil
-}
-
// validateSTRM 验证 STRM 记录。
func (s *STRMService) validateSTRM(record *model.STRMRecord) error {
if record.Title == "" {
diff --git a/internal/service/strm_svc_test.go b/internal/service/strm_svc_test.go
index 733348a..ca6b2b2 100644
--- a/internal/service/strm_svc_test.go
+++ b/internal/service/strm_svc_test.go
@@ -7,10 +7,8 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"github.com/golang-jwt/jwt/v5"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -18,13 +16,7 @@ import (
)
func TestGenerateSTRMForLibraryWritesFilesAndRecords(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{})
repos := repository.New(db)
lib := model.Library{Name: "电影", Path: "cloud://openlist/电影", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
@@ -90,13 +82,7 @@ func TestGenerateSTRMForLibraryWritesFilesAndRecords(t *testing.T) {
}
func TestGenerateSTRMForLibrarySignsDefaultPlaybackToken(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{}, &model.User{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{}, &model.User{})
repos := repository.New(db)
admin := model.User{Username: "admin", PasswordHash: "x", Role: "admin", Tier: "plus", IsActive: true}
if err := repos.User.Create(t.Context(), &admin); err != nil {
@@ -147,6 +133,162 @@ func TestGenerateSTRMForLibrarySignsDefaultPlaybackToken(t *testing.T) {
}
}
+func TestGenerateSTRMForLibraryCleanupStaleFilesAndRecords(t *testing.T) {
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{})
+ repos := repository.New(db)
+ lib := model.Library{Name: "电影", Path: "cloud://openlist/电影", Type: "movie", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{Base: model.Base{ID: "cloud-media"}, LibraryID: lib.ID, Title: "云盘电影", Year: 2026, Path: "cloud://openlist/电影/云盘电影.mkv", STRMURL: "/api/cloud/play/openlist?ref=movie"}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ outDir := filepath.Join(t.TempDir(), "strm")
+ stalePath := filepath.Join(outDir, "旧电影", "旧电影.strm")
+ if err := os.MkdirAll(filepath.Dir(stalePath), 0o755); err != nil {
+ t.Fatal(err)
+ }
+ if err := os.WriteFile(stalePath, []byte("http://old.example/stream\n"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ staleRecord := model.STRMRecord{Title: "旧电影", URL: "http://old.example/stream", FilePath: stalePath, Protocol: "http", MediaID: "missing-media"}
+ if err := repos.DB.Create(&staleRecord).Error; err != nil {
+ t.Fatal(err)
+ }
+
+ svc := NewSTRMService(zap.NewNop(), repos, &config.Config{})
+ res, err := svc.GenerateForLibrary(t.Context(), GenerateSTRMOptions{
+ LibraryID: lib.ID,
+ OutputDir: outDir,
+ BaseURL: "http://nas.example:18080",
+ IncludeLocal: true,
+ Overwrite: true,
+ PlaybackToken: "strm-token",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if res.Cleaned == 0 {
+ t.Fatalf("cleaned = %d, want stale file/record cleaned", res.Cleaned)
+ }
+ if _, err := os.Stat(stalePath); !os.IsNotExist(err) {
+ t.Fatalf("stale strm file should be removed, stat err=%v", err)
+ }
+ var count int64
+ if err := repos.DB.Model(&model.STRMRecord{}).Where("media_id = ?", "missing-media").Count(&count).Error; err != nil {
+ t.Fatal(err)
+ }
+ if count != 0 {
+ t.Fatalf("stale strm record count = %d, want 0", count)
+ }
+}
+
+func TestSTRMLibraryOutputSubdirUsesLibraryCategoryPath(t *testing.T) {
+ tests := []struct {
+ name string
+ lib model.Library
+ want string
+ }{
+ {
+ name: "cloud nested tv category",
+ lib: model.Library{Name: "OpenList · 欧美剧", Path: BuildCloudLibraryPath("openlist", "/电视剧/欧美剧", "/电视剧/欧美剧"), Type: "tv"},
+ want: filepath.Join("电视剧", "欧美剧"),
+ },
+ {
+ name: "cloud second-level category without root",
+ lib: model.Library{Name: "OpenList · 国产剧", Path: BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv"},
+ want: filepath.Join("电视剧", "国产剧"),
+ },
+ {
+ name: "local nested tv category",
+ lib: model.Library{Name: "欧美剧", Path: `F:\media\电视剧\欧美剧`, Type: "tv"},
+ want: filepath.Join("电视剧", "欧美剧"),
+ },
+ {
+ name: "fallback to type root",
+ lib: model.Library{Name: "Archive", Path: `F:\archive`, Type: "movie"},
+ want: "电影",
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := strmLibraryOutputSubdir(tt.lib); got != tt.want {
+ t.Fatalf("strmLibraryOutputSubdir() = %q, want %q", got, tt.want)
+ }
+ })
+ }
+}
+
+func TestGenerateSTRMForLibraryUsesCategoryDefaultOutputDir(t *testing.T) {
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{})
+ repos := repository.New(db)
+ dataDir := t.TempDir()
+ lib := model.Library{Name: "OpenList · 欧美剧", Path: BuildCloudLibraryPath("openlist", "/电视剧/欧美剧", "/电视剧/欧美剧"), Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &lib); err != nil {
+ t.Fatal(err)
+ }
+ media := model.Media{Base: model.Base{ID: "show-1"}, LibraryID: lib.ID, Title: "第一集", Path: "cloud://openlist/电视剧/欧美剧/Show/S01E01.mkv", STRMURL: "/api/cloud/play/openlist?ref=show", SeasonNum: 1, EpisodeNum: 1}
+ if err := repos.DB.Create(&media).Error; err != nil {
+ t.Fatal(err)
+ }
+ svc := NewSTRMService(zap.NewNop(), repos, &config.Config{App: config.AppConfig{DataDir: dataDir}})
+
+ res, err := svc.GenerateForLibrary(t.Context(), GenerateSTRMOptions{
+ LibraryID: lib.ID,
+ BaseURL: "http://nas.example:18080",
+ IncludeLocal: true,
+ PlaybackToken: "strm-token",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ wantDir := filepath.Join(dataDir, "strm", "电视剧", "欧美剧")
+ if res.OutputDir != wantDir {
+ t.Fatalf("output dir = %q, want %q", res.OutputDir, wantDir)
+ }
+ assertFileContains(t, filepath.Join(wantDir, "Show", "Season 01", "Show - S01E01.strm"), "http://nas.example:18080/api/stream/show-1?token=strm-token")
+}
+
+func TestGenerateSTRMForAllLibrariesWritesPerLibraryFolders(t *testing.T) {
+ db := newServiceTestDB(t, &model.Library{}, &model.Media{}, &model.STRMRecord{}, &model.Setting{})
+ repos := repository.New(db)
+ movieLib := model.Library{Name: "电影", Path: "cloud://openlist/电影", Type: "movie", Enabled: true}
+ tvLib := model.Library{Name: "欧美剧", Path: BuildCloudLibraryPath("openlist", "/电视剧/欧美剧", "/电视剧/欧美剧"), Type: "tv", Enabled: true}
+ if err := repos.Library.Create(t.Context(), &movieLib); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Library.Create(t.Context(), &tvLib); err != nil {
+ t.Fatal(err)
+ }
+ rows := []model.Media{
+ {Base: model.Base{ID: "movie-1"}, LibraryID: movieLib.ID, Title: "云盘电影", Year: 2026, Path: "cloud://openlist/电影/云盘电影.mkv", STRMURL: "/api/cloud/play/openlist?ref=movie"},
+ {Base: model.Base{ID: "show-1"}, LibraryID: tvLib.ID, Title: "第一集", Path: "cloud://openlist/电视剧/欧美剧/Show/S01E01.mkv", STRMURL: "/api/cloud/play/openlist?ref=show", SeasonNum: 1, EpisodeNum: 1},
+ }
+ for i := range rows {
+ if err := repos.DB.Create(&rows[i]).Error; err != nil {
+ t.Fatal(err)
+ }
+ }
+
+ outDir := filepath.Join(t.TempDir(), "strm-all")
+ svc := NewSTRMService(zap.NewNop(), repos, &config.Config{})
+ res, err := svc.GenerateForAllLibraries(t.Context(), GenerateSTRMOptions{
+ OutputDir: outDir,
+ BaseURL: "http://nas.example:18080",
+ IncludeLocal: true,
+ PlaybackToken: "strm-token",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if res.Generated != 2 {
+ t.Fatalf("generated = %d, want 2", res.Generated)
+ }
+ assertFileContains(t, filepath.Join(outDir, "电影", "云盘电影 (2026)", "云盘电影 (2026).strm"), "http://nas.example:18080/api/stream/movie-1?token=strm-token")
+ assertFileContains(t, filepath.Join(outDir, "电视剧", "欧美剧", "Show", "Season 01", "Show - S01E01.strm"), "http://nas.example:18080/api/stream/show-1?token=strm-token")
+}
+
func assertFileContains(t *testing.T, path, want string) {
t.Helper()
if got := readSTRM(t, path); got != want {
diff --git a/internal/service/strm_url.go b/internal/service/strm_url.go
new file mode 100644
index 0000000..7390258
--- /dev/null
+++ b/internal/service/strm_url.go
@@ -0,0 +1,138 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "path/filepath"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/golang-jwt/jwt/v5"
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *STRMService) strmPlaybackURL(ctx context.Context, media model.Media, baseURL, playbackToken string) string {
+ if media.ID == "" {
+ return ""
+ }
+ query := url.Values{}
+ token := strings.TrimSpace(playbackToken)
+ if token == "" {
+ token = s.defaultSTRMPlaybackToken(ctx)
+ }
+ if token != "" {
+ query.Set("token", token)
+ }
+ return buildAbsoluteSTRMAPIURL(firstNonEmpty(baseURL, PublicServerURL(ctx, s.repo, s.cfg)), "/api/stream/"+url.PathEscape(media.ID), query)
+}
+
+func (s *STRMService) defaultSTRMPlaybackToken(ctx context.Context) string {
+ if s == nil || s.repo == nil || s.repo.User == nil || s.cfg == nil || strings.TrimSpace(s.cfg.Secrets.JWTSecret) == "" {
+ return ""
+ }
+ admin, err := s.repo.User.FirstAdmin(ctx)
+ if err != nil || admin == nil {
+ if err != nil && s.log != nil {
+ s.log.Warn("generate strm playback token failed", zap.Error(err))
+ }
+ return ""
+ }
+ token, err := signSTRMPlaybackToken(admin, s.cfg.Secrets.JWTSecret)
+ if err != nil {
+ if s.log != nil {
+ s.log.Warn("sign strm playback token failed", zap.Error(err))
+ }
+ return ""
+ }
+ return token
+}
+
+func signSTRMPlaybackToken(u *model.User, secret string) (string, error) {
+ if u == nil || strings.TrimSpace(u.ID) == "" || strings.TrimSpace(secret) == "" {
+ return "", ErrSTRMURLInvalid
+ }
+ claims := Claims{
+ UserID: u.ID,
+ Role: u.Role,
+ Tier: u.Tier,
+ RegisteredClaims: jwt.RegisteredClaims{
+ IssuedAt: jwt.NewNumericDate(time.Now()),
+ ExpiresAt: jwt.NewNumericDate(time.Now().Add(EmbyTokenDuration)),
+ Issuer: "mediastationgo",
+ Subject: u.ID,
+ },
+ }
+ t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
+ return t.SignedString([]byte(secret))
+}
+
+func (s *STRMService) strmRelativePath(lib model.Library, media model.Media) string {
+ title := strings.TrimSpace(media.Title)
+ if title == "" {
+ title = strings.TrimSuffix(filepath.Base(media.Path), filepath.Ext(media.Path))
+ }
+ if title == "" {
+ return ""
+ }
+ seriesLike := isSeriesLibraryType(lib.Type) || media.SeasonNum > 0 || media.EpisodeNum > 0
+ if seriesLike {
+ show := inferSeriesNameFromPath(media.Path)
+ if show == "" {
+ show = title
+ }
+ season := media.SeasonNum
+ if season <= 0 {
+ season = 1
+ }
+ name := title
+ if media.EpisodeNum > 0 {
+ name = fmt.Sprintf("%s - S%02dE%02d", show, season, media.EpisodeNum)
+ }
+ return filepath.Join(sanitizeFilename(show), fmt.Sprintf("Season %02d", season), sanitizeFilename(name)+".strm")
+ }
+ folder := title
+ if media.Year > 0 && !strings.Contains(folder, strconv.Itoa(media.Year)) {
+ folder = fmt.Sprintf("%s (%d)", folder, media.Year)
+ }
+ safe := sanitizeFilename(folder)
+ return filepath.Join(safe, safe+".strm")
+}
+
+func absolutizeSTRMURL(raw, baseURL string) string {
+ raw = strings.TrimSpace(raw)
+ if raw == "" || strings.HasPrefix(raw, "//") {
+ return raw
+ }
+ u, err := url.Parse(raw)
+ if err == nil && u.IsAbs() {
+ return raw
+ }
+ return buildAbsoluteSTRMAPIURL(baseURL, raw, nil)
+}
+
+func buildAbsoluteSTRMAPIURL(baseURL, apiPath string, query url.Values) string {
+ apiPath = "/" + strings.TrimLeft(strings.TrimSpace(apiPath), "/")
+ if query != nil && len(query) > 0 {
+ apiPath += "?" + query.Encode()
+ }
+ baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
+ if baseURL == "" {
+ return apiPath
+ }
+ base, err := url.Parse(baseURL)
+ if err != nil || base.Scheme == "" || base.Host == "" {
+ return apiPath
+ }
+ target, err := url.Parse(apiPath)
+ if err != nil {
+ return apiPath
+ }
+ base.Path = strings.TrimRight(base.Path, "/") + "/" + strings.TrimLeft(target.Path, "/")
+ base.RawQuery = target.RawQuery
+ base.Fragment = ""
+ return base.String()
+}
diff --git a/internal/service/subscription.go b/internal/service/subscription.go
index f1b96d3..e824dd4 100644
--- a/internal/service/subscription.go
+++ b/internal/service/subscription.go
@@ -8,13 +8,8 @@ package service
import (
"context"
- "encoding/xml"
"errors"
"fmt"
- "io"
- "net/http"
- "net/url"
- "regexp"
"strings"
"time"
@@ -68,24 +63,6 @@ func (s *SubscriptionService) Start(ctx context.Context) {
// Stop shuts the loop down.
func (s *SubscriptionService) Stop() { close(s.stop) }
-// rssFeed is the minimal RSS subset we need to decode.
-type rssFeed struct {
- XMLName xml.Name `xml:"rss"`
- Channel struct {
- Items []rssItem `xml:"item"`
- } `xml:"channel"`
-}
-
-type rssItem struct {
- Title string `xml:"title"`
- Link string `xml:"link"`
- GUID string `xml:"guid"`
- Description string `xml:"description"`
- Enclosure struct {
- URL string `xml:"url,attr"`
- } `xml:"enclosure"`
-}
-
// Create persists a new subscription.
func (s *SubscriptionService) Create(ctx context.Context, sub *model.Subscription) error {
if sub.Name == "" || sub.FeedURL == "" {
@@ -125,46 +102,6 @@ func (s *SubscriptionService) List(ctx context.Context) ([]model.Subscription, e
return s.repo.Subscription.List(ctx)
}
-// History returns completed/archived subscription rules.
-func (s *SubscriptionService) History(ctx context.Context) ([]model.Subscription, error) {
- return s.repo.Subscription.History(ctx)
-}
-
-// Restore moves an archived subscription back to the active management list.
-// It also clears the per-subscription seen state so an unfinished historical
-// rule can match resources again when it is run next.
-func (s *SubscriptionService) Restore(ctx context.Context, id string) (*model.Subscription, error) {
- var sub model.Subscription
- if err := s.repo.DB.WithContext(ctx).Where("id = ?", id).First(&sub).Error; err != nil {
- return nil, err
- }
- if err := s.repo.DB.WithContext(ctx).Model(&model.Subscription{}).
- Where("id = ?", id).
- Updates(map[string]any{
- "enabled": true,
- "archive_reason": "",
- // 重置为 0:此前可能被 feed 低估并锁死(updateSubscriptionTotalEpisodes
- // 只增不减,resolveSubscriptionTotalEpisodes 见 >0 即不再回查元数据)。
- // 归零后下次 run 会从 TMDb/豆瓣等权威源重算真实总集数,避免恢复后
- // 因"误判已无缺集"而不再搜索资源。
- "total_episodes": 0,
- }).Error; err != nil {
- return nil, err
- }
- if err := s.repo.DB.WithContext(ctx).
- Exec("UPDATE subscriptions SET archived_at = NULL WHERE id = ?", id).Error; err != nil {
- return nil, err
- }
- if s.repo.Setting != nil {
- _ = s.repo.Setting.Delete(ctx, fmt.Sprintf("subscription.%s.seen", id))
- }
- var restored model.Subscription
- if err := s.repo.DB.WithContext(ctx).Where("id = ?", id).First(&restored).Error; err != nil {
- return nil, err
- }
- return &restored, nil
-}
-
// Delete removes a subscription.
func (s *SubscriptionService) Delete(ctx context.Context, id string) error {
var sub model.Subscription
@@ -332,870 +269,3 @@ func (s *SubscriptionService) runOne(ctx context.Context, sub *model.Subscriptio
}
return queued, nil
}
-
-func (s *SubscriptionService) runSiteSearch(ctx context.Context, sub *model.Subscription) (int, error) {
- if s.site == nil {
- if s.log != nil {
- s.log.Warn("site-search subscription service unavailable", subscriptionSiteSearchLogFields(sub, "")...)
- }
- return 0, errors.New("site search service unavailable")
- }
- keywords := siteSearchKeywords(sub)
- keyword := ""
- if len(keywords) > 0 {
- keyword = keywords[0]
- }
- if keyword == "" {
- if s.log != nil {
- s.log.Warn("site-search subscription keyword missing", subscriptionSiteSearchLogFields(sub, "")...)
- }
- return 0, errors.New("site-search subscription keyword required")
- }
- if s.log != nil {
- s.log.Info("site-search subscription run started", subscriptionSiteSearchLogFields(sub, keyword)...)
- }
-
- var (
- results []SearchResult
- lastSearchErr error
- searchErrors int
- )
- for _, searchKeyword := range keywords {
- found, err := s.site.Search(ctx, searchKeyword)
- if err != nil {
- lastSearchErr = err
- searchErrors++
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, searchKeyword)
- fields = append(fields, zap.Error(err))
- s.log.Warn("site-search subscription search failed", fields...)
- }
- continue
- }
- results = append(results, found...)
- }
- results = dedupeSiteSearchResults(results)
- if len(results) == 0 && lastSearchErr != nil && searchErrors == len(keywords) {
- return 0, lastSearchErr
- }
- if len(results) == 0 {
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = append(fields, zap.Int("results_count", 0))
- s.log.Info("site-search subscription no results", fields...)
- }
- now := time.Now()
- _ = s.repo.DB.Model(sub).Updates(map[string]any{"last_run_at": &now}).Error
- return 0, nil
- }
- s.updateSubscriptionTotalEpisodes(ctx, sub, s.resolveSubscriptionTotalEpisodes(ctx, sub, inferSearchTotalEpisodes(results, sub)))
-
- guidKey := fmt.Sprintf("subscription.%s.seen", sub.ID)
- seenRaw, _ := s.repo.Setting.Get(ctx, guidKey)
- seen := splitNonEmpty(seenRaw)
- seenSet := make(map[string]struct{}, len(seen))
- for _, g := range seen {
- seenSet[g] = struct{}{}
- }
-
- availability := mergeLocalAvailability(
- SubscriptionLocalAvailability(ctx, s.repo, sub),
- s.pendingDownloadAvailability(ctx, sub),
- )
- candidates, selectionStats := selectSiteSearchCandidatesWithStats(results, sub, seenSet, availability)
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = appendSiteSearchSelectionLogFields(fields, selectionStats)
- fields = appendAvailabilityLogFields(fields, availability)
- s.log.Info("site-search subscription selection summary", fields...)
- }
- var lastEnqueueErr error
- queued := 0
- var resources []string
- for _, candidate := range candidates {
- item := candidate.Item
- matchText := subscriptionSearchResultText(item)
- mediaType, mediaCategory := s.classifySubscriptionItem(ctx, sub, matchText, item.Category)
- if s.shouldSkipExistingTorrent(ctx, mediaType, candidate) {
- addSiteSearchCandidateAvailability(candidate, &availability)
- seen = append(seen, candidate.GUID)
- seenSet[candidate.GUID] = struct{}{}
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = append(fields,
- zap.String("reason", "existing_torrent"),
- zap.String("title", item.Title),
- zap.String("subtitle", item.Subtitle),
- zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
- zap.String("site_category", item.Category),
- zap.Int("season", candidate.Season),
- zap.Int("episode", candidate.Episode),
- zap.Bool("pack", candidate.Pack),
- zap.String("media_type", mediaType),
- )
- s.log.Info("site-search subscription candidate skipped", fields...)
- }
- continue
- }
- realURL := s.site.ResolveDownloadURL(ctx, candidate.Download)
- savePath := s.resolveSubscriptionSavePath(ctx, sub, mediaType, mediaCategory)
- if s.downloadPathHasCandidate(ctx, sub, matchText, savePath) {
- addSiteSearchCandidateAvailability(candidate, &availability)
- seen = append(seen, candidate.GUID)
- seenSet[candidate.GUID] = struct{}{}
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = append(fields,
- zap.String("reason", "download_path_has_candidate"),
- zap.String("title", item.Title),
- zap.String("subtitle", item.Subtitle),
- zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
- zap.String("site_category", item.Category),
- zap.Int("season", candidate.Season),
- zap.Int("episode", candidate.Episode),
- zap.Bool("pack", candidate.Pack),
- zap.String("media_type", mediaType),
- zap.String("media_category", mediaCategory),
- zap.String("save_path", savePath),
- )
- s.log.Info("site-search subscription candidate skipped", fields...)
- }
- continue
- }
- if _, err := s.downloads.AddDownloadWithMeta(ctx, sub.UserID, realURL, savePath, DownloadTaskMeta{
- SubscriptionID: sub.ID,
- Title: firstNonEmpty(item.Title, sub.Name),
- PosterURL: sub.PosterURL,
- BackdropURL: sub.BackdropURL,
- Overview: sub.Overview,
- MediaType: mediaType,
- MediaCategory: mediaCategory,
- SourceCategory: item.Category,
- AllowExistingLibrary: sub.WashEnabled,
- }); err != nil {
- if IsDownloadDedupError(err) {
- addSiteSearchCandidateAvailability(candidate, &availability)
- seen = append(seen, candidate.GUID)
- seenSet[candidate.GUID] = struct{}{}
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = append(fields,
- zap.String("reason", "download_dedup"),
- zap.String("title", item.Title),
- zap.String("subtitle", item.Subtitle),
- zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
- zap.String("site_category", item.Category),
- zap.Int("season", candidate.Season),
- zap.Int("episode", candidate.Episode),
- zap.Bool("pack", candidate.Pack),
- zap.String("media_type", mediaType),
- zap.String("media_category", mediaCategory),
- zap.String("save_path", savePath),
- )
- s.log.Info("site-search subscription candidate skipped", fields...)
- }
- continue
- }
- lastEnqueueErr = err
- s.log.Warn("site-search subscription enqueue failed",
- zap.String("subscription_id", sub.ID),
- zap.String("subscription", sub.Name),
- zap.String("keyword", keyword),
- zap.String("title", item.Title),
- zap.String("subtitle", item.Subtitle),
- zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
- zap.String("site_category", item.Category),
- zap.String("media_type", mediaType),
- zap.String("media_category", mediaCategory),
- zap.String("save_path", savePath),
- zap.Error(err))
- continue
- }
- queued++
- addSiteSearchCandidateAvailability(candidate, &availability)
- resources = append(resources, item.Title)
- seen = append(seen, candidate.GUID)
- seenSet[candidate.GUID] = struct{}{}
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = append(fields,
- zap.String("title", item.Title),
- zap.String("subtitle", item.Subtitle),
- zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
- zap.String("site_category", item.Category),
- zap.Int("season", candidate.Season),
- zap.Int("episode", candidate.Episode),
- zap.Bool("pack", candidate.Pack),
- zap.Int("score", candidate.Score),
- zap.String("media_type", mediaType),
- zap.String("media_category", mediaCategory),
- zap.String("save_path", savePath),
- )
- s.log.Info("site-search subscription candidate queued", fields...)
- }
- }
- availability = s.finalizePendingAvailability(sub, availability)
- if len(seen) > 200 {
- seen = seen[len(seen)-200:]
- }
- _ = s.repo.Setting.Set(ctx, guidKey, strings.Join(seen, "\n"))
- now := time.Now()
- _ = s.repo.DB.Model(sub).Updates(map[string]any{"last_run_at": &now}).Error
- _ = s.archiveCompletedSubscription(ctx, sub, availability)
- if queued > 0 {
- s.hub.Publish("subscription", map[string]any{
- "id": sub.ID,
- "name": sub.Name,
- "queued": queued,
- "keyword": keyword,
- "resources": resources,
- })
- s.notifySubscriptionHit(sub, queued, resources)
- return queued, nil
- }
- if lastEnqueueErr != nil {
- return 0, fmt.Errorf("找到 PT 资源但加入下载器失败: %w", lastEnqueueErr)
- }
- if s.log != nil {
- fields := subscriptionSiteSearchLogFields(sub, keyword)
- fields = appendSiteSearchSelectionLogFields(fields, selectionStats)
- fields = appendAvailabilityLogFields(fields, availability)
- fields = append(fields, zap.Int("queued", queued))
- s.log.Info("site-search subscription no candidate queued", fields...)
- }
- return 0, nil
-}
-
-func subscriptionSiteSearchLogFields(sub *model.Subscription, keyword string) []zap.Field {
- fields := []zap.Field{zap.String("keyword", keyword), zap.Strings("search_keywords", siteSearchKeywords(sub))}
- if sub == nil {
- return fields
- }
- fields = append(fields,
- zap.String("subscription_id", sub.ID),
- zap.String("subscription", sub.Name),
- zap.String("filter", sub.Filter),
- zap.String("media_type", sub.MediaType),
- zap.String("media_category", sub.MediaCategory),
- zap.String("search_mode", sub.SearchMode),
- zap.String("imdb_id", sub.IMDBID),
- zap.Bool("wash_enabled", sub.WashEnabled),
- zap.String("wash_priority", sub.WashPriority),
- zap.Int("total_episodes", sub.TotalEpisodes),
- )
- return fields
-}
-
-func appendSiteSearchSelectionLogFields(fields []zap.Field, stats siteSearchSelectionStats) []zap.Field {
- return append(fields,
- zap.Int("results_count", stats.Total),
- zap.Int("query_mismatch_count", stats.QueryMismatch),
- zap.Int("relaxed_query_match_count", stats.RelaxedQueryMatch),
- zap.Int("rule_mismatch_count", stats.RuleMismatch),
- zap.Int("missing_download_count", stats.MissingDownload),
- zap.Int("seen_count", stats.Seen),
- zap.Int("prepared_count", stats.Prepared),
- zap.Int("selected_count", stats.Selected),
- zap.Bool("local_already_satisfied", stats.LocalAlreadySatisfied),
- zap.Bool("local_series_pack_present", stats.LocalSeriesPackPresent),
- zap.Bool("series_complete", stats.SeriesComplete),
- zap.Int("existing_episode_skipped_count", stats.ExistingEpisodeSkipped),
- zap.Int("not_missing_episode_skipped_count", stats.NotMissingEpisodeSkipped),
- zap.Int("no_episode_skipped_count", stats.NoEpisodeSkipped),
- zap.Bool("pack_fallback_available", stats.PackFallbackAvailable),
- zap.Bool("pack_fallback_used", stats.PackFallbackUsed),
- )
-}
-
-func appendAvailabilityLogFields(fields []zap.Field, availability LocalAvailability) []zap.Field {
- missingSample, missingMore := limitedEpisodeSample(availability.MissingEpisodes, 20)
- return append(fields,
- zap.Int("local_media_count", availability.LocalMediaCount),
- zap.Bool("in_library", availability.InLibrary),
- zap.Bool("has_series_pack", availability.HasSeriesPack),
- zap.Int("downloaded_episodes", availability.DownloadedEpisodes),
- zap.Int("availability_total_episodes", availability.TotalEpisodes),
- zap.Int("missing_episode_count", len(availability.MissingEpisodes)),
- zap.Ints("missing_episodes", missingSample),
- zap.Int("missing_episodes_more", missingMore),
- )
-}
-
-func limitedEpisodeSample(values []int, limit int) ([]int, int) {
- if limit <= 0 || len(values) == 0 {
- return nil, len(values)
- }
- if len(values) <= limit {
- out := append([]int(nil), values...)
- return out, 0
- }
- out := append([]int(nil), values[:limit]...)
- return out, len(values) - limit
-}
-
-func (s *SubscriptionService) notifySubscriptionHit(sub *model.Subscription, queued int, resources []string) {
- if s == nil || s.notify == nil || sub == nil || queued <= 0 {
- return
- }
- body := fmt.Sprintf("订阅:%s\n新增资源:%d", sub.Name, queued)
- if len(resources) > 0 {
- body += "\n资源:\n- " + strings.Join(resources, "\n- ")
- }
- go func() {
- ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
- defer cancel()
- data := map[string]interface{}{}
- if strings.TrimSpace(sub.PosterURL) != "" {
- data["poster_url"] = sub.PosterURL
- }
- if strings.TrimSpace(sub.BackdropURL) != "" {
- data["backdrop_url"] = sub.BackdropURL
- }
- if strings.TrimSpace(sub.MediaType) != "" {
- data["media_type"] = sub.MediaType
- }
- if strings.TrimSpace(sub.MediaCategory) != "" {
- data["media_category"] = sub.MediaCategory
- }
- // 补充媒体通知模板(formatTelegramMediaNotification)所需字段:片名 / 原名 /
- // 语言 / 年份 / 评分 / 类型 / 简介 / 外链 / 资源标题(供模板提取季集 + 版本)。
- // 仅填现成可用的,缺失项模板会自动略过。
- if strings.TrimSpace(sub.Name) != "" {
- data["title"] = sub.Name
- }
- if strings.TrimSpace(sub.OriginalName) != "" {
- data["original_title"] = sub.OriginalName
- }
- if strings.TrimSpace(sub.OriginalLanguage) != "" {
- data["original_language"] = sub.OriginalLanguage
- }
- if sub.Year > 0 {
- data["year"] = sub.Year
- }
- if sub.Rating > 0 {
- data["rating"] = sub.Rating
- }
- if strings.TrimSpace(sub.Genres) != "" {
- data["genres"] = sub.Genres
- }
- if strings.TrimSpace(sub.Overview) != "" {
- data["overview"] = sub.Overview
- }
- if id := strings.TrimSpace(sub.IMDBID); id != "" {
- data["imdb_url"] = "https://www.imdb.com/title/" + id + "/"
- }
- if len(resources) > 0 {
- data["resource_title"] = resources[0]
- }
- s.notify.BroadcastEvent(ctx, NotifyEvent{
- Type: EventSubscriptionHit,
- Title: "MediaStationGo 订阅命中新资源",
- Message: body,
- Data: data,
- })
- }()
-}
-
-func (s *SubscriptionService) archiveCompletedSubscription(ctx context.Context, sub *model.Subscription, availability LocalAvailability) error {
- if s == nil || s.repo == nil || s.repo.Subscription == nil || sub == nil {
- return nil
- }
- if !subscriptionShouldArchive(sub, availability) {
- return nil
- }
- now := time.Now()
- reason := subscriptionArchiveReason(sub, availability)
- if err := s.repo.Subscription.Archive(ctx, sub.ID, reason, now); err != nil {
- return err
- }
- sub.Enabled = false
- sub.ArchivedAt = &now
- sub.ArchiveReason = reason
- if s.log != nil {
- s.log.Info("subscription completed, moved to history",
- zap.String("id", sub.ID),
- zap.String("name", sub.Name),
- zap.String("reason", reason))
- }
- if s.hub != nil {
- s.hub.Publish("subscription", map[string]any{
- "id": sub.ID,
- "name": sub.Name,
- "archived": true,
- "reason": reason,
- })
- }
- return nil
-}
-
-func subscriptionShouldArchive(sub *model.Subscription, availability LocalAvailability) bool {
- if sub == nil || sub.WashEnabled || sub.ArchivedAt != nil {
- return false
- }
- mediaType := strings.ToLower(strings.TrimSpace(sub.MediaType))
- if !isSubscriptionSeriesType(mediaType) {
- return availability.InLibrary || availability.LocalMediaCount > 0 || availability.DownloadedEpisodes > 0
- }
- if availability.HasSeriesPack {
- return true
- }
- total := sub.TotalEpisodes
- if total <= 0 {
- total = availability.TotalEpisodes
- }
- if total > 0 {
- return availability.DownloadedEpisodes >= total && len(availability.MissingEpisodes) == 0
- }
- return subscriptionLooksSingleEpisode(sub) && availability.DownloadedEpisodes > 0
-}
-
-func subscriptionArchiveReason(sub *model.Subscription, availability LocalAvailability) string {
- if sub != nil && sub.WashEnabled {
- return ""
- }
- if availability.HasSeriesPack {
- return "整季资源已加入下载/入库"
- }
- if availability.TotalEpisodes > 0 {
- return fmt.Sprintf("订阅完成:%d/%d", availability.DownloadedEpisodes, availability.TotalEpisodes)
- }
- if availability.DownloadedEpisodes > 0 {
- return "单集订阅已加入下载/入库"
- }
- return "订阅媒体已加入下载/入库"
-}
-
-func (s *SubscriptionService) updateSubscriptionTotalEpisodes(ctx context.Context, sub *model.Subscription, total int) {
- if s == nil || s.repo == nil || s.repo.DB == nil || sub == nil || total <= sub.TotalEpisodes {
- return
- }
- sub.TotalEpisodes = total
- _ = s.repo.DB.WithContext(ctx).Model(sub).Update("total_episodes", total).Error
-}
-
-func inferRSSTotalEpisodes(items []rssItem, sub *model.Subscription, filter *regexp.Regexp) int {
- if !subscriptionShouldInferTotal(sub) {
- return 0
- }
- maxEpisode := 0
- for _, item := range items {
- title := strings.TrimSpace(item.Title)
- if title == "" {
- continue
- }
- if filter != nil && !filter.MatchString(title) {
- continue
- }
- if !subscriptionTitleMatchesQuery(sub, title) {
- continue
- }
- if !matchesSubscriptionRules(sub, title) {
- continue
- }
- _, episode := ParseEpisode(title)
- if episode > maxEpisode {
- maxEpisode = episode
- }
- }
- return maxEpisode
-}
-
-func inferSearchTotalEpisodes(results []SearchResult, sub *model.Subscription) int {
- if !subscriptionShouldInferTotal(sub) {
- return 0
- }
- maxEpisode := 0
- for _, item := range results {
- matchText := subscriptionSearchResultText(item)
- if !subscriptionTitleMatchesQuery(sub, matchText) {
- continue
- }
- if !matchesSubscriptionRules(sub, matchText) {
- continue
- }
- _, episode := ParseEpisode(matchText)
- if episode > maxEpisode {
- maxEpisode = episode
- }
- }
- return maxEpisode
-}
-
-func subscriptionShouldInferTotal(sub *model.Subscription) bool {
- if sub == nil {
- return false
- }
- mediaType := normalizeMediaType(sub.MediaType, sub.Name+" "+sub.Filter, "")
- return isSubscriptionSeriesType(mediaType)
-}
-
-func (s *SubscriptionService) resolveSubscriptionTotalEpisodes(ctx context.Context, sub *model.Subscription, fallback int) int {
- if !subscriptionShouldInferTotal(sub) {
- return 0
- }
- if sub.TotalEpisodes > 0 {
- return sub.TotalEpisodes
- }
- if total := s.resolveSubscriptionMetadataTotalEpisodes(ctx, sub); total > 0 {
- return total
- }
- return fallback
-}
-
-func (s *SubscriptionService) resolveSubscriptionMetadataTotalEpisodes(ctx context.Context, sub *model.Subscription) int {
- if s == nil || s.scraper == nil || sub == nil {
- return 0
- }
- queries := subscriptionEpisodeMetadataQueries(sub)
-
- // Priority: TMDb -> Douban -> Bangumi -> TheTVDB -> Fanart -> title fallback.
- // Fanart.tv is artwork-only in MediaStationGo, so it intentionally does not
- // claim episode counts and lets the title fallback handle the final layer.
- if s.scraper.tmdb != nil {
- if id := subscriptionExplicitTMDbID(sub); id > 0 {
- if total, err := s.scraper.tmdb.GetTVEpisodeCount(ctx, id); err == nil && total > 0 {
- return total
- } else if err != nil && s.log != nil {
- s.log.Debug("subscription tmdb episode count failed", zap.Int("tmdb_id", id), zap.Error(err))
- }
- }
- for _, query := range queries {
- match, err := s.scraper.tmdb.SearchTV(ctx, query, 0)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription tmdb search failed", zap.String("query", query), zap.Error(err))
- }
- continue
- }
- if match == nil || match.TMDbID <= 0 {
- continue
- }
- total, err := s.scraper.tmdb.GetTVEpisodeCount(ctx, match.TMDbID)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription tmdb episode count failed", zap.Int("tmdb_id", match.TMDbID), zap.Error(err))
- }
- continue
- }
- if total > 0 {
- return total
- }
- }
- }
-
- if s.scraper.douban != nil {
- for _, query := range queries {
- total, err := s.scraper.douban.GetEpisodeCount(ctx, query)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription douban episode count failed", zap.String("query", query), zap.Error(err))
- }
- continue
- }
- if total > 0 {
- return total
- }
- }
- }
-
- if s.scraper.bangumi != nil {
- for _, query := range queries {
- match, err := s.scraper.bangumi.Search(ctx, query)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription bangumi search failed", zap.String("query", query), zap.Error(err))
- }
- continue
- }
- if match == nil || match.BangumiID <= 0 {
- continue
- }
- total, err := s.scraper.bangumi.GetEpisodeCount(ctx, match.BangumiID)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription bangumi episode count failed", zap.Int("bangumi_id", match.BangumiID), zap.Error(err))
- }
- continue
- }
- if total > 0 {
- return total
- }
- }
- }
-
- if s.scraper.thetvdb != nil {
- for _, query := range queries {
- match, err := s.scraper.thetvdb.SearchSeries(ctx, query)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription thetvdb search failed", zap.String("query", query), zap.Error(err))
- }
- continue
- }
- if match == nil || strings.TrimSpace(match.TheTVDBID) == "" {
- continue
- }
- total, err := s.scraper.thetvdb.GetSeriesEpisodeCount(ctx, match.TheTVDBID)
- if err != nil {
- if s.log != nil {
- s.log.Debug("subscription thetvdb episode count failed", zap.String("thetvdb_id", match.TheTVDBID), zap.Error(err))
- }
- continue
- }
- if total > 0 {
- return total
- }
- }
- }
-
- return 0
-}
-
-func subscriptionTitleMatchesQuery(sub *model.Subscription, title string) bool {
- if strings.TrimSpace(title) == "" {
- return false
- }
- for _, query := range subscriptionTitleMatchQueries(sub) {
- if strings.Contains(normalizeAvailabilityComparable(title), normalizeAvailabilityComparable(query)) {
- return true
- }
- }
- return len(subscriptionTitleMatchQueries(sub)) == 0
-}
-
-func subscriptionTitleMatchQueries(sub *model.Subscription) []string {
- if sub == nil {
- return nil
- }
- values := []string{
- availabilityQuery(subscriptionName(sub), subscriptionFilter(sub)),
- cleanAvailabilityTitle(subscriptionFilter(sub)),
- cleanAvailabilityTitle(subscriptionName(sub)),
- }
- for _, alias := range subscriptionFeedAliases(sub) {
- values = append(values, alias, cleanAvailabilityTitle(alias))
- }
- return compactUniqueStrings(values...)
-}
-
-func subscriptionEpisodeMetadataQueries(sub *model.Subscription) []string {
- if sub == nil {
- return nil
- }
- raw := []string{
- siteSearchKeyword(sub),
- sub.Filter,
- sub.Name,
- availabilityQuery(subscriptionName(sub), subscriptionFilter(sub)),
- }
- out := make([]string, 0, len(raw)*2)
- for _, value := range raw {
- value = cleanAvailabilityTitle(value)
- if value == "" {
- continue
- }
- if cleaned, _ := CleanQuery(value); cleaned != "" {
- out = append(out, cleaned)
- }
- out = append(out, value)
- }
- return compactUniqueStrings(out...)
-}
-
-func subscriptionExplicitTMDbID(sub *model.Subscription) int {
- if sub == nil {
- return 0
- }
- values := []string{sub.Name, sub.Filter, sub.FeedURL}
- for _, raw := range values {
- for _, pattern := range []string{`(?i)\btmdb[_:\-\s=]+(\d{2,})`, `(?i)\btmdbid[_:\-\s=]+(\d{2,})`} {
- if m := regexp.MustCompile(pattern).FindStringSubmatch(raw); len(m) >= 2 {
- var id int
- if _, err := fmt.Sscanf(m[1], "%d", &id); err == nil && id > 0 {
- return id
- }
- }
- }
- if u, err := url.Parse(raw); err == nil {
- for _, key := range []string{"tmdb_id", "tmdb", "tmdbid"} {
- var id int
- if _, err := fmt.Sscanf(u.Query().Get(key), "%d", &id); err == nil && id > 0 {
- return id
- }
- }
- }
- }
- return 0
-}
-
-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 subscriptionLooksSingleEpisode(sub *model.Subscription) bool {
- if sub == nil {
- return false
- }
- for _, value := range []string{sub.Name, sub.Filter} {
- _, episode := ParseEpisode(value)
- if episode > 0 {
- return true
- }
- }
- return false
-}
-
-func (s *SubscriptionService) shouldSkipExistingTorrent(ctx context.Context, mediaType string, candidate siteSearchCandidate) bool {
- if s == nil || s.downloads == nil {
- return false
- }
- if isSubscriptionSeriesType(mediaType) && !candidate.Pack && candidate.Episode > 0 {
- return false
- }
- return s.downloads.TorrentExistsByName(ctx, candidate.Item.Title)
-}
-
-func siteSearchKeywords(sub *model.Subscription) []string {
- if sub == nil {
- return nil
- }
- values := make([]string, 0, 8)
- if strings.EqualFold(strings.TrimSpace(sub.SearchMode), "imdb") && strings.TrimSpace(sub.IMDBID) != "" {
- values = append(values, strings.TrimSpace(sub.IMDBID))
- }
- if u, err := url.Parse(sub.FeedURL); err == nil {
- if keyword := strings.TrimSpace(u.Query().Get("keyword")); keyword != "" {
- values = append(values, keyword)
- }
- }
- if strings.TrimSpace(sub.Filter) != "" {
- values = append(values, sub.Filter)
- }
- if len(values) == 0 && strings.TrimSpace(sub.Name) != "" {
- values = append(values, sub.Name)
- }
- values = append(values, subscriptionFeedAliases(sub)...)
- for _, value := range append([]string(nil), values...) {
- if cleaned := cleanAvailabilityTitle(value); cleaned != "" {
- values = append(values, cleaned)
- }
- }
- return compactUniqueStrings(values...)
-}
-
-func siteSearchKeyword(sub *model.Subscription) string {
- keywords := siteSearchKeywords(sub)
- if len(keywords) == 0 {
- return ""
- }
- return keywords[0]
-}
-
-func subscriptionFeedAliases(sub *model.Subscription) []string {
- if sub == nil {
- return nil
- }
- u, err := url.Parse(sub.FeedURL)
- if err != nil {
- return nil
- }
- q := u.Query()
- values := make([]string, 0, len(q["alias"])+2)
- values = append(values, q["alias"]...)
- for _, raw := range q["aliases"] {
- for _, part := range strings.FieldsFunc(raw, func(r rune) bool {
- return r == '|' || r == '\n' || r == '\r' || r == '\t'
- }) {
- values = append(values, part)
- }
- }
- return compactUniqueStrings(values...)
-}
-
-func dedupeSiteSearchResults(results []SearchResult) []SearchResult {
- if len(results) < 2 {
- return results
- }
- seen := make(map[string]struct{}, len(results))
- out := make([]SearchResult, 0, len(results))
- for _, item := range results {
- download := strings.TrimSpace(item.DownloadURL)
- if download == "" {
- download = strings.TrimSpace(item.TorrentURL)
- }
- key := stableSiteSearchGUID(item, download)
- if _, ok := seen[key]; ok {
- continue
- }
- seen[key] = struct{}{}
- out = append(out, item)
- }
- return out
-}
-
-func (s *SubscriptionService) fetch(ctx context.Context, feedURL string) (*rssFeed, error) {
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, feedURL, nil)
- if err != nil {
- return nil, err
- }
- req.Header.Set("User-Agent", "MediaStationGo/0.1")
- resp, err := http.DefaultClient.Do(req)
- if err != nil {
- return nil, err
- }
- defer resp.Body.Close()
- if resp.StatusCode >= 400 {
- return nil, fmt.Errorf("rss %s: %d", feedURL, resp.StatusCode)
- }
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, err
- }
- var f rssFeed
- if err := xml.Unmarshal(body, &f); err != nil {
- return nil, err
- }
- return &f, nil
-}
-
-func compileFilter(pat string) *regexp.Regexp {
- pat = strings.TrimSpace(pat)
- if pat == "" {
- return nil
- }
- if r, err := regexp.Compile("(?i)" + pat); err == nil {
- return r
- }
- return nil
-}
-
-func splitNonEmpty(s string) []string {
- if s == "" {
- return nil
- }
- out := make([]string, 0)
- for _, p := range strings.Split(s, "\n") {
- p = strings.TrimSpace(p)
- if p != "" {
- out = append(out, p)
- }
- }
- return out
-}
diff --git a/internal/service/subscription_archive.go b/internal/service/subscription_archive.go
new file mode 100644
index 0000000..7080802
--- /dev/null
+++ b/internal/service/subscription_archive.go
@@ -0,0 +1,134 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// History returns completed/archived subscription rules.
+func (s *SubscriptionService) History(ctx context.Context) ([]model.Subscription, error) {
+ return s.repo.Subscription.History(ctx)
+}
+
+// Restore moves an archived subscription back to the active management list.
+// It also clears the per-subscription seen state so an unfinished historical
+// rule can match resources again when it is run next.
+func (s *SubscriptionService) Restore(ctx context.Context, id string) (*model.Subscription, error) {
+ var sub model.Subscription
+ if err := s.repo.DB.WithContext(ctx).Where("id = ?", id).First(&sub).Error; err != nil {
+ return nil, err
+ }
+ if err := s.repo.DB.WithContext(ctx).Model(&model.Subscription{}).
+ Where("id = ?", id).
+ Updates(map[string]any{
+ "enabled": true,
+ "archive_reason": "",
+ // 重置为 0:此前可能被 feed 低估并锁死(updateSubscriptionTotalEpisodes
+ // 只增不减,resolveSubscriptionTotalEpisodes 见 >0 即不再回查元数据)。
+ // 归零后下次 run 会从 TMDb/豆瓣等权威源重算真实总集数,避免恢复后
+ // 因"误判已无缺集"而不再搜索资源。
+ "total_episodes": 0,
+ }).Error; err != nil {
+ return nil, err
+ }
+ if err := s.repo.DB.WithContext(ctx).
+ Exec("UPDATE subscriptions SET archived_at = NULL WHERE id = ?", id).Error; err != nil {
+ return nil, err
+ }
+ if s.repo.Setting != nil {
+ _ = s.repo.Setting.Delete(ctx, fmt.Sprintf("subscription.%s.seen", id))
+ }
+ var restored model.Subscription
+ if err := s.repo.DB.WithContext(ctx).Where("id = ?", id).First(&restored).Error; err != nil {
+ return nil, err
+ }
+ return &restored, nil
+}
+
+func (s *SubscriptionService) archiveCompletedSubscription(ctx context.Context, sub *model.Subscription, availability LocalAvailability) error {
+ if s == nil || s.repo == nil || s.repo.Subscription == nil || sub == nil {
+ return nil
+ }
+ if !subscriptionShouldArchive(sub, availability) {
+ return nil
+ }
+ now := time.Now()
+ reason := subscriptionArchiveReason(sub, availability)
+ if err := s.repo.Subscription.Archive(ctx, sub.ID, reason, now); err != nil {
+ return err
+ }
+ sub.Enabled = false
+ sub.ArchivedAt = &now
+ sub.ArchiveReason = reason
+ if s.log != nil {
+ s.log.Info("subscription completed, moved to history",
+ zap.String("id", sub.ID),
+ zap.String("name", sub.Name),
+ zap.String("reason", reason))
+ }
+ if s.hub != nil {
+ s.hub.Publish("subscription", map[string]any{
+ "id": sub.ID,
+ "name": sub.Name,
+ "archived": true,
+ "reason": reason,
+ })
+ }
+ return nil
+}
+
+func subscriptionShouldArchive(sub *model.Subscription, availability LocalAvailability) bool {
+ if sub == nil || sub.WashEnabled || sub.ArchivedAt != nil {
+ return false
+ }
+ mediaType := strings.ToLower(strings.TrimSpace(sub.MediaType))
+ if !isSubscriptionSeriesType(mediaType) {
+ return availability.InLibrary || availability.LocalMediaCount > 0 || availability.DownloadedEpisodes > 0
+ }
+ if availability.HasSeriesPack {
+ return true
+ }
+ total := sub.TotalEpisodes
+ if total <= 0 {
+ total = availability.TotalEpisodes
+ }
+ if total > 0 {
+ return availability.DownloadedEpisodes >= total && len(availability.MissingEpisodes) == 0
+ }
+ return subscriptionLooksSingleEpisode(sub) && availability.DownloadedEpisodes > 0
+}
+
+func subscriptionArchiveReason(sub *model.Subscription, availability LocalAvailability) string {
+ if sub != nil && sub.WashEnabled {
+ return ""
+ }
+ if availability.HasSeriesPack {
+ return "整季资源已加入下载/入库"
+ }
+ if availability.TotalEpisodes > 0 {
+ return fmt.Sprintf("订阅完成:%d/%d", availability.DownloadedEpisodes, availability.TotalEpisodes)
+ }
+ if availability.DownloadedEpisodes > 0 {
+ return "单集订阅已加入下载/入库"
+ }
+ return "订阅媒体已加入下载/入库"
+}
+
+func subscriptionLooksSingleEpisode(sub *model.Subscription) bool {
+ if sub == nil {
+ return false
+ }
+ for _, value := range []string{sub.Name, sub.Filter} {
+ _, episode := ParseEpisode(value)
+ if episode > 0 {
+ return true
+ }
+ }
+ return false
+}
diff --git a/internal/service/subscription_availability_test.go b/internal/service/subscription_availability_test.go
new file mode 100644
index 0000000..66a60a3
--- /dev/null
+++ b/internal/service/subscription_availability_test.go
@@ -0,0 +1,143 @@
+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 TestSubscriptionEnrichProgressIncludesPendingDownloads(t *testing.T) {
+ db := newServiceTestDB(t, &model.DownloadTask{}, &model.Media{})
+ repos := repository.New(db)
+ if err := repos.Download.Create(t.Context(), &model.DownloadTask{
+ Source: "qbittorrent",
+ URL: "magnet:?xt=urn:btih:4444444444444444444444444444444444444444",
+ Title: "Inception 2010 1080p",
+ SavePath: "/downloads/movies",
+ Status: "completed",
+ Progress: 1,
+ }); err != nil {
+ t.Fatal(err)
+ }
+ svc := NewSubscriptionService(nil, nil, repos, nil, nil, nil)
+ items := []model.Subscription{{
+ Name: "Inception 2010",
+ Filter: "Inception 2010",
+ MediaType: "movie",
+ SavePath: "/downloads/movies",
+ }}
+
+ svc.EnrichProgress(t.Context(), items)
+ if items[0].InLibrary {
+ t.Fatal("pending download should not be reported as in-library media")
+ }
+ if items[0].DownloadedEpisodes != 1 || items[0].LocalMediaCount != 1 || items[0].TotalEpisodes != 1 {
+ t.Fatalf("unexpected enriched progress: %+v", items[0])
+ }
+}
+
+func TestSubscriptionLocalAvailabilityMatchesMediaPath(t *testing.T) {
+ db := newServiceTestDB(t, &model.Media{})
+ repos := repository.New(db)
+ if err := db.Create(&model.Media{
+ Title: "Scraped English Title",
+ Path: "/media/电视剧/国产剧/凡人修仙传/Season 01/凡人修仙传 - S01E146.mkv",
+ SeasonNum: 1,
+ EpisodeNum: 146,
+ }).Error; err != nil {
+ t.Fatal(err)
+ }
+ sub := &model.Subscription{
+ Name: "凡人修仙传 年番",
+ Filter: "凡人修仙传",
+ MediaType: "tv",
+ TotalEpisodes: 146,
+ }
+
+ availability := SubscriptionLocalAvailability(t.Context(), repos, sub)
+ if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 146)]; !ok {
+ t.Fatalf("missing path-matched E146 key: %#v", availability.ExistingEpisodeKeys)
+ }
+ results := []SearchResult{
+ {Title: "凡人修仙传 年番 - 146 1080p", DownloadURL: "https://pt/download/146", Seeders: 80},
+ }
+ got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
+ if len(got) != 0 {
+ t.Fatalf("selected %#v, want none because path-matched local episode exists", got)
+ }
+}
+
+func TestSubscriptionPendingDownloadAvailabilityIgnoresDeletedTasks(t *testing.T) {
+ db := newServiceTestDB(t, &model.DownloadTask{})
+ repos := repository.New(db)
+ if err := repos.Download.Create(t.Context(), &model.DownloadTask{
+ Source: "qbittorrent",
+ URL: "magnet:?xt=urn:btih:3333333333333333333333333333333333333333",
+ Title: "间谍过家家 S01E02 1080p",
+ SavePath: "/downloads/tv",
+ Status: "deleted",
+ }); err != nil {
+ t.Fatal(err)
+ }
+ svc := NewSubscriptionService(nil, nil, repos, nil, nil, nil)
+ sub := &model.Subscription{
+ Name: "间谍过家家 自动订阅",
+ Filter: "间谍过家家",
+ MediaType: "tv",
+ SavePath: "/downloads/tv",
+ TotalEpisodes: 3,
+ }
+
+ availability := svc.pendingDownloadAvailability(t.Context(), sub)
+ if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 2)]; ok {
+ t.Fatalf("deleted E02 task should not count as available: %#v", availability.ExistingEpisodeKeys)
+ }
+ results := []SearchResult{
+ {Title: "间谍过家家 S01E02 1080p WEB-DL", DownloadURL: "https://pt/download/2", Seeders: 80},
+ {Title: "间谍过家家 S01E03 1080p WEB-DL", DownloadURL: "https://pt/download/3", Seeders: 70},
+ }
+ got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
+ if len(got) != 2 || got[0].Episode != 2 || got[1].Episode != 3 {
+ t.Fatalf("selected %#v, want deleted episode 2 and new episode 3", got)
+ }
+}
+
+func TestSubscriptionPendingDownloadAvailabilityIncludesLiveQBTorrents(t *testing.T) {
+ 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":"间谍过家家 S01E01 1080p","state":"downloading","progress":0.2}]`))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer qb.Close()
+
+ db := newServiceTestDB(t, &model.DownloadTask{})
+ repos := repository.New(db)
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ downloads.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"})
+ svc := NewSubscriptionService(nil, nil, repos, downloads, nil, nil)
+ sub := &model.Subscription{
+ Name: "间谍过家家 自动订阅",
+ Filter: "间谍过家家",
+ MediaType: "tv",
+ SavePath: "/downloads/tv",
+ TotalEpisodes: 2,
+ }
+
+ availability := svc.pendingDownloadAvailability(t.Context(), sub)
+ if availability.DownloadedEpisodes != 1 {
+ t.Fatalf("downloaded episodes = %d, want 1", availability.DownloadedEpisodes)
+ }
+ if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 1)]; !ok {
+ t.Fatalf("missing live qB E01 key: %#v", availability.ExistingEpisodeKeys)
+ }
+}
diff --git a/internal/service/subscription_episode_totals.go b/internal/service/subscription_episode_totals.go
new file mode 100644
index 0000000..271d0dc
--- /dev/null
+++ b/internal/service/subscription_episode_totals.go
@@ -0,0 +1,279 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "regexp"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *SubscriptionService) updateSubscriptionTotalEpisodes(ctx context.Context, sub *model.Subscription, total int) {
+ if s == nil || s.repo == nil || s.repo.DB == nil || sub == nil || total <= sub.TotalEpisodes {
+ return
+ }
+ sub.TotalEpisodes = total
+ _ = s.repo.DB.WithContext(ctx).Model(sub).Update("total_episodes", total).Error
+}
+
+func inferRSSTotalEpisodes(items []rssItem, sub *model.Subscription, filter *regexp.Regexp) int {
+ if !subscriptionShouldInferTotal(sub) {
+ return 0
+ }
+ maxEpisode := 0
+ for _, item := range items {
+ title := strings.TrimSpace(item.Title)
+ if title == "" {
+ continue
+ }
+ if filter != nil && !filter.MatchString(title) {
+ continue
+ }
+ if !subscriptionTitleMatchesQuery(sub, title) {
+ continue
+ }
+ if !matchesSubscriptionRules(sub, title) {
+ continue
+ }
+ _, episode := ParseEpisode(title)
+ if episode > maxEpisode {
+ maxEpisode = episode
+ }
+ }
+ return maxEpisode
+}
+
+func inferSearchTotalEpisodes(results []SearchResult, sub *model.Subscription) int {
+ if !subscriptionShouldInferTotal(sub) {
+ return 0
+ }
+ maxEpisode := 0
+ for _, item := range results {
+ matchText := subscriptionSearchResultText(item)
+ if !subscriptionTitleMatchesQuery(sub, matchText) {
+ continue
+ }
+ if !matchesSubscriptionRules(sub, matchText) {
+ continue
+ }
+ _, episode := ParseEpisode(matchText)
+ if episode > maxEpisode {
+ maxEpisode = episode
+ }
+ }
+ return maxEpisode
+}
+
+func subscriptionShouldInferTotal(sub *model.Subscription) bool {
+ if sub == nil {
+ return false
+ }
+ mediaType := normalizeMediaType(sub.MediaType, sub.Name+" "+sub.Filter, "")
+ return isSubscriptionSeriesType(mediaType)
+}
+
+func (s *SubscriptionService) resolveSubscriptionTotalEpisodes(ctx context.Context, sub *model.Subscription, fallback int) int {
+ if !subscriptionShouldInferTotal(sub) {
+ return 0
+ }
+ if sub.TotalEpisodes > 0 {
+ return sub.TotalEpisodes
+ }
+ if total := s.resolveSubscriptionMetadataTotalEpisodes(ctx, sub); total > 0 {
+ return total
+ }
+ return fallback
+}
+
+func (s *SubscriptionService) resolveSubscriptionMetadataTotalEpisodes(ctx context.Context, sub *model.Subscription) int {
+ if s == nil || s.scraper == nil || sub == nil {
+ return 0
+ }
+ queries := subscriptionEpisodeMetadataQueries(sub)
+
+ // Priority: TMDb -> Douban -> Bangumi -> TheTVDB -> Fanart -> title fallback.
+ // Fanart.tv is artwork-only in MediaStationGo, so it intentionally does not
+ // claim episode counts and lets the title fallback handle the final layer.
+ if s.scraper.tmdb != nil {
+ if id := subscriptionExplicitTMDbID(sub); id > 0 {
+ if total, err := s.scraper.tmdb.GetTVEpisodeCount(ctx, id); err == nil && total > 0 {
+ return total
+ } else if err != nil && s.log != nil {
+ s.log.Debug("subscription tmdb episode count failed", zap.Int("tmdb_id", id), zap.Error(err))
+ }
+ }
+ for _, query := range queries {
+ match, err := s.scraper.tmdb.SearchTV(ctx, query, 0)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription tmdb search failed", zap.String("query", query), zap.Error(err))
+ }
+ continue
+ }
+ if match == nil || match.TMDbID <= 0 {
+ continue
+ }
+ total, err := s.scraper.tmdb.GetTVEpisodeCount(ctx, match.TMDbID)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription tmdb episode count failed", zap.Int("tmdb_id", match.TMDbID), zap.Error(err))
+ }
+ continue
+ }
+ if total > 0 {
+ return total
+ }
+ }
+ }
+
+ if s.scraper.douban != nil {
+ for _, query := range queries {
+ total, err := s.scraper.douban.GetEpisodeCount(ctx, query)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription douban episode count failed", zap.String("query", query), zap.Error(err))
+ }
+ continue
+ }
+ if total > 0 {
+ return total
+ }
+ }
+ }
+
+ if s.scraper.bangumi != nil {
+ for _, query := range queries {
+ match, err := s.scraper.bangumi.Search(ctx, query)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription bangumi search failed", zap.String("query", query), zap.Error(err))
+ }
+ continue
+ }
+ if match == nil || match.BangumiID <= 0 {
+ continue
+ }
+ total, err := s.scraper.bangumi.GetEpisodeCount(ctx, match.BangumiID)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription bangumi episode count failed", zap.Int("bangumi_id", match.BangumiID), zap.Error(err))
+ }
+ continue
+ }
+ if total > 0 {
+ return total
+ }
+ }
+ }
+
+ if s.scraper.thetvdb != nil {
+ for _, query := range queries {
+ match, err := s.scraper.thetvdb.SearchSeries(ctx, query)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription thetvdb search failed", zap.String("query", query), zap.Error(err))
+ }
+ continue
+ }
+ if match == nil || strings.TrimSpace(match.TheTVDBID) == "" {
+ continue
+ }
+ total, err := s.scraper.thetvdb.GetSeriesEpisodeCount(ctx, match.TheTVDBID)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("subscription thetvdb episode count failed", zap.String("thetvdb_id", match.TheTVDBID), zap.Error(err))
+ }
+ continue
+ }
+ if total > 0 {
+ return total
+ }
+ }
+ }
+
+ return 0
+}
+
+func subscriptionTitleMatchesQuery(sub *model.Subscription, title string) bool {
+ if strings.TrimSpace(title) == "" {
+ return false
+ }
+ for _, query := range subscriptionTitleMatchQueries(sub) {
+ if strings.Contains(normalizeAvailabilityComparable(title), normalizeAvailabilityComparable(query)) {
+ return true
+ }
+ }
+ return len(subscriptionTitleMatchQueries(sub)) == 0
+}
+
+func subscriptionTitleMatchQueries(sub *model.Subscription) []string {
+ if sub == nil {
+ return nil
+ }
+ values := []string{
+ availabilityQuery(subscriptionName(sub), subscriptionFilter(sub)),
+ cleanAvailabilityTitle(subscriptionFilter(sub)),
+ cleanAvailabilityTitle(subscriptionName(sub)),
+ }
+ for _, alias := range subscriptionFeedAliases(sub) {
+ values = append(values, alias, cleanAvailabilityTitle(alias))
+ }
+ for _, alias := range subscriptionMetadataAliases(sub) {
+ values = append(values, alias, cleanAvailabilityTitle(alias))
+ }
+ return compactUniqueStrings(values...)
+}
+
+func subscriptionEpisodeMetadataQueries(sub *model.Subscription) []string {
+ if sub == nil {
+ return nil
+ }
+ raw := []string{
+ siteSearchKeyword(sub),
+ sub.Filter,
+ sub.Name,
+ availabilityQuery(subscriptionName(sub), subscriptionFilter(sub)),
+ }
+ out := make([]string, 0, len(raw)*2)
+ for _, value := range raw {
+ value = cleanAvailabilityTitle(value)
+ if value == "" {
+ continue
+ }
+ if cleaned, _ := CleanQuery(value); cleaned != "" {
+ out = append(out, cleaned)
+ }
+ out = append(out, value)
+ }
+ return compactUniqueStrings(out...)
+}
+
+func subscriptionExplicitTMDbID(sub *model.Subscription) int {
+ if sub == nil {
+ return 0
+ }
+ values := []string{sub.Name, sub.Filter, sub.FeedURL}
+ for _, raw := range values {
+ for _, pattern := range []string{`(?i)\btmdb[_:\-\s=]+(\d{2,})`, `(?i)\btmdbid[_:\-\s=]+(\d{2,})`} {
+ if m := regexp.MustCompile(pattern).FindStringSubmatch(raw); len(m) >= 2 {
+ var id int
+ if _, err := fmt.Sscanf(m[1], "%d", &id); err == nil && id > 0 {
+ return id
+ }
+ }
+ }
+ if u, err := url.Parse(raw); err == nil {
+ for _, key := range []string{"tmdb_id", "tmdb", "tmdbid"} {
+ var id int
+ if _, err := fmt.Sscanf(u.Query().Get(key), "%d", &id); err == nil && id > 0 {
+ return id
+ }
+ }
+ }
+ }
+ return 0
+}
diff --git a/internal/service/subscription_notification.go b/internal/service/subscription_notification.go
new file mode 100644
index 0000000..5a55db7
--- /dev/null
+++ b/internal/service/subscription_notification.go
@@ -0,0 +1,73 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *SubscriptionService) notifySubscriptionHit(sub *model.Subscription, queued int, resources []string) {
+ if s == nil || s.notify == nil || sub == nil || queued <= 0 {
+ return
+ }
+ body := fmt.Sprintf("订阅:%s\n新增资源:%d", sub.Name, queued)
+ if len(resources) > 0 {
+ body += "\n资源:\n- " + strings.Join(resources, "\n- ")
+ }
+ go func() {
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ data := map[string]interface{}{}
+ if strings.TrimSpace(sub.PosterURL) != "" {
+ data["poster_url"] = sub.PosterURL
+ }
+ if strings.TrimSpace(sub.BackdropURL) != "" {
+ data["backdrop_url"] = sub.BackdropURL
+ }
+ if strings.TrimSpace(sub.MediaType) != "" {
+ data["media_type"] = sub.MediaType
+ }
+ if strings.TrimSpace(sub.MediaCategory) != "" {
+ data["media_category"] = sub.MediaCategory
+ }
+ // 补充媒体通知模板(formatTelegramMediaNotification)所需字段:片名 / 原名 /
+ // 语言 / 年份 / 评分 / 类型 / 简介 / 外链 / 资源标题(供模板提取季集 + 版本)。
+ // 仅填现成可用的,缺失项模板会自动略过。
+ if strings.TrimSpace(sub.Name) != "" {
+ data["title"] = sub.Name
+ }
+ if strings.TrimSpace(sub.OriginalName) != "" {
+ data["original_title"] = sub.OriginalName
+ }
+ if strings.TrimSpace(sub.OriginalLanguage) != "" {
+ data["original_language"] = sub.OriginalLanguage
+ }
+ if sub.Year > 0 {
+ data["year"] = sub.Year
+ }
+ if sub.Rating > 0 {
+ data["rating"] = sub.Rating
+ }
+ if strings.TrimSpace(sub.Genres) != "" {
+ data["genres"] = sub.Genres
+ }
+ if strings.TrimSpace(sub.Overview) != "" {
+ data["overview"] = sub.Overview
+ }
+ if id := strings.TrimSpace(sub.IMDBID); id != "" {
+ data["imdb_url"] = "https://www.imdb.com/title/" + id + "/"
+ }
+ if len(resources) > 0 {
+ data["resource_title"] = resources[0]
+ }
+ s.notify.BroadcastEvent(ctx, NotifyEvent{
+ Type: EventSubscriptionHit,
+ Title: "MediaStationGo 订阅命中新资源",
+ Message: body,
+ Data: data,
+ })
+ }()
+}
diff --git a/internal/service/subscription_planner.go b/internal/service/subscription_planner.go
index 098e878..20da6f8 100644
--- a/internal/service/subscription_planner.go
+++ b/internal/service/subscription_planner.go
@@ -30,6 +30,7 @@ type siteSearchCandidate struct {
type siteSearchSelectionStats struct {
Total int
QueryMismatch int
+ QueryMismatchExamples []string
RelaxedQueryMatch int
RuleMismatch int
MissingDownload int
@@ -105,6 +106,7 @@ func collectSiteSearchCandidates(results []SearchResult, sub *model.Subscription
stats.RelaxedQueryMatch++
} else {
stats.QueryMismatch++
+ stats.QueryMismatchExamples = appendLimitedStrings(stats.QueryMismatchExamples, matchText, 5)
continue
}
}
@@ -141,6 +143,14 @@ func collectSiteSearchCandidates(results []SearchResult, sub *model.Subscription
return candidates
}
+func appendLimitedStrings(values []string, value string, limit int) []string {
+ value = strings.TrimSpace(value)
+ if value == "" || limit <= 0 || len(values) >= limit {
+ return values
+ }
+ return append(values, value)
+}
+
func shouldRelaxSiteSearchQueryMatch(sub *model.Subscription, local LocalAvailability) bool {
if sub == nil {
return false
diff --git a/internal/service/subscription_rss.go b/internal/service/subscription_rss.go
new file mode 100644
index 0000000..f765551
--- /dev/null
+++ b/internal/service/subscription_rss.go
@@ -0,0 +1,79 @@
+package service
+
+import (
+ "context"
+ "encoding/xml"
+ "fmt"
+ "io"
+ "net/http"
+ "regexp"
+ "strings"
+)
+
+// rssFeed is the minimal RSS subset we need to decode.
+type rssFeed struct {
+ XMLName xml.Name `xml:"rss"`
+ Channel struct {
+ Items []rssItem `xml:"item"`
+ } `xml:"channel"`
+}
+
+type rssItem struct {
+ Title string `xml:"title"`
+ Link string `xml:"link"`
+ GUID string `xml:"guid"`
+ Description string `xml:"description"`
+ Enclosure struct {
+ URL string `xml:"url,attr"`
+ } `xml:"enclosure"`
+}
+
+func (s *SubscriptionService) fetch(ctx context.Context, feedURL string) (*rssFeed, error) {
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, feedURL, nil)
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("User-Agent", "MediaStationGo/0.1")
+ resp, err := http.DefaultClient.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+ if resp.StatusCode >= 400 {
+ return nil, fmt.Errorf("rss %s: %d", feedURL, resp.StatusCode)
+ }
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ return nil, err
+ }
+ var f rssFeed
+ if err := xml.Unmarshal(body, &f); err != nil {
+ return nil, err
+ }
+ return &f, nil
+}
+
+func compileFilter(pat string) *regexp.Regexp {
+ pat = strings.TrimSpace(pat)
+ if pat == "" {
+ return nil
+ }
+ if r, err := regexp.Compile("(?i)" + pat); err == nil {
+ return r
+ }
+ return nil
+}
+
+func splitNonEmpty(s string) []string {
+ if s == "" {
+ return nil
+ }
+ out := make([]string, 0)
+ for _, p := range strings.Split(s, "\n") {
+ p = strings.TrimSpace(p)
+ if p != "" {
+ out = append(out, p)
+ }
+ }
+ return out
+}
diff --git a/internal/service/subscription_rules_test.go b/internal/service/subscription_rules_test.go
new file mode 100644
index 0000000..b343a92
--- /dev/null
+++ b/internal/service/subscription_rules_test.go
@@ -0,0 +1,100 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func TestMatchesSubscriptionRulesUserExcludeWords(t *testing.T) {
+ sub := &model.Subscription{ExcludeWords: "10bit,dolby vision,杜比"}
+ cases := []struct {
+ title string
+ want bool
+ }{
+ {"Movie 2024 1080p WEB-DL", true},
+ {"Movie 2024 2160p 10bit HEVC", false},
+ {"Movie 2024 2160p Dolby Vision", false},
+ {"电影 2024 杜比全景声", false},
+ }
+ for _, c := range cases {
+ if got := matchesSubscriptionRules(sub, c.title); got != c.want {
+ t.Errorf("matchesSubscriptionRules(%q) = %v, want %v", c.title, got, c.want)
+ }
+ }
+}
+
+func TestMatchesSubscriptionRulesDefaultExcludesJunkReleases(t *testing.T) {
+ sub := &model.Subscription{}
+ for _, title := range []string{
+ "Some Movie 2024 CAM",
+ "Some Movie 2024 HDTS",
+ "某电影 2024 枪版",
+ "Some Movie 2024 TELESYNC",
+ "Some Show 预告",
+ } {
+ if matchesSubscriptionRules(sub, title) {
+ t.Errorf("expected default rules to exclude junk release %q", title)
+ }
+ }
+}
+
+func TestMatchesSubscriptionRulesWordBoundaryAvoidsFalsePositives(t *testing.T) {
+ sub := &model.Subscription{}
+ // "ts" / "cam" / "tc" 作为子串出现在合法标题里时不应被默认排除误伤。
+ for _, title := range []string{
+ "Tsukihime 2024 1080p WEB-DL",
+ "Camp Rock 2024 1080p BluRay",
+ "Catch Me 2024 1080p WEB-DL",
+ } {
+ if !matchesSubscriptionRules(sub, title) {
+ t.Errorf("word-boundary match wrongly excluded %q", title)
+ }
+ }
+}
+
+func TestSelectSiteSearchCandidatesSkipsExistingMovieWhenNotWashing(t *testing.T) {
+ sub := &model.Subscription{Name: "Inception 自动订阅", Filter: "Inception 2010", MediaType: "movie"}
+ results := []SearchResult{
+ {Title: "Inception 2010 2160p 10bit Dolby Vision Atmos", DownloadURL: "https://pt/download/dovi", Seeders: 500},
+ {Title: "Inception 2010 1080p WEB-DL", DownloadURL: "https://pt/download/web", Seeders: 90},
+ }
+ availability := LocalAvailability{LocalMediaCount: 1, InLibrary: true, DownloadedEpisodes: 1, TotalEpisodes: 1}
+
+ got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
+ if len(got) != 0 {
+ t.Fatalf("selected %#v, want none (movie already in library, wash disabled)", got)
+ }
+}
+
+func TestSelectSiteSearchCandidatesAllowsMovieWashUpgrade(t *testing.T) {
+ sub := &model.Subscription{Name: "Inception 自动订阅", Filter: "Inception 2010", MediaType: "movie", WashEnabled: true, WashPriority: "resolution"}
+ results := []SearchResult{
+ {Title: "Inception 2010 2160p REMUX", DownloadURL: "https://pt/download/2160", Seeders: 80},
+ {Title: "Inception 2010 1080p WEB-DL", DownloadURL: "https://pt/download/1080", Seeders: 200},
+ }
+ availability := LocalAvailability{LocalMediaCount: 1, InLibrary: true, DownloadedEpisodes: 1, TotalEpisodes: 1}
+
+ got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
+ if len(got) != 1 || got[0].Download != "https://pt/download/2160" {
+ t.Fatalf("selected %#v, want 2160p upgrade allowed when washing", got)
+ }
+}
+
+func TestSubscriptionItemAlreadyAvailable(t *testing.T) {
+ movieSub := &model.Subscription{MediaType: "movie"}
+ if !subscriptionItemAlreadyAvailable(movieSub, LocalAvailability{LocalMediaCount: 1}, "Inception 2010 2160p") {
+ t.Fatal("movie already in library should be reported available")
+ }
+ if subscriptionItemAlreadyAvailable(movieSub, LocalAvailability{}, "Inception 2010 2160p") {
+ t.Fatal("empty library should not be reported available")
+ }
+ tvSub := &model.Subscription{MediaType: "tv"}
+ avail := LocalAvailability{LocalMediaCount: 1, ExistingEpisodeKeys: map[string]struct{}{episodeKey(1, 2): {}}}
+ if !subscriptionItemAlreadyAvailable(tvSub, avail, "Show S01E02 1080p") {
+ t.Fatal("existing episode should be reported available")
+ }
+ if subscriptionItemAlreadyAvailable(tvSub, avail, "Show S01E03 1080p") {
+ t.Fatal("missing episode should not be reported available")
+ }
+}
diff --git a/internal/service/subscription_run_test.go b/internal/service/subscription_run_test.go
new file mode 100644
index 0000000..6b80b0e
--- /dev/null
+++ b/internal/service/subscription_run_test.go
@@ -0,0 +1,610 @@
+package service
+
+import (
+ "fmt"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/config"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "github.com/ShukeBta/MediaStationGo/internal/repository"
+)
+
+func TestSubscriptionRunOneArchivesCompletedMovieRSS(t *testing.T) {
+ rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/rss+xml")
+ _, _ = w.Write([]byte(`
+
+ -
+ Dune 2021 1080p WEB-DL
+ dune-1080-web
+ magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Dune+2021+1080p+WEB-DL
+
+`))
+ }))
+ defer rss.Close()
+
+ var addCalls int32
+ var added bool
+ 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 added {
+ _, _ = w.Write([]byte(`[{"hash":"dunehash","name":"Dune 2021 1080p WEB-DL","state":"downloading","progress":0.1}]`))
+ return
+ }
+ _, _ = w.Write([]byte(`[]`))
+ case "/api/v2/torrents/add":
+ added = true
+ atomic.AddInt32(&addCalls, 1)
+ _, _ = w.Write([]byte("Ok."))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer qb.Close()
+
+ db := newServiceTestDB(t, &model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{})
+ repos := repository.New(db)
+ configureTestDefaultQB(t, repos, qb.URL)
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
+
+ sub := &model.Subscription{
+ Name: "Dune 自动订阅",
+ FeedURL: rss.URL,
+ Filter: "Dune 2021",
+ MediaType: "movie",
+ SavePath: "/downloads/movies",
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+ queued, err := svc.runOne(t.Context(), sub)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if queued != 1 {
+ t.Fatalf("queued = %d, want 1", queued)
+ }
+ if got := atomic.LoadInt32(&addCalls); got != 1 {
+ t.Fatalf("qb add calls = %d, want 1", got)
+ }
+ active, err := repos.Subscription.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(active) != 0 {
+ t.Fatalf("active subscriptions = %d, want 0 after completion", len(active))
+ }
+ history, err := repos.Subscription.History(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(history) != 1 || history[0].ArchivedAt == nil {
+ t.Fatalf("history = %#v, want one archived subscription", history)
+ }
+}
+
+func TestSubscriptionArchiveCompletedSingleEpisodeTV(t *testing.T) {
+ db := newServiceTestDB(t, &model.Subscription{})
+ repos := repository.New(db)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
+ sub := &model.Subscription{
+ Name: "Some Show S01E01 自动订阅",
+ FeedURL: "site-search://search?keyword=Some%20Show%20S01E01",
+ Filter: "Some Show S01E01",
+ MediaType: "tv",
+ Enabled: true,
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+
+ if err := svc.archiveCompletedSubscription(t.Context(), sub, LocalAvailability{
+ DownloadedEpisodes: 1,
+ LocalMediaCount: 1,
+ InLibrary: true,
+ ExistingEpisodeKeys: map[string]struct{}{
+ episodeKey(1, 1): {},
+ },
+ }); err != nil {
+ t.Fatal(err)
+ }
+ active, err := repos.Subscription.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(active) != 0 {
+ t.Fatalf("active subscriptions = %d, want 0", len(active))
+ }
+ history, err := repos.Subscription.History(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(history) != 1 || history[0].ArchiveReason == "" {
+ t.Fatalf("history = %#v, want archived single episode", history)
+ }
+}
+
+func TestSubscriptionArchiveKeepsGenericUnknownTotalSeriesActive(t *testing.T) {
+ sub := &model.Subscription{
+ Name: "Some Show 自动订阅",
+ Filter: "Some Show",
+ MediaType: "tv",
+ }
+ availability := LocalAvailability{
+ DownloadedEpisodes: 1,
+ LocalMediaCount: 1,
+ InLibrary: true,
+ ExistingEpisodeKeys: map[string]struct{}{
+ episodeKey(1, 1): {},
+ },
+ }
+ if subscriptionShouldArchive(sub, availability) {
+ t.Fatal("generic series with unknown total should stay active for incremental episodes")
+ }
+}
+
+func TestInferSubscriptionTotalEpisodesFromSearchAndRSS(t *testing.T) {
+ sub := &model.Subscription{Name: "Some Show 自动订阅", Filter: "Some Show", MediaType: "tv"}
+ results := []SearchResult{
+ {Title: "Some Show S01E01 1080p"},
+ {Title: "Some Show S01E12 1080p"},
+ {Title: "Other Show S01E99 1080p"},
+ }
+ if got := inferSearchTotalEpisodes(results, sub); got != 12 {
+ t.Fatalf("search inferred total = %d, want 12", got)
+ }
+ subtitleResults := []SearchResult{
+ {Title: "Smoking Behind the Supermarket with You", Subtitle: "躲在超市后门抽烟的两人 S01E12"},
+ }
+ subtitleSub := &model.Subscription{Name: "躲在超市后门抽烟的两人 自动订阅", Filter: "躲在超市后门抽烟的两人", MediaType: "tv"}
+ if got := inferSearchTotalEpisodes(subtitleResults, subtitleSub); got != 12 {
+ t.Fatalf("subtitle search inferred total = %d, want 12", got)
+ }
+ items := []rssItem{
+ {Title: "Some Show S01E02 WEB-DL"},
+ {Title: "Some Show S01E10 WEB-DL"},
+ }
+ if got := inferRSSTotalEpisodes(items, sub, compileFilter("Some Show")); got != 10 {
+ t.Fatalf("rss inferred total = %d, want 10", got)
+ }
+}
+
+func TestResolveSubscriptionTotalEpisodesPrefersTMDbOverTitleFallback(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.URL.Path {
+ case "/search/tv":
+ _, _ = w.Write([]byte(`{"results":[{"id":42,"name":"Some Show","first_air_date":"2026-01-01"}]}`))
+ case "/tv/42":
+ _, _ = w.Write([]byte(`{"number_of_episodes":13}`))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer upstream.Close()
+
+ cfg := &config.Config{}
+ cfg.Secrets.TMDbAPIKey = "test"
+ cfg.Secrets.TMDbAPIProxy = upstream.URL
+ tmdb := NewTMDbProvider(cfg, zap.NewNop(), nil)
+ svc := NewSubscriptionService(cfg, zap.NewNop(), nil, nil, nil, NewHub(zap.NewNop()))
+ svc.SetScraper(NewScraperService(cfg, zap.NewNop(), nil, tmdb, nil, nil, nil, NewHub(zap.NewNop())))
+
+ sub := &model.Subscription{Name: "Some Show 自动订阅", Filter: "Some Show", MediaType: "tv"}
+ if got := svc.resolveSubscriptionTotalEpisodes(t.Context(), sub, 10); got != 13 {
+ t.Fatalf("resolved total = %d, want TMDb total 13", got)
+ }
+}
+
+func TestSubscriptionArchiveKeepsWashSubscriptionActive(t *testing.T) {
+ db := newServiceTestDB(t, &model.Subscription{})
+ repos := repository.New(db)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
+ sub := &model.Subscription{
+ Name: "Dune 自动订阅",
+ FeedURL: "site-search://search?keyword=Dune",
+ Filter: "Dune 2021",
+ MediaType: "movie",
+ WashEnabled: true,
+ Enabled: true,
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+
+ if err := svc.archiveCompletedSubscription(t.Context(), sub, LocalAvailability{
+ DownloadedEpisodes: 1,
+ LocalMediaCount: 1,
+ InLibrary: true,
+ }); err != nil {
+ t.Fatal(err)
+ }
+ active, err := repos.Subscription.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(active) != 1 {
+ t.Fatalf("active subscriptions = %d, want wash subscription to stay active", len(active))
+ }
+ history, err := repos.Subscription.History(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(history) != 0 {
+ t.Fatalf("history subscriptions = %d, want 0", len(history))
+ }
+}
+
+func TestRestoreArchivedSubscriptionReturnsToActiveAndClearsSeenState(t *testing.T) {
+ db := newServiceTestDB(t, &model.Subscription{}, &model.Setting{})
+ repos := repository.New(db)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
+ sub := &model.Subscription{
+ Name: "南部档案 自动订阅",
+ FeedURL: "https://rss.example/feed",
+ Filter: "南部档案",
+ MediaType: "tv",
+ TotalEpisodes: 33,
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+ archivedAt := time.Now()
+ if err := repos.Subscription.Archive(t.Context(), sub.ID, "已下载 1/33 集,缺 33 集", archivedAt); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.Setting.Set(t.Context(), "subscription."+sub.ID+".seen", "old-guid"); err != nil {
+ t.Fatal(err)
+ }
+ restored, err := svc.Restore(t.Context(), sub.ID)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if restored.ArchivedAt != nil || restored.ArchiveReason != "" || !restored.Enabled {
+ t.Fatalf("restored subscription not active: archived=%v reason=%q enabled=%v", restored.ArchivedAt, restored.ArchiveReason, restored.Enabled)
+ }
+ if restored.TotalEpisodes != 0 {
+ t.Fatalf("restored total_episodes = %d, want 0 so it gets recomputed from authoritative metadata", restored.TotalEpisodes)
+ }
+ active, err := repos.Subscription.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(active) != 1 || active[0].ID != sub.ID {
+ t.Fatalf("active subscriptions = %#v, want restored subscription", active)
+ }
+ history, err := repos.Subscription.History(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(history) != 0 {
+ t.Fatalf("history subscriptions = %d, want 0 after restore", len(history))
+ }
+ seen, err := repos.Setting.Get(t.Context(), "subscription."+sub.ID+".seen")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if seen != "" {
+ t.Fatalf("seen state = %q, want cleared", seen)
+ }
+}
+
+func TestSubscriptionRunOneDeduplicatesDuplicateRSSGUIDInSameFeed(t *testing.T) {
+ rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/rss+xml")
+ _, _ = w.Write([]byte(`
+
+ -
+ Some Show S01E01 1080p
+ episode-1
+ magnet:?xt=urn:btih:1111111111111111111111111111111111111111&dn=Some+Show+S01E01
+
+ -
+ Some Show S01E01 1080p
+ episode-1
+ magnet:?xt=urn:btih:1111111111111111111111111111111111111111&dn=Some+Show+S01E01
+
+`))
+ }))
+ defer rss.Close()
+
+ 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":"Some Show S01E01 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.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{})
+ repos := repository.New(db)
+ configureTestDefaultQB(t, repos, qb.URL)
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
+
+ sub := &model.Subscription{
+ Name: "Some Show 自动订阅",
+ FeedURL: rss.URL,
+ Filter: "Some Show",
+ MediaType: "tv",
+ SavePath: "/downloads/tv",
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+ queued, err := svc.runOne(t.Context(), sub)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if queued != 1 {
+ t.Fatalf("queued = %d, want 1", queued)
+ }
+ if got := atomic.LoadInt32(&addCalls); got != 1 {
+ t.Fatalf("qb add calls = %d, want 1", got)
+ }
+ rows, err := repos.Download.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 1 {
+ t.Fatalf("download rows = %d, want 1", len(rows))
+ }
+}
+
+func TestSubscriptionRunOneSkipsSameEpisodeAddedEarlierInFeed(t *testing.T) {
+ rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/rss+xml")
+ _, _ = w.Write([]byte(`
+
+ -
+ Some Show S01E01 1080p
+ episode-1-a
+ magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Some+Show+S01E01+1080p
+
+ -
+ Some Show S01E01 WEB-DL
+ episode-1-b
+ magnet:?xt=urn:btih:bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb&dn=Some+Show+S01E01+WEB-DL
+
+`))
+ }))
+ defer rss.Close()
+
+ 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":"Some Show S01E01 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.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{})
+ repos := repository.New(db)
+ configureTestDefaultQB(t, repos, qb.URL)
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
+
+ sub := &model.Subscription{
+ Name: "Some Show 自动订阅",
+ FeedURL: rss.URL,
+ Filter: "Some Show",
+ MediaType: "tv",
+ SavePath: "/downloads/tv",
+ TotalEpisodes: 12,
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+ queued, err := svc.runOne(t.Context(), sub)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if queued != 1 {
+ t.Fatalf("queued = %d, want 1", queued)
+ }
+ if got := atomic.LoadInt32(&addCalls); got != 1 {
+ t.Fatalf("qb add calls = %d, want 1", got)
+ }
+ rows, err := repos.Download.List(t.Context())
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(rows) != 1 {
+ t.Fatalf("download rows = %d, want 1", len(rows))
+ }
+}
+
+func TestSubscriptionRunOneRSSWashQueuesOnlyBestMovieVariant(t *testing.T) {
+ rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/rss+xml")
+ _, _ = w.Write([]byte(`
+
+ -
+ Dune 2021 1080p WEB-DL
+ dune-1080-web
+ magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Dune+2021+1080p+WEB-DL
+
+ -
+ Dune 2021 2160p UHD BluRay REMUX HDR
+ dune-2160-remux
+ magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=Dune+2021+2160p+REMUX
+
+ -
+ Dune 2021 720p HDTV
+ dune-720-hdtv
+ magnet:?xt=urn:btih:ffffffffffffffffffffffffffffffffffffffff&dn=Dune+2021+720p+HDTV
+
+`))
+ }))
+ defer rss.Close()
+
+ var addCalls int32
+ var addedTitles []string
+ addedHashes := make([]string, 0, 3)
+ 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 len(addedHashes) == 0 {
+ _, _ = w.Write([]byte(`[]`))
+ return
+ }
+ var items []string
+ for _, hash := range addedHashes {
+ items = append(items, `{"hash":"`+hash+`","name":"Dune 2021","state":"downloading","progress":0.1}`)
+ }
+ _, _ = w.Write([]byte(`[` + strings.Join(items, ",") + `]`))
+ case "/api/v2/torrents/add":
+ call := atomic.AddInt32(&addCalls, 1)
+ _ = r.ParseMultipartForm(10 << 20)
+ addedTitles = append(addedTitles, r.FormValue("urls"))
+ addedHashes = append(addedHashes, strings.Repeat(fmt.Sprintf("%x", call), 40))
+ _, _ = w.Write([]byte("Ok."))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer qb.Close()
+
+ db := newServiceTestDB(t, &model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{})
+ repos := repository.New(db)
+ configureTestDefaultQB(t, repos, qb.URL)
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
+
+ sub := &model.Subscription{
+ Name: "Dune 自动订阅",
+ FeedURL: rss.URL,
+ Filter: "Dune 2021",
+ MediaType: "movie",
+ WashEnabled: true,
+ WashPriority: "resolution",
+ SavePath: "/downloads/movies",
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+ queued, err := svc.runOne(t.Context(), sub)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if queued != 1 {
+ t.Fatalf("queued = %d, want 1 best movie variant", queued)
+ }
+ if got := atomic.LoadInt32(&addCalls); got != 1 {
+ t.Fatalf("qb add calls = %d, want 1", got)
+ }
+ if len(addedTitles) != 1 || !strings.Contains(addedTitles[0], "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee") {
+ t.Fatalf("added %#v, want 2160p REMUX variant only", addedTitles)
+ }
+}
+
+func TestSubscriptionRunOneDoesNotUseDeletedDownloader(t *testing.T) {
+ rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/rss+xml")
+ _, _ = w.Write([]byte(`
+
+ -
+ Deleted Downloader Show S01E01 1080p
+ deleted-downloader-episode-1
+ magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Deleted+Downloader+Show+S01E01
+
+`))
+ }))
+ defer rss.Close()
+
+ var qbCalls int32
+ qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ atomic.AddInt32(&qbCalls, 1)
+ 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":
+ _, _ = w.Write([]byte("Ok."))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer qb.Close()
+
+ db := newServiceTestDB(t, &model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{})
+ repos := repository.New(db)
+ client := &model.DownloadClient{Name: "qB deleted", 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.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil {
+ t.Fatal(err)
+ }
+ if err := repos.DownloadClient.Delete(t.Context(), client.ID); err != nil {
+ t.Fatal(err)
+ }
+
+ downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
+ svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
+ sub := &model.Subscription{
+ Name: "Deleted Downloader Show 自动订阅",
+ FeedURL: rss.URL,
+ Filter: "Deleted Downloader Show",
+ MediaType: "tv",
+ SavePath: "/downloads/tv",
+ }
+ if err := repos.Subscription.Create(t.Context(), sub); err != nil {
+ t.Fatal(err)
+ }
+
+ queued, err := svc.runOne(t.Context(), sub)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if queued != 0 {
+ t.Fatalf("queued = %d, want 0 when default downloader was deleted", queued)
+ }
+ if got := atomic.LoadInt32(&qbCalls); got != 0 {
+ t.Fatalf("qB calls = %d, want 0 after downloader deletion", got)
+ }
+ 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))
+ }
+}
diff --git a/internal/service/subscription_site_search.go b/internal/service/subscription_site_search.go
new file mode 100644
index 0000000..ec83044
--- /dev/null
+++ b/internal/service/subscription_site_search.go
@@ -0,0 +1,246 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *SubscriptionService) runSiteSearch(ctx context.Context, sub *model.Subscription) (int, error) {
+ if s.site == nil {
+ if s.log != nil {
+ s.log.Warn("site-search subscription service unavailable", subscriptionSiteSearchLogFields(sub, "")...)
+ }
+ return 0, errors.New("site search service unavailable")
+ }
+ keywords := siteSearchKeywords(sub)
+ keyword := ""
+ if len(keywords) > 0 {
+ keyword = keywords[0]
+ }
+ if keyword == "" {
+ if s.log != nil {
+ s.log.Warn("site-search subscription keyword missing", subscriptionSiteSearchLogFields(sub, "")...)
+ }
+ return 0, errors.New("site-search subscription keyword required")
+ }
+ if s.log != nil {
+ s.log.Info("site-search subscription run started", subscriptionSiteSearchLogFields(sub, keyword)...)
+ }
+
+ var (
+ results []SearchResult
+ lastSearchErr error
+ searchErrors int
+ )
+ for _, searchKeyword := range keywords {
+ found, err := s.site.Search(ctx, searchKeyword)
+ if err != nil {
+ lastSearchErr = err
+ searchErrors++
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, searchKeyword)
+ fields = append(fields, zap.Error(err))
+ s.log.Warn("site-search subscription search failed", fields...)
+ }
+ continue
+ }
+ results = append(results, found...)
+ }
+ results = dedupeSiteSearchResults(results)
+ if len(results) == 0 && lastSearchErr != nil && searchErrors == len(keywords) {
+ return 0, lastSearchErr
+ }
+ if len(results) == 0 {
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = append(fields, zap.Int("results_count", 0))
+ s.log.Info("site-search subscription no results", fields...)
+ }
+ now := time.Now()
+ _ = s.repo.DB.Model(sub).Updates(map[string]any{"last_run_at": &now}).Error
+ return 0, nil
+ }
+ s.updateSubscriptionTotalEpisodes(ctx, sub, s.resolveSubscriptionTotalEpisodes(ctx, sub, inferSearchTotalEpisodes(results, sub)))
+
+ guidKey := fmt.Sprintf("subscription.%s.seen", sub.ID)
+ seenRaw, _ := s.repo.Setting.Get(ctx, guidKey)
+ seen := splitNonEmpty(seenRaw)
+ seenSet := make(map[string]struct{}, len(seen))
+ for _, g := range seen {
+ seenSet[g] = struct{}{}
+ }
+
+ availability := mergeLocalAvailability(
+ SubscriptionLocalAvailability(ctx, s.repo, sub),
+ s.pendingDownloadAvailability(ctx, sub),
+ )
+ candidates, selectionStats := selectSiteSearchCandidatesWithStats(results, sub, seenSet, availability)
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = appendSiteSearchSelectionLogFields(fields, selectionStats)
+ fields = appendAvailabilityLogFields(fields, availability)
+ s.log.Info("site-search subscription selection summary", fields...)
+ }
+ var lastEnqueueErr error
+ queued := 0
+ var resources []string
+ for _, candidate := range candidates {
+ item := candidate.Item
+ matchText := subscriptionSearchResultText(item)
+ mediaType, mediaCategory := s.classifySubscriptionItem(ctx, sub, matchText, item.Category)
+ if s.shouldSkipExistingTorrent(ctx, mediaType, candidate) {
+ addSiteSearchCandidateAvailability(candidate, &availability)
+ seen = append(seen, candidate.GUID)
+ seenSet[candidate.GUID] = struct{}{}
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = append(fields,
+ zap.String("reason", "existing_torrent"),
+ zap.String("title", item.Title),
+ zap.String("subtitle", item.Subtitle),
+ zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
+ zap.String("site_category", item.Category),
+ zap.Int("season", candidate.Season),
+ zap.Int("episode", candidate.Episode),
+ zap.Bool("pack", candidate.Pack),
+ zap.String("media_type", mediaType),
+ )
+ s.log.Info("site-search subscription candidate skipped", fields...)
+ }
+ continue
+ }
+ realURL := s.site.ResolveDownloadURL(ctx, candidate.Download)
+ savePath := s.resolveSubscriptionSavePath(ctx, sub, mediaType, mediaCategory)
+ if s.downloadPathHasCandidate(ctx, sub, matchText, savePath) {
+ addSiteSearchCandidateAvailability(candidate, &availability)
+ seen = append(seen, candidate.GUID)
+ seenSet[candidate.GUID] = struct{}{}
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = append(fields,
+ zap.String("reason", "download_path_has_candidate"),
+ zap.String("title", item.Title),
+ zap.String("subtitle", item.Subtitle),
+ zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
+ zap.String("site_category", item.Category),
+ zap.Int("season", candidate.Season),
+ zap.Int("episode", candidate.Episode),
+ zap.Bool("pack", candidate.Pack),
+ zap.String("media_type", mediaType),
+ zap.String("media_category", mediaCategory),
+ zap.String("save_path", savePath),
+ )
+ s.log.Info("site-search subscription candidate skipped", fields...)
+ }
+ continue
+ }
+ if _, err := s.downloads.AddDownloadWithMeta(ctx, sub.UserID, realURL, savePath, DownloadTaskMeta{
+ SubscriptionID: sub.ID,
+ Title: firstNonEmpty(item.Title, sub.Name),
+ PosterURL: sub.PosterURL,
+ BackdropURL: sub.BackdropURL,
+ Overview: sub.Overview,
+ MediaType: mediaType,
+ MediaCategory: mediaCategory,
+ SourceCategory: item.Category,
+ AllowExistingLibrary: sub.WashEnabled,
+ }); err != nil {
+ if IsDownloadDedupError(err) {
+ addSiteSearchCandidateAvailability(candidate, &availability)
+ seen = append(seen, candidate.GUID)
+ seenSet[candidate.GUID] = struct{}{}
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = append(fields,
+ zap.String("reason", "download_dedup"),
+ zap.String("title", item.Title),
+ zap.String("subtitle", item.Subtitle),
+ zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
+ zap.String("site_category", item.Category),
+ zap.Int("season", candidate.Season),
+ zap.Int("episode", candidate.Episode),
+ zap.Bool("pack", candidate.Pack),
+ zap.String("media_type", mediaType),
+ zap.String("media_category", mediaCategory),
+ zap.String("save_path", savePath),
+ )
+ s.log.Info("site-search subscription candidate skipped", fields...)
+ }
+ continue
+ }
+ lastEnqueueErr = err
+ s.log.Warn("site-search subscription enqueue failed",
+ zap.String("subscription_id", sub.ID),
+ zap.String("subscription", sub.Name),
+ zap.String("keyword", keyword),
+ zap.String("title", item.Title),
+ zap.String("subtitle", item.Subtitle),
+ zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
+ zap.String("site_category", item.Category),
+ zap.String("media_type", mediaType),
+ zap.String("media_category", mediaCategory),
+ zap.String("save_path", savePath),
+ zap.Error(err))
+ continue
+ }
+ queued++
+ addSiteSearchCandidateAvailability(candidate, &availability)
+ resources = append(resources, item.Title)
+ seen = append(seen, candidate.GUID)
+ seenSet[candidate.GUID] = struct{}{}
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = append(fields,
+ zap.String("title", item.Title),
+ zap.String("subtitle", item.Subtitle),
+ zap.String("site", firstNonEmpty(item.SiteName, item.SiteID)),
+ zap.String("site_category", item.Category),
+ zap.Int("season", candidate.Season),
+ zap.Int("episode", candidate.Episode),
+ zap.Bool("pack", candidate.Pack),
+ zap.Int("score", candidate.Score),
+ zap.String("media_type", mediaType),
+ zap.String("media_category", mediaCategory),
+ zap.String("save_path", savePath),
+ )
+ s.log.Info("site-search subscription candidate queued", fields...)
+ }
+ }
+ availability = s.finalizePendingAvailability(sub, availability)
+ if len(seen) > 200 {
+ seen = seen[len(seen)-200:]
+ }
+ _ = s.repo.Setting.Set(ctx, guidKey, strings.Join(seen, "\n"))
+ now := time.Now()
+ _ = s.repo.DB.Model(sub).Updates(map[string]any{"last_run_at": &now}).Error
+ _ = s.archiveCompletedSubscription(ctx, sub, availability)
+ if queued > 0 {
+ s.hub.Publish("subscription", map[string]any{
+ "id": sub.ID,
+ "name": sub.Name,
+ "queued": queued,
+ "keyword": keyword,
+ "resources": resources,
+ })
+ s.notifySubscriptionHit(sub, queued, resources)
+ return queued, nil
+ }
+ if lastEnqueueErr != nil {
+ return 0, fmt.Errorf("找到 PT 资源但加入下载器失败: %w", lastEnqueueErr)
+ }
+ if s.log != nil {
+ fields := subscriptionSiteSearchLogFields(sub, keyword)
+ fields = appendSiteSearchSelectionLogFields(fields, selectionStats)
+ fields = appendAvailabilityLogFields(fields, availability)
+ fields = append(fields, zap.Int("queued", queued))
+ s.log.Info("site-search subscription no candidate queued", fields...)
+ }
+ return 0, nil
+}
diff --git a/internal/service/subscription_site_search_helpers.go b/internal/service/subscription_site_search_helpers.go
new file mode 100644
index 0000000..ee86648
--- /dev/null
+++ b/internal/service/subscription_site_search_helpers.go
@@ -0,0 +1,197 @@
+package service
+
+import (
+ "context"
+ "net/url"
+ "strings"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func subscriptionSiteSearchLogFields(sub *model.Subscription, keyword string) []zap.Field {
+ fields := []zap.Field{zap.String("keyword", keyword), zap.Strings("search_keywords", siteSearchKeywords(sub))}
+ if sub == nil {
+ return fields
+ }
+ fields = append(fields,
+ zap.String("subscription_id", sub.ID),
+ zap.String("subscription", sub.Name),
+ zap.String("filter", sub.Filter),
+ zap.String("media_type", sub.MediaType),
+ zap.String("media_category", sub.MediaCategory),
+ zap.String("search_mode", sub.SearchMode),
+ zap.String("imdb_id", sub.IMDBID),
+ zap.Bool("wash_enabled", sub.WashEnabled),
+ zap.String("wash_priority", sub.WashPriority),
+ zap.Int("total_episodes", sub.TotalEpisodes),
+ )
+ return fields
+}
+
+func appendSiteSearchSelectionLogFields(fields []zap.Field, stats siteSearchSelectionStats) []zap.Field {
+ return append(fields,
+ zap.Int("results_count", stats.Total),
+ zap.Int("query_mismatch_count", stats.QueryMismatch),
+ zap.Strings("query_mismatch_examples", stats.QueryMismatchExamples),
+ zap.Int("relaxed_query_match_count", stats.RelaxedQueryMatch),
+ zap.Int("rule_mismatch_count", stats.RuleMismatch),
+ zap.Int("missing_download_count", stats.MissingDownload),
+ zap.Int("seen_count", stats.Seen),
+ zap.Int("prepared_count", stats.Prepared),
+ zap.Int("selected_count", stats.Selected),
+ zap.Bool("local_already_satisfied", stats.LocalAlreadySatisfied),
+ zap.Bool("local_series_pack_present", stats.LocalSeriesPackPresent),
+ zap.Bool("series_complete", stats.SeriesComplete),
+ zap.Int("existing_episode_skipped_count", stats.ExistingEpisodeSkipped),
+ zap.Int("not_missing_episode_skipped_count", stats.NotMissingEpisodeSkipped),
+ zap.Int("no_episode_skipped_count", stats.NoEpisodeSkipped),
+ zap.Bool("pack_fallback_available", stats.PackFallbackAvailable),
+ zap.Bool("pack_fallback_used", stats.PackFallbackUsed),
+ )
+}
+
+func appendAvailabilityLogFields(fields []zap.Field, availability LocalAvailability) []zap.Field {
+ missingSample, missingMore := limitedEpisodeSample(availability.MissingEpisodes, 20)
+ return append(fields,
+ zap.Int("local_media_count", availability.LocalMediaCount),
+ zap.Bool("in_library", availability.InLibrary),
+ zap.Bool("has_series_pack", availability.HasSeriesPack),
+ zap.Int("downloaded_episodes", availability.DownloadedEpisodes),
+ zap.Int("availability_total_episodes", availability.TotalEpisodes),
+ zap.Int("missing_episode_count", len(availability.MissingEpisodes)),
+ zap.Ints("missing_episodes", missingSample),
+ zap.Int("missing_episodes_more", missingMore),
+ )
+}
+
+func limitedEpisodeSample(values []int, limit int) ([]int, int) {
+ if limit <= 0 || len(values) == 0 {
+ return nil, len(values)
+ }
+ if len(values) <= limit {
+ out := append([]int(nil), values...)
+ return out, 0
+ }
+ out := append([]int(nil), values[:limit]...)
+ return out, len(values) - limit
+}
+
+func (s *SubscriptionService) shouldSkipExistingTorrent(ctx context.Context, mediaType string, candidate siteSearchCandidate) bool {
+ if s == nil || s.downloads == nil {
+ return false
+ }
+ if isSubscriptionSeriesType(mediaType) && !candidate.Pack && candidate.Episode > 0 {
+ return false
+ }
+ return s.downloads.TorrentExistsByName(ctx, candidate.Item.Title)
+}
+
+func siteSearchKeywords(sub *model.Subscription) []string {
+ if sub == nil {
+ return nil
+ }
+ values := make([]string, 0, 8)
+ if strings.EqualFold(strings.TrimSpace(sub.SearchMode), "imdb") && strings.TrimSpace(sub.IMDBID) != "" {
+ values = append(values, strings.TrimSpace(sub.IMDBID))
+ }
+ if u, err := url.Parse(sub.FeedURL); err == nil {
+ if keyword := strings.TrimSpace(u.Query().Get("keyword")); keyword != "" {
+ values = append(values, keyword)
+ }
+ }
+ if strings.TrimSpace(sub.Filter) != "" {
+ values = append(values, sub.Filter)
+ }
+ if len(values) == 0 && strings.TrimSpace(sub.Name) != "" {
+ values = append(values, sub.Name)
+ }
+ values = append(values, subscriptionFeedAliases(sub)...)
+ values = append(values, subscriptionMetadataAliases(sub)...)
+ for _, value := range append([]string(nil), values...) {
+ if cleaned := cleanAvailabilityTitle(value); cleaned != "" {
+ values = append(values, cleaned)
+ }
+ }
+ return compactUniqueStrings(values...)
+}
+
+func siteSearchKeyword(sub *model.Subscription) string {
+ keywords := siteSearchKeywords(sub)
+ if len(keywords) == 0 {
+ return ""
+ }
+ return keywords[0]
+}
+
+func subscriptionFeedAliases(sub *model.Subscription) []string {
+ if sub == nil {
+ return nil
+ }
+ u, err := url.Parse(sub.FeedURL)
+ if err != nil {
+ return nil
+ }
+ q := u.Query()
+ values := make([]string, 0, len(q["alias"])+2)
+ values = append(values, q["alias"]...)
+ for _, raw := range q["aliases"] {
+ for _, part := range strings.FieldsFunc(raw, func(r rune) bool {
+ return r == '|' || r == '\n' || r == '\r' || r == '\t'
+ }) {
+ values = append(values, part)
+ }
+ }
+ return compactUniqueStrings(values...)
+}
+
+func subscriptionMetadataAliases(sub *model.Subscription) []string {
+ if sub == nil {
+ return nil
+ }
+ title := cleanAvailabilityTitle(firstNonEmpty(sub.Filter, sub.Name))
+ return buildSubscribeAliases(title, sub.OriginalName, sub.Year)
+}
+
+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 dedupeSiteSearchResults(results []SearchResult) []SearchResult {
+ if len(results) < 2 {
+ return results
+ }
+ seen := make(map[string]struct{}, len(results))
+ out := make([]SearchResult, 0, len(results))
+ for _, item := range results {
+ download := strings.TrimSpace(item.DownloadURL)
+ if download == "" {
+ download = strings.TrimSpace(item.TorrentURL)
+ }
+ key := stableSiteSearchGUID(item, download)
+ if _, ok := seen[key]; ok {
+ continue
+ }
+ seen[key] = struct{}{}
+ out = append(out, item)
+ }
+ return out
+}
diff --git a/internal/service/subscription_test.go b/internal/service/subscription_test.go
index a0ca609..376ac17 100644
--- a/internal/service/subscription_test.go
+++ b/internal/service/subscription_test.go
@@ -1,7 +1,6 @@
package service
import (
- "fmt"
"net/http"
"net/http/httptest"
"os"
@@ -9,13 +8,9 @@ import (
"strings"
"sync/atomic"
"testing"
- "time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
- "github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
@@ -114,6 +109,24 @@ func TestSelectSiteSearchCandidatesMatchesFeedAlias(t *testing.T) {
}
}
+func TestSelectSiteSearchCandidatesMatchesSubscriptionOriginalNameAlias(t *testing.T) {
+ sub := &model.Subscription{
+ Name: "玩具总动员 5 自动订阅",
+ Filter: "玩具总动员 5 2026",
+ OriginalName: "Toy Story 5",
+ Year: 2026,
+ MediaType: "movie",
+ }
+ results := []SearchResult{
+ {Title: "Toy Story 5 2026 1080p WEB-DL", DownloadURL: "https://pt/download/right", Seeders: 90},
+ }
+
+ got := selectSiteSearchCandidates(results, sub, map[string]struct{}{})
+ if len(got) != 1 || got[0].Download != "https://pt/download/right" {
+ t.Fatalf("selected %#v, want original-name alias match", got)
+ }
+}
+
func TestSelectSiteSearchCandidatesDoesNotWashByDefault(t *testing.T) {
sub := &model.Subscription{Name: "Inception 自动订阅", Filter: "Inception 2010", MediaType: "movie", WashPriority: "resolution"}
results := []SearchResult{
@@ -174,6 +187,27 @@ func TestSiteSearchKeywordsIncludeAliasesAndCleanedKeywords(t *testing.T) {
}
}
+func TestSiteSearchKeywordsUseCleanMetadataAliases(t *testing.T) {
+ sub := &model.Subscription{
+ Name: "玩具总动员 4 自动订阅",
+ Filter: "玩具总动员 4 2019",
+ OriginalName: "Toy Story 4",
+ Year: 2019,
+ }
+
+ got := siteSearchKeywords(sub)
+ for _, want := range []string{"玩具总动员 4 2019", "Toy Story 4", "Toy Story 4 2019", "玩具总动员 4"} {
+ if !containsString(got, want) {
+ t.Fatalf("keywords = %#v, missing %q", got, want)
+ }
+ }
+ for _, unwanted := range []string{"玩具总动员 4 自动订阅", "玩具总动员 4 自动订阅 2019", "玩具总动员 4 2019 2019"} {
+ if containsString(got, unwanted) {
+ t.Fatalf("keywords = %#v, should not contain %q", got, unwanted)
+ }
+ }
+}
+
func containsString(values []string, want string) bool {
for _, value := range values {
if value == want {
@@ -225,6 +259,9 @@ func TestSelectSiteSearchCandidatesWithStatsExplainsFiltering(t *testing.T) {
stats.Selected != 1 {
t.Fatalf("unexpected stats: %#v", stats)
}
+ if len(stats.QueryMismatchExamples) != 1 || stats.QueryMismatchExamples[0] != "Different Show S01E01 1080p" {
+ t.Fatalf("query mismatch examples = %#v", stats.QueryMismatchExamples)
+ }
}
func TestDeleteSubscriptionRemovesDownloaderTaskAndSeenState(t *testing.T) {
@@ -249,13 +286,7 @@ func TestDeleteSubscriptionRemovesDownloaderTaskAndSeenState(t *testing.T) {
}))
defer qb.Close()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.Subscription{}, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{})
repos := repository.New(db)
configureTestDefaultQB(t, repos, qb.URL)
downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
@@ -563,13 +594,7 @@ func TestSubscriptionPendingDownloadAvailabilitySkipsUnorganizedEpisodes(t *test
}
func TestSubscriptionPendingDownloadAvailabilityIncludesQueuedTasks(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{})
repos := repository.New(db)
if err := repos.Download.Create(t.Context(), &model.DownloadTask{
Source: "qbittorrent",
@@ -611,13 +636,7 @@ func TestSubscriptionPendingDownloadAvailabilityIncludesQueuedTasks(t *testing.T
}
func TestSubscriptionPendingDownloadAvailabilityIncludesLinkedAliasTask(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.DownloadTask{})
repos := repository.New(db)
sub := &model.Subscription{
Base: model.Base{ID: "sub-qiao-chu"},
@@ -650,894 +669,3 @@ func TestSubscriptionPendingDownloadAvailabilityIncludesLinkedAliasTask(t *testi
t.Fatalf("selected %#v, want linked alias task to satisfy E21", got)
}
}
-
-func TestSubscriptionEnrichProgressIncludesPendingDownloads(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}, &model.Media{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- if err := repos.Download.Create(t.Context(), &model.DownloadTask{
- Source: "qbittorrent",
- URL: "magnet:?xt=urn:btih:4444444444444444444444444444444444444444",
- Title: "Inception 2010 1080p",
- SavePath: "/downloads/movies",
- Status: "completed",
- Progress: 1,
- }); err != nil {
- t.Fatal(err)
- }
- svc := NewSubscriptionService(nil, nil, repos, nil, nil, nil)
- items := []model.Subscription{{
- Name: "Inception 2010",
- Filter: "Inception 2010",
- MediaType: "movie",
- SavePath: "/downloads/movies",
- }}
-
- svc.EnrichProgress(t.Context(), items)
- if items[0].InLibrary {
- t.Fatal("pending download should not be reported as in-library media")
- }
- if items[0].DownloadedEpisodes != 1 || items[0].LocalMediaCount != 1 || items[0].TotalEpisodes != 1 {
- t.Fatalf("unexpected enriched progress: %+v", items[0])
- }
-}
-
-func TestSubscriptionLocalAvailabilityMatchesMediaPath(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Media{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- if err := db.Create(&model.Media{
- Title: "Scraped English Title",
- Path: "/media/电视剧/国产剧/凡人修仙传/Season 01/凡人修仙传 - S01E146.mkv",
- SeasonNum: 1,
- EpisodeNum: 146,
- }).Error; err != nil {
- t.Fatal(err)
- }
- sub := &model.Subscription{
- Name: "凡人修仙传 年番",
- Filter: "凡人修仙传",
- MediaType: "tv",
- TotalEpisodes: 146,
- }
-
- availability := SubscriptionLocalAvailability(t.Context(), repos, sub)
- if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 146)]; !ok {
- t.Fatalf("missing path-matched E146 key: %#v", availability.ExistingEpisodeKeys)
- }
- results := []SearchResult{
- {Title: "凡人修仙传 年番 - 146 1080p", DownloadURL: "https://pt/download/146", Seeders: 80},
- }
- got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
- if len(got) != 0 {
- t.Fatalf("selected %#v, want none because path-matched local episode exists", got)
- }
-}
-
-func TestSubscriptionPendingDownloadAvailabilityIgnoresDeletedTasks(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- if err := repos.Download.Create(t.Context(), &model.DownloadTask{
- Source: "qbittorrent",
- URL: "magnet:?xt=urn:btih:3333333333333333333333333333333333333333",
- Title: "间谍过家家 S01E02 1080p",
- SavePath: "/downloads/tv",
- Status: "deleted",
- }); err != nil {
- t.Fatal(err)
- }
- svc := NewSubscriptionService(nil, nil, repos, nil, nil, nil)
- sub := &model.Subscription{
- Name: "间谍过家家 自动订阅",
- Filter: "间谍过家家",
- MediaType: "tv",
- SavePath: "/downloads/tv",
- TotalEpisodes: 3,
- }
-
- availability := svc.pendingDownloadAvailability(t.Context(), sub)
- if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 2)]; ok {
- t.Fatalf("deleted E02 task should not count as available: %#v", availability.ExistingEpisodeKeys)
- }
- results := []SearchResult{
- {Title: "间谍过家家 S01E02 1080p WEB-DL", DownloadURL: "https://pt/download/2", Seeders: 80},
- {Title: "间谍过家家 S01E03 1080p WEB-DL", DownloadURL: "https://pt/download/3", Seeders: 70},
- }
- got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
- if len(got) != 2 || got[0].Episode != 2 || got[1].Episode != 3 {
- t.Fatalf("selected %#v, want deleted episode 2 and new episode 3", got)
- }
-}
-
-func TestSubscriptionPendingDownloadAvailabilityIncludesLiveQBTorrents(t *testing.T) {
- 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":"间谍过家家 S01E01 1080p","state":"downloading","progress":0.2}]`))
- default:
- http.NotFound(w, r)
- }
- }))
- defer qb.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.DownloadTask{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- downloads.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"})
- svc := NewSubscriptionService(nil, nil, repos, downloads, nil, nil)
- sub := &model.Subscription{
- Name: "间谍过家家 自动订阅",
- Filter: "间谍过家家",
- MediaType: "tv",
- SavePath: "/downloads/tv",
- TotalEpisodes: 2,
- }
-
- availability := svc.pendingDownloadAvailability(t.Context(), sub)
- if availability.DownloadedEpisodes != 1 {
- t.Fatalf("downloaded episodes = %d, want 1", availability.DownloadedEpisodes)
- }
- if _, ok := availability.ExistingEpisodeKeys[episodeKey(1, 1)]; !ok {
- t.Fatalf("missing live qB E01 key: %#v", availability.ExistingEpisodeKeys)
- }
-}
-
-func TestSubscriptionRunOneArchivesCompletedMovieRSS(t *testing.T) {
- rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
- w.Header().Set("Content-Type", "application/rss+xml")
- _, _ = w.Write([]byte(`
-
- -
- Dune 2021 1080p WEB-DL
- dune-1080-web
- magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Dune+2021+1080p+WEB-DL
-
-`))
- }))
- defer rss.Close()
-
- var addCalls int32
- var added bool
- 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 added {
- _, _ = w.Write([]byte(`[{"hash":"dunehash","name":"Dune 2021 1080p WEB-DL","state":"downloading","progress":0.1}]`))
- return
- }
- _, _ = w.Write([]byte(`[]`))
- case "/api/v2/torrents/add":
- added = true
- atomic.AddInt32(&addCalls, 1)
- _, _ = w.Write([]byte("Ok."))
- default:
- http.NotFound(w, r)
- }
- }))
- defer qb.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- configureTestDefaultQB(t, repos, qb.URL)
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
-
- sub := &model.Subscription{
- Name: "Dune 自动订阅",
- FeedURL: rss.URL,
- Filter: "Dune 2021",
- MediaType: "movie",
- SavePath: "/downloads/movies",
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
- queued, err := svc.runOne(t.Context(), sub)
- if err != nil {
- t.Fatal(err)
- }
- if queued != 1 {
- t.Fatalf("queued = %d, want 1", queued)
- }
- if got := atomic.LoadInt32(&addCalls); got != 1 {
- t.Fatalf("qb add calls = %d, want 1", got)
- }
- active, err := repos.Subscription.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(active) != 0 {
- t.Fatalf("active subscriptions = %d, want 0 after completion", len(active))
- }
- history, err := repos.Subscription.History(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(history) != 1 || history[0].ArchivedAt == nil {
- t.Fatalf("history = %#v, want one archived subscription", history)
- }
-}
-
-func TestSubscriptionArchiveCompletedSingleEpisodeTV(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
- sub := &model.Subscription{
- Name: "Some Show S01E01 自动订阅",
- FeedURL: "site-search://search?keyword=Some%20Show%20S01E01",
- Filter: "Some Show S01E01",
- MediaType: "tv",
- Enabled: true,
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
-
- err = svc.archiveCompletedSubscription(t.Context(), sub, LocalAvailability{
- DownloadedEpisodes: 1,
- LocalMediaCount: 1,
- InLibrary: true,
- ExistingEpisodeKeys: map[string]struct{}{
- episodeKey(1, 1): {},
- },
- })
- if err != nil {
- t.Fatal(err)
- }
- active, err := repos.Subscription.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(active) != 0 {
- t.Fatalf("active subscriptions = %d, want 0", len(active))
- }
- history, err := repos.Subscription.History(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(history) != 1 || history[0].ArchiveReason == "" {
- t.Fatalf("history = %#v, want archived single episode", history)
- }
-}
-
-func TestSubscriptionArchiveKeepsGenericUnknownTotalSeriesActive(t *testing.T) {
- sub := &model.Subscription{
- Name: "Some Show 自动订阅",
- Filter: "Some Show",
- MediaType: "tv",
- }
- availability := LocalAvailability{
- DownloadedEpisodes: 1,
- LocalMediaCount: 1,
- InLibrary: true,
- ExistingEpisodeKeys: map[string]struct{}{
- episodeKey(1, 1): {},
- },
- }
- if subscriptionShouldArchive(sub, availability) {
- t.Fatal("generic series with unknown total should stay active for incremental episodes")
- }
-}
-
-func TestInferSubscriptionTotalEpisodesFromSearchAndRSS(t *testing.T) {
- sub := &model.Subscription{Name: "Some Show 自动订阅", Filter: "Some Show", MediaType: "tv"}
- results := []SearchResult{
- {Title: "Some Show S01E01 1080p"},
- {Title: "Some Show S01E12 1080p"},
- {Title: "Other Show S01E99 1080p"},
- }
- if got := inferSearchTotalEpisodes(results, sub); got != 12 {
- t.Fatalf("search inferred total = %d, want 12", got)
- }
- subtitleResults := []SearchResult{
- {Title: "Smoking Behind the Supermarket with You", Subtitle: "躲在超市后门抽烟的两人 S01E12"},
- }
- subtitleSub := &model.Subscription{Name: "躲在超市后门抽烟的两人 自动订阅", Filter: "躲在超市后门抽烟的两人", MediaType: "tv"}
- if got := inferSearchTotalEpisodes(subtitleResults, subtitleSub); got != 12 {
- t.Fatalf("subtitle search inferred total = %d, want 12", got)
- }
- items := []rssItem{
- {Title: "Some Show S01E02 WEB-DL"},
- {Title: "Some Show S01E10 WEB-DL"},
- }
- if got := inferRSSTotalEpisodes(items, sub, compileFilter("Some Show")); got != 10 {
- t.Fatalf("rss inferred total = %d, want 10", got)
- }
-}
-
-func TestResolveSubscriptionTotalEpisodesPrefersTMDbOverTitleFallback(t *testing.T) {
- upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- switch r.URL.Path {
- case "/search/tv":
- _, _ = w.Write([]byte(`{"results":[{"id":42,"name":"Some Show","first_air_date":"2026-01-01"}]}`))
- case "/tv/42":
- _, _ = w.Write([]byte(`{"number_of_episodes":13}`))
- default:
- http.NotFound(w, r)
- }
- }))
- defer upstream.Close()
-
- cfg := &config.Config{}
- cfg.Secrets.TMDbAPIKey = "test"
- cfg.Secrets.TMDbAPIProxy = upstream.URL
- tmdb := NewTMDbProvider(cfg, zap.NewNop(), nil)
- svc := NewSubscriptionService(cfg, zap.NewNop(), nil, nil, nil, NewHub(zap.NewNop()))
- svc.SetScraper(NewScraperService(cfg, zap.NewNop(), nil, tmdb, nil, nil, nil, NewHub(zap.NewNop())))
-
- sub := &model.Subscription{Name: "Some Show 自动订阅", Filter: "Some Show", MediaType: "tv"}
- if got := svc.resolveSubscriptionTotalEpisodes(t.Context(), sub, 10); got != 13 {
- t.Fatalf("resolved total = %d, want TMDb total 13", got)
- }
-}
-
-func TestSubscriptionArchiveKeepsWashSubscriptionActive(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
- sub := &model.Subscription{
- Name: "Dune 自动订阅",
- FeedURL: "site-search://search?keyword=Dune",
- Filter: "Dune 2021",
- MediaType: "movie",
- WashEnabled: true,
- Enabled: true,
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
-
- err = svc.archiveCompletedSubscription(t.Context(), sub, LocalAvailability{
- DownloadedEpisodes: 1,
- LocalMediaCount: 1,
- InLibrary: true,
- })
- if err != nil {
- t.Fatal(err)
- }
- active, err := repos.Subscription.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(active) != 1 {
- t.Fatalf("active subscriptions = %d, want wash subscription to stay active", len(active))
- }
- history, err := repos.Subscription.History(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(history) != 0 {
- t.Fatalf("history subscriptions = %d, want 0", len(history))
- }
-}
-
-func TestRestoreArchivedSubscriptionReturnsToActiveAndClearsSeenState(t *testing.T) {
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop()))
- sub := &model.Subscription{
- Name: "南部档案 自动订阅",
- FeedURL: "https://rss.example/feed",
- Filter: "南部档案",
- MediaType: "tv",
- TotalEpisodes: 33,
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
- archivedAt := time.Now()
- if err := repos.Subscription.Archive(t.Context(), sub.ID, "已下载 1/33 集,缺 33 集", archivedAt); err != nil {
- t.Fatal(err)
- }
- if err := repos.Setting.Set(t.Context(), "subscription."+sub.ID+".seen", "old-guid"); err != nil {
- t.Fatal(err)
- }
- restored, err := svc.Restore(t.Context(), sub.ID)
- if err != nil {
- t.Fatal(err)
- }
- if restored.ArchivedAt != nil || restored.ArchiveReason != "" || !restored.Enabled {
- t.Fatalf("restored subscription not active: archived=%v reason=%q enabled=%v", restored.ArchivedAt, restored.ArchiveReason, restored.Enabled)
- }
- if restored.TotalEpisodes != 0 {
- t.Fatalf("restored total_episodes = %d, want 0 so it gets recomputed from authoritative metadata", restored.TotalEpisodes)
- }
- active, err := repos.Subscription.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(active) != 1 || active[0].ID != sub.ID {
- t.Fatalf("active subscriptions = %#v, want restored subscription", active)
- }
- history, err := repos.Subscription.History(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(history) != 0 {
- t.Fatalf("history subscriptions = %d, want 0 after restore", len(history))
- }
- seen, err := repos.Setting.Get(t.Context(), "subscription."+sub.ID+".seen")
- if err != nil {
- t.Fatal(err)
- }
- if seen != "" {
- t.Fatalf("seen state = %q, want cleared", seen)
- }
-}
-
-func TestSubscriptionRunOneDeduplicatesDuplicateRSSGUIDInSameFeed(t *testing.T) {
- rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
- w.Header().Set("Content-Type", "application/rss+xml")
- _, _ = w.Write([]byte(`
-
- -
- Some Show S01E01 1080p
- episode-1
- magnet:?xt=urn:btih:1111111111111111111111111111111111111111&dn=Some+Show+S01E01
-
- -
- Some Show S01E01 1080p
- episode-1
- magnet:?xt=urn:btih:1111111111111111111111111111111111111111&dn=Some+Show+S01E01
-
-`))
- }))
- defer rss.Close()
-
- 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":"Some Show S01E01 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- configureTestDefaultQB(t, repos, qb.URL)
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
-
- sub := &model.Subscription{
- Name: "Some Show 自动订阅",
- FeedURL: rss.URL,
- Filter: "Some Show",
- MediaType: "tv",
- SavePath: "/downloads/tv",
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
- queued, err := svc.runOne(t.Context(), sub)
- if err != nil {
- t.Fatal(err)
- }
- if queued != 1 {
- t.Fatalf("queued = %d, want 1", queued)
- }
- if got := atomic.LoadInt32(&addCalls); got != 1 {
- t.Fatalf("qb add calls = %d, want 1", got)
- }
- rows, err := repos.Download.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(rows) != 1 {
- t.Fatalf("download rows = %d, want 1", len(rows))
- }
-}
-
-func TestSubscriptionRunOneSkipsSameEpisodeAddedEarlierInFeed(t *testing.T) {
- rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
- w.Header().Set("Content-Type", "application/rss+xml")
- _, _ = w.Write([]byte(`
-
- -
- Some Show S01E01 1080p
- episode-1-a
- magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Some+Show+S01E01+1080p
-
- -
- Some Show S01E01 WEB-DL
- episode-1-b
- magnet:?xt=urn:btih:bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb&dn=Some+Show+S01E01+WEB-DL
-
-`))
- }))
- defer rss.Close()
-
- 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":"Some Show S01E01 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, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- configureTestDefaultQB(t, repos, qb.URL)
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
-
- sub := &model.Subscription{
- Name: "Some Show 自动订阅",
- FeedURL: rss.URL,
- Filter: "Some Show",
- MediaType: "tv",
- SavePath: "/downloads/tv",
- TotalEpisodes: 12,
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
- queued, err := svc.runOne(t.Context(), sub)
- if err != nil {
- t.Fatal(err)
- }
- if queued != 1 {
- t.Fatalf("queued = %d, want 1", queued)
- }
- if got := atomic.LoadInt32(&addCalls); got != 1 {
- t.Fatalf("qb add calls = %d, want 1", got)
- }
- rows, err := repos.Download.List(t.Context())
- if err != nil {
- t.Fatal(err)
- }
- if len(rows) != 1 {
- t.Fatalf("download rows = %d, want 1", len(rows))
- }
-}
-
-func TestSubscriptionRunOneRSSWashQueuesOnlyBestMovieVariant(t *testing.T) {
- rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
- w.Header().Set("Content-Type", "application/rss+xml")
- _, _ = w.Write([]byte(`
-
- -
- Dune 2021 1080p WEB-DL
- dune-1080-web
- magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Dune+2021+1080p+WEB-DL
-
- -
- Dune 2021 2160p UHD BluRay REMUX HDR
- dune-2160-remux
- magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=Dune+2021+2160p+REMUX
-
- -
- Dune 2021 720p HDTV
- dune-720-hdtv
- magnet:?xt=urn:btih:ffffffffffffffffffffffffffffffffffffffff&dn=Dune+2021+720p+HDTV
-
-`))
- }))
- defer rss.Close()
-
- var addCalls int32
- var addedTitles []string
- addedHashes := make([]string, 0, 3)
- 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 len(addedHashes) == 0 {
- _, _ = w.Write([]byte(`[]`))
- return
- }
- var items []string
- for _, hash := range addedHashes {
- items = append(items, `{"hash":"`+hash+`","name":"Dune 2021","state":"downloading","progress":0.1}`)
- }
- _, _ = w.Write([]byte(`[` + strings.Join(items, ",") + `]`))
- case "/api/v2/torrents/add":
- call := atomic.AddInt32(&addCalls, 1)
- _ = r.ParseMultipartForm(10 << 20)
- addedTitles = append(addedTitles, r.FormValue("urls"))
- addedHashes = append(addedHashes, strings.Repeat(fmt.Sprintf("%x", call), 40))
- _, _ = w.Write([]byte("Ok."))
- default:
- http.NotFound(w, r)
- }
- }))
- defer qb.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- configureTestDefaultQB(t, repos, qb.URL)
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
-
- sub := &model.Subscription{
- Name: "Dune 自动订阅",
- FeedURL: rss.URL,
- Filter: "Dune 2021",
- MediaType: "movie",
- WashEnabled: true,
- WashPriority: "resolution",
- SavePath: "/downloads/movies",
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
- queued, err := svc.runOne(t.Context(), sub)
- if err != nil {
- t.Fatal(err)
- }
- if queued != 1 {
- t.Fatalf("queued = %d, want 1 best movie variant", queued)
- }
- if got := atomic.LoadInt32(&addCalls); got != 1 {
- t.Fatalf("qb add calls = %d, want 1", got)
- }
- if len(addedTitles) != 1 || !strings.Contains(addedTitles[0], "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee") {
- t.Fatalf("added %#v, want 2160p REMUX variant only", addedTitles)
- }
-}
-
-func TestSubscriptionRunOneDoesNotUseDeletedDownloader(t *testing.T) {
- rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
- w.Header().Set("Content-Type", "application/rss+xml")
- _, _ = w.Write([]byte(`
-
- -
- Deleted Downloader Show S01E01 1080p
- deleted-downloader-episode-1
- magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Deleted+Downloader+Show+S01E01
-
-`))
- }))
- defer rss.Close()
-
- var qbCalls int32
- qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- atomic.AddInt32(&qbCalls, 1)
- 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":
- _, _ = w.Write([]byte("Ok."))
- default:
- http.NotFound(w, r)
- }
- }))
- defer qb.Close()
-
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil {
- t.Fatal(err)
- }
- repos := repository.New(db)
- client := &model.DownloadClient{Name: "qB deleted", 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.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil {
- t.Fatal(err)
- }
- if err := repos.DownloadClient.Delete(t.Context(), client.ID); err != nil {
- t.Fatal(err)
- }
-
- downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
- svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
- sub := &model.Subscription{
- Name: "Deleted Downloader Show 自动订阅",
- FeedURL: rss.URL,
- Filter: "Deleted Downloader Show",
- MediaType: "tv",
- SavePath: "/downloads/tv",
- }
- if err := repos.Subscription.Create(t.Context(), sub); err != nil {
- t.Fatal(err)
- }
-
- queued, err := svc.runOne(t.Context(), sub)
- if err != nil {
- t.Fatal(err)
- }
- if queued != 0 {
- t.Fatalf("queued = %d, want 0 when default downloader was deleted", queued)
- }
- if got := atomic.LoadInt32(&qbCalls); got != 0 {
- t.Fatalf("qB calls = %d, want 0 after downloader deletion", got)
- }
- 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 TestMatchesSubscriptionRulesUserExcludeWords(t *testing.T) {
- sub := &model.Subscription{ExcludeWords: "10bit,dolby vision,杜比"}
- cases := []struct {
- title string
- want bool
- }{
- {"Movie 2024 1080p WEB-DL", true},
- {"Movie 2024 2160p 10bit HEVC", false},
- {"Movie 2024 2160p Dolby Vision", false},
- {"电影 2024 杜比全景声", false},
- }
- for _, c := range cases {
- if got := matchesSubscriptionRules(sub, c.title); got != c.want {
- t.Errorf("matchesSubscriptionRules(%q) = %v, want %v", c.title, got, c.want)
- }
- }
-}
-
-func TestMatchesSubscriptionRulesDefaultExcludesJunkReleases(t *testing.T) {
- sub := &model.Subscription{}
- for _, title := range []string{
- "Some Movie 2024 CAM",
- "Some Movie 2024 HDTS",
- "某电影 2024 枪版",
- "Some Movie 2024 TELESYNC",
- "Some Show 预告",
- } {
- if matchesSubscriptionRules(sub, title) {
- t.Errorf("expected default rules to exclude junk release %q", title)
- }
- }
-}
-
-func TestMatchesSubscriptionRulesWordBoundaryAvoidsFalsePositives(t *testing.T) {
- sub := &model.Subscription{}
- // "ts" / "cam" / "tc" 作为子串出现在合法标题里时不应被默认排除误伤。
- for _, title := range []string{
- "Tsukihime 2024 1080p WEB-DL",
- "Camp Rock 2024 1080p BluRay",
- "Catch Me 2024 1080p WEB-DL",
- } {
- if !matchesSubscriptionRules(sub, title) {
- t.Errorf("word-boundary match wrongly excluded %q", title)
- }
- }
-}
-
-func TestSelectSiteSearchCandidatesSkipsExistingMovieWhenNotWashing(t *testing.T) {
- sub := &model.Subscription{Name: "Inception 自动订阅", Filter: "Inception 2010", MediaType: "movie"}
- results := []SearchResult{
- {Title: "Inception 2010 2160p 10bit Dolby Vision Atmos", DownloadURL: "https://pt/download/dovi", Seeders: 500},
- {Title: "Inception 2010 1080p WEB-DL", DownloadURL: "https://pt/download/web", Seeders: 90},
- }
- availability := LocalAvailability{LocalMediaCount: 1, InLibrary: true, DownloadedEpisodes: 1, TotalEpisodes: 1}
-
- got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
- if len(got) != 0 {
- t.Fatalf("selected %#v, want none (movie already in library, wash disabled)", got)
- }
-}
-
-func TestSelectSiteSearchCandidatesAllowsMovieWashUpgrade(t *testing.T) {
- sub := &model.Subscription{Name: "Inception 自动订阅", Filter: "Inception 2010", MediaType: "movie", WashEnabled: true, WashPriority: "resolution"}
- results := []SearchResult{
- {Title: "Inception 2010 2160p REMUX", DownloadURL: "https://pt/download/2160", Seeders: 80},
- {Title: "Inception 2010 1080p WEB-DL", DownloadURL: "https://pt/download/1080", Seeders: 200},
- }
- availability := LocalAvailability{LocalMediaCount: 1, InLibrary: true, DownloadedEpisodes: 1, TotalEpisodes: 1}
-
- got := selectSiteSearchCandidates(results, sub, map[string]struct{}{}, availability)
- if len(got) != 1 || got[0].Download != "https://pt/download/2160" {
- t.Fatalf("selected %#v, want 2160p upgrade allowed when washing", got)
- }
-}
-
-func TestSubscriptionItemAlreadyAvailable(t *testing.T) {
- movieSub := &model.Subscription{MediaType: "movie"}
- if !subscriptionItemAlreadyAvailable(movieSub, LocalAvailability{LocalMediaCount: 1}, "Inception 2010 2160p") {
- t.Fatal("movie already in library should be reported available")
- }
- if subscriptionItemAlreadyAvailable(movieSub, LocalAvailability{}, "Inception 2010 2160p") {
- t.Fatal("empty library should not be reported available")
- }
- tvSub := &model.Subscription{MediaType: "tv"}
- avail := LocalAvailability{LocalMediaCount: 1, ExistingEpisodeKeys: map[string]struct{}{episodeKey(1, 2): {}}}
- if !subscriptionItemAlreadyAvailable(tvSub, avail, "Show S01E02 1080p") {
- t.Fatal("existing episode should be reported available")
- }
- if subscriptionItemAlreadyAvailable(tvSub, avail, "Show S01E03 1080p") {
- t.Fatal("missing episode should not be reported available")
- }
-}
diff --git a/internal/service/subtitle.go b/internal/service/subtitle.go
index b4c2dc0..5974f3d 100644
--- a/internal/service/subtitle.go
+++ b/internal/service/subtitle.go
@@ -19,6 +19,7 @@ import (
"errors"
"fmt"
"io"
+ "net/url"
"os"
"path/filepath"
"regexp"
@@ -26,18 +27,31 @@ import (
"go.uber.org/zap"
+ "github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
+ "github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
// SubtitleService is the discovery + conversion entry point.
type SubtitleService struct {
- log *zap.Logger
- repo *repository.Container
+ log *zap.Logger
+ repo *repository.Container
+ storage *StorageConfigService
}
// NewSubtitleService is the constructor.
-func NewSubtitleService(log *zap.Logger, repo *repository.Container) *SubtitleService {
- return &SubtitleService{log: log, repo: repo}
+func NewSubtitleService(log *zap.Logger, repo *repository.Container, storage ...*StorageConfigService) *SubtitleService {
+ s := &SubtitleService{log: log, repo: repo}
+ if len(storage) > 0 {
+ s.storage = storage[0]
+ }
+ return s
+}
+
+func (s *SubtitleService) SetStorageConfig(storage *StorageConfigService) {
+ if s != nil {
+ s.storage = storage
+ }
}
// SubtitleTrack describes one external subtitle file.
@@ -68,6 +82,9 @@ func (s *SubtitleService) Discover(ctx context.Context, mediaID string) ([]Subti
if m == nil {
return nil, errors.New("media not found")
}
+ if strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://") {
+ return s.discoverCloud(ctx, *m), nil
+ }
dir := filepath.Dir(m.Path)
base := strings.TrimSuffix(filepath.Base(m.Path), filepath.Ext(m.Path))
@@ -111,6 +128,126 @@ func (s *SubtitleService) Discover(ctx context.Context, mediaID string) ([]Subti
return tracks, nil
}
+func (s *SubtitleService) discoverCloud(ctx context.Context, m model.Media) []SubtitleTrack {
+ if s == nil || s.storage == nil {
+ return []SubtitleTrack{}
+ }
+ typ, mediaRef, ok := cloudSubtitleMediaRef(m)
+ if !ok {
+ return []SubtitleTrack{}
+ }
+ dirRef, mediaName := splitCloudRef(mediaRef)
+ if mediaName == "" {
+ return []SubtitleTrack{}
+ }
+ base := strings.TrimSuffix(mediaName, filepath.Ext(mediaName))
+ entries, err := s.storage.CloudList(ctx, typ, dirRef)
+ if err != nil {
+ if s.log != nil {
+ s.log.Debug("list cloud subtitles failed", zap.String("provider", typ), zap.String("dir", dirRef), zap.Error(err))
+ }
+ return []SubtitleTrack{}
+ }
+ tracks := cloudSubtitleTracks(typ, entries, base, false)
+ for _, entry := range entries {
+ if !entry.IsDir || !isSubtitleDirectory(entry.Name) || strings.TrimSpace(entry.ID) == "" {
+ continue
+ }
+ subEntries, err := s.storage.CloudList(ctx, typ, entry.ID)
+ if err != nil {
+ continue
+ }
+ tracks = append(tracks, cloudSubtitleTracks(typ, subEntries, base, true)...)
+ }
+ return tracks
+}
+
+func cloudSubtitleTracks(typ string, entries []cloud.FileEntry, base string, subdir bool) []SubtitleTrack {
+ tracks := make([]SubtitleTrack, 0)
+ baseLower := strings.ToLower(base)
+ for _, entry := range entries {
+ if entry.IsDir {
+ continue
+ }
+ ext := strings.ToLower(filepath.Ext(entry.Name))
+ codec, ok := extToCodec[ext]
+ if !ok {
+ continue
+ }
+ fullName := strings.TrimSuffix(entry.Name, ext)
+ if !subdir && !strings.HasPrefix(strings.ToLower(fullName), baseLower) {
+ continue
+ }
+ ref := cloudEntryRef(typ, entry.ID, entry.PickCode)
+ if ref == "" {
+ continue
+ }
+ lang := detectLang(fullName, base)
+ tracks = append(tracks, SubtitleTrack{
+ Lang: lang,
+ Label: lang,
+ Path: buildCloudSubtitlePath(typ, ref, entry.Name),
+ Codec: codec,
+ })
+ }
+ return tracks
+}
+
+func cloudSubtitleMediaRef(m model.Media) (typ, ref string, ok bool) {
+ if info, parsed := ParseCloudLibraryMount(m.Path); parsed && strings.TrimSpace(info.DisplayDir) != "" {
+ return info.Provider, info.DisplayDir, true
+ }
+ if typ, ref, parsed := parseCloudMediaPlaybackURL(m.STRMURL); parsed {
+ return typ, ref, true
+ }
+ return "", "", false
+}
+
+func splitCloudRef(ref string) (dir, name string) {
+ ref = strings.Trim(strings.ReplaceAll(strings.TrimSpace(ref), "\\", "/"), "/")
+ if ref == "" {
+ return "", ""
+ }
+ idx := strings.LastIndex(ref, "/")
+ if idx < 0 {
+ return "", ref
+ }
+ return ref[:idx], ref[idx+1:]
+}
+
+func isSubtitleDirectory(name string) bool {
+ switch strings.ToLower(strings.TrimSpace(name)) {
+ case "subs", "sub", ".sub", "subtitles", "subtitle":
+ return true
+ default:
+ return false
+ }
+}
+
+func buildCloudSubtitlePath(typ, ref, name string) string {
+ u := url.URL{
+ Scheme: "cloud",
+ Host: strings.TrimSpace(typ),
+ Path: "/" + strings.TrimLeft(strings.TrimSpace(ref), "/"),
+ }
+ q := u.Query()
+ q.Set("name", strings.TrimSpace(name))
+ u.RawQuery = q.Encode()
+ return u.String()
+}
+
+func parseCloudSubtitlePath(raw string) (typ, ref, name string, ok bool) {
+ u, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil || strings.ToLower(u.Scheme) != "cloud" || strings.TrimSpace(u.Host) == "" {
+ return "", "", "", false
+ }
+ ref = strings.TrimLeft(u.EscapedPath(), "/")
+ if decoded, err := url.PathUnescape(ref); err == nil {
+ ref = decoded
+ }
+ return strings.TrimSpace(u.Host), strings.TrimSpace(ref), strings.TrimSpace(u.Query().Get("name")), ref != ""
+}
+
// langTag matches the .zh / .zh-cn / .chs language sub-extensions.
var langTag = regexp.MustCompile(`(?i)\.([a-z]{2,3}(?:[-_][a-z]{2,4})?)$`)
@@ -134,6 +271,9 @@ func (s *SubtitleService) Serve(ctx context.Context, mediaID, sub string, w io.W
if err != nil || m == nil {
return errors.New("media not found")
}
+ if typ, ref, name, ok := parseCloudSubtitlePath(sub); ok {
+ return s.serveCloud(ctx, *m, typ, ref, name, w)
+ }
abs, err := filepath.Abs(sub)
if err != nil {
return err
@@ -166,6 +306,42 @@ func (s *SubtitleService) Serve(ctx context.Context, mediaID, sub string, w io.W
return err
}
+func (s *SubtitleService) serveCloud(ctx context.Context, m model.Media, typ, ref, name string, w io.Writer) error {
+ if s == nil || s.storage == nil {
+ return errors.New("cloud storage service unavailable")
+ }
+ mediaTyp, _, ok := cloudSubtitleMediaRef(m)
+ if !ok || mediaTyp != typ {
+ return ErrCloudPlaybackUnavailable
+ }
+ allowed := false
+ for _, track := range s.discoverCloud(ctx, m) {
+ if track.Path == buildCloudSubtitlePath(typ, ref, name) {
+ allowed = true
+ break
+ }
+ }
+ if !allowed {
+ return fmt.Errorf("path escape")
+ }
+ body, err := s.storage.CloudReadText(ctx, typ, ref, 8<<20)
+ if err != nil {
+ return err
+ }
+ ext := strings.ToLower(filepath.Ext(firstNonEmpty(name, ref)))
+ switch ext {
+ case ".vtt":
+ _, err = io.WriteString(w, body)
+ case ".srt":
+ _, err = io.WriteString(w, srtToVTT(body))
+ case ".ass", ".ssa":
+ _, err = io.WriteString(w, assToVTT(body))
+ default:
+ return errors.New("unsupported subtitle format")
+ }
+ return err
+}
+
// srtToVTT performs the minimal SRT → WebVTT transformation: prepend
// "WEBVTT\n\n" and replace ',' with '.' in the timecode separators.
func srtToVTT(body string) string {
diff --git a/internal/service/task_tracker.go b/internal/service/task_tracker.go
index b448e6e..d3d4cb5 100644
--- a/internal/service/task_tracker.go
+++ b/internal/service/task_tracker.go
@@ -244,10 +244,11 @@ func OrganizeTaskMetrics(res *OrganizeResult) map[string]int64 {
return nil
}
metrics := map[string]int64{
- "organized": int64(res.Organized),
- "replaced": int64(res.Replaced),
- "skipped": int64(res.Skipped),
- "errors": int64(len(res.Errors)),
+ "organized": int64(res.Organized),
+ "replaced": int64(res.Replaced),
+ "reclassified": int64(res.Reclassified),
+ "skipped": int64(res.Skipped),
+ "errors": int64(len(res.Errors)),
}
var scanVisited, scanAdded, scanUpdated, scanRemoved int64
for _, scan := range res.Scans {
@@ -269,9 +270,10 @@ func OrganizeTaskMetrics(res *OrganizeResult) map[string]int64 {
for reason, count := range OrganizeSkipReasonCounts(res) {
metrics["skip_"+organizeMetricKey(reason)] = int64(count)
}
- var scrapeMatched int64
+ var scrapeMatched, scrapeProcessed int64
for _, scrape := range res.Scrapes {
scrapeMatched += int64(scrape.Matched)
+ scrapeProcessed += int64(scrape.Processed)
if scrape.Error != "" {
metrics["scrape_errors"]++
}
@@ -282,6 +284,7 @@ func OrganizeTaskMetrics(res *OrganizeResult) map[string]int64 {
if len(res.Scrapes) > 0 {
metrics["scrapes"] = int64(len(res.Scrapes))
metrics["scrape_matched"] = scrapeMatched
+ metrics["scrape_processed"] = scrapeProcessed
}
return metrics
}
@@ -302,7 +305,7 @@ func OrganizeTaskDetails(res *OrganizeResult, limit int) []string {
}
}
for _, item := range res.Items {
- if item.Action != "error" && item.Action != "skip" {
+ if item.Action != "error" && item.Action != "skip" && item.Action != "reclassify" && item.Action != "cleanup" {
continue
}
line := strings.TrimSpace(item.Source)
diff --git a/internal/service/task_tracker_test.go b/internal/service/task_tracker_test.go
new file mode 100644
index 0000000..b9d3079
--- /dev/null
+++ b/internal/service/task_tracker_test.go
@@ -0,0 +1,29 @@
+package service
+
+import "testing"
+
+func TestOrganizeTaskMetricsIncludesScrapeProcessed(t *testing.T) {
+ metrics := OrganizeTaskMetrics(&OrganizeResult{
+ Scrapes: []OrganizeScrapeSummary{
+ {Name: "A", Processed: 3, Matched: 2},
+ {Name: "B", Processed: 4, Matched: 1, Error: "failed"},
+ {Name: "C", Skipped: true},
+ },
+ })
+
+ if metrics["scrapes"] != 3 {
+ t.Fatalf("scrapes = %d, want 3", metrics["scrapes"])
+ }
+ if metrics["scrape_processed"] != 7 {
+ t.Fatalf("scrape_processed = %d, want 7", metrics["scrape_processed"])
+ }
+ if metrics["scrape_matched"] != 3 {
+ t.Fatalf("scrape_matched = %d, want 3", metrics["scrape_matched"])
+ }
+ if metrics["scrape_errors"] != 1 {
+ t.Fatalf("scrape_errors = %d, want 1", metrics["scrape_errors"])
+ }
+ if metrics["scrape_skipped"] != 1 {
+ t.Fatalf("scrape_skipped = %d, want 1", metrics["scrape_skipped"])
+ }
+}
diff --git a/internal/service/telegram_admin_codes.go b/internal/service/telegram_admin_codes.go
new file mode 100644
index 0000000..9a86057
--- /dev/null
+++ b/internal/service/telegram_admin_codes.go
@@ -0,0 +1,133 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *TelegramBotService) replyCapacity(ctx context.Context) telegramCommandReply {
+ c := s.loadCapacity(ctx)
+ quota := "未开放"
+ if c.OpenRegOn {
+ if c.OpenRegLimit > 0 {
+ quota = fmt.Sprintf("已开放(%d/%d 名额)", c.OpenRegUsed, c.OpenRegLimit)
+ } else {
+ quota = "已开放(不限名额,受授权上限约束)"
+ }
+ }
+ text := fmt.Sprintf("容量 / 状态\n\n授权上限:%d 人(随凭证授权实时变化)\n已用:%d 人\n剩余可注册:%d 人\n开注状态:%s",
+ c.MaxUsers, c.UsedUsers, c.Remaining(), quota)
+ return telegramCommandReply{Text: text, Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}
+}
+
+func (s *TelegramBotService) replyOpenRegMenu(ctx context.Context) telegramCommandReply {
+ c := s.loadCapacity(ctx)
+ state := "未开放"
+ if c.OpenRegOn {
+ state = fmt.Sprintf("已开放(%d/%d)", c.OpenRegUsed, c.OpenRegLimit)
+ }
+ return telegramCommandReply{
+ Text: "开注设置\n当前:" + state + "\n选择要开放的名额:",
+ Buttons: [][]telegramInlineButton{
+ {{Text: "5 个", Data: "adm_openreg_set:5"}, {Text: "10 个", Data: "adm_openreg_set:10"}, {Text: "20 个", Data: "adm_openreg_set:20"}},
+ {{Text: "不限名额", Data: "adm_openreg_set:0"}, {Text: "关闭注册", Data: "adm_openreg_close"}},
+ {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
+ },
+ }
+}
+
+func (s *TelegramBotService) replyGenCodeMenu() telegramCommandReply {
+ return telegramCommandReply{
+ Text: "生成兑换码\n选择类型与时长:",
+ Buttons: [][]telegramInlineButton{
+ {{Text: "注册码·30天", Data: "gc:register:30"}, {Text: "注册码·永久", Data: "gc:register:0"}},
+ {{Text: "续期码·30天", Data: "gc:renew:30"}, {Text: "续期码·90天", Data: "gc:renew:90"}},
+ {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
+ },
+ }
+}
+
+func (s *TelegramBotService) replyGenCode(ctx context.Context, msg *TelegramMessage, data string) telegramCommandReply {
+ parts := strings.Split(data, ":") // gc::
+ if len(parts) != 3 {
+ return telegramCommandReply{Text: "参数错误。"}
+ }
+ kind := parts[1]
+ days, _ := strconv.Atoi(parts[2])
+ createdBy := ""
+ if u := s.boundUser(ctx, msg.From.ID); u != nil {
+ createdBy = u.ID
+ }
+ code, err := s.generateCode(ctx, kind, days, 0, createdBy)
+ if err != nil {
+ return telegramCommandReply{Text: "生成失败:" + err.Error()}
+ }
+ kindLabel := map[string]string{model.RegistrationCodeRegister: "注册码", model.RegistrationCodeRenew: "续期码"}[code.Kind]
+ dur := "永久"
+ if days > 0 {
+ dur = fmt.Sprintf("%d 天", days)
+ }
+ return telegramCommandReply{
+ Text: fmt.Sprintf("已生成%s(%s):\n\n%s\n\n发给用户在 Bot 中兑换即可。", kindLabel, dur, code.Code),
+ Buttons: [][]telegramInlineButton{{{Text: "再生成一个", Data: "adm_gencode"}, {Text: "⬅️ 返回菜单", Data: "menu_main"}}},
+ }
+}
+
+func (s *TelegramBotService) cmdGenCode(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "用法:/gencode register|renew 天数 [有效天数] [可用次数]\n示例:/gencode register 30、/gencode renew 90 7 5"}
+ }
+ kind := strings.ToLower(strings.TrimSpace(args[0]))
+ switch kind {
+ case "reg", "register", "注册码":
+ kind = model.RegistrationCodeRegister
+ case "renew", "续期", "续期码":
+ kind = model.RegistrationCodeRenew
+ default:
+ return telegramCommandReply{Text: "类型无效,只支持 register / renew。"}
+ }
+ days, err := strconv.Atoi(args[1])
+ if err != nil || days < 0 {
+ return telegramCommandReply{Text: "天数必须是非负整数,0 表示永久。"}
+ }
+ validDays := 0
+ if len(args) > 2 {
+ validDays, err = strconv.Atoi(args[2])
+ if err != nil || validDays < 0 {
+ return telegramCommandReply{Text: "有效天数必须是非负整数。"}
+ }
+ }
+ maxUses := 1
+ if len(args) > 3 {
+ maxUses, err = strconv.Atoi(args[3])
+ if err != nil || maxUses <= 0 {
+ return telegramCommandReply{Text: "可用次数必须是正整数。"}
+ }
+ }
+ createdBy := ""
+ if u := s.boundUser(ctx, msg.From.ID); u != nil {
+ createdBy = u.ID
+ }
+ code, err := s.generateCodeWithUses(ctx, kind, days, validDays, maxUses, createdBy)
+ if err != nil {
+ return telegramCommandReply{Text: "生成失败:" + err.Error()}
+ }
+ kindLabel := map[string]string{model.RegistrationCodeRegister: "注册码", model.RegistrationCodeRenew: "续期码"}[code.Kind]
+ dur := "永久"
+ if days > 0 {
+ dur = fmt.Sprintf("%d 天", days)
+ }
+ valid := "长期有效"
+ if validDays > 0 && code.ExpiresAt != nil {
+ valid = "有效至 " + code.ExpiresAt.Format("2006-01-02 15:04")
+ }
+ uses := "单次使用"
+ if code.EffectiveMaxUses() > 1 {
+ uses = fmt.Sprintf("最多 %d 次", code.EffectiveMaxUses())
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已生成%s(%s,%s,%s):\n\n%s", kindLabel, dur, valid, uses, code.Code)}
+}
diff --git a/internal/service/telegram_admin_users.go b/internal/service/telegram_admin_users.go
new file mode 100644
index 0000000..2653af2
--- /dev/null
+++ b/internal/service/telegram_admin_users.go
@@ -0,0 +1,186 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+)
+
+func (s *TelegramBotService) replyUserList(ctx context.Context) telegramCommandReply {
+ users, err := s.repo.User.List(ctx)
+ if err != nil {
+ return telegramCommandReply{Text: "读取用户失败:" + err.Error()}
+ }
+ if len(users) == 0 {
+ return telegramCommandReply{Text: "暂无用户。"}
+ }
+ var rows [][]telegramInlineButton
+ limit := len(users)
+ if limit > 12 {
+ limit = 12
+ }
+ for i := 0; i < limit; i++ {
+ u := users[i]
+ flag := ""
+ if !u.IsActive {
+ flag = "🚫"
+ }
+ if u.Role == "admin" {
+ flag = "👑"
+ }
+ rows = append(rows, []telegramInlineButton{{Text: flag + " " + u.Username, Data: "usr:" + u.ID}})
+ }
+ rows = append(rows, []telegramInlineButton{{Text: "⬅️ 返回菜单", Data: "menu_main"}})
+ return telegramCommandReply{Text: fmt.Sprintf("用户管理(共 %d 人,显示前 %d)\n点击用户进行操作:", len(users), limit), Buttons: rows}
+}
+
+func (s *TelegramBotService) replyUserActions(ctx context.Context, userID string) telegramCommandReply {
+ u, err := s.repo.User.FindByID(ctx, userID)
+ if err != nil || u == nil {
+ return telegramCommandReply{Text: "用户不存在。"}
+ }
+ protected := UserIsProtectedAccount(ctx, s.repo, u)
+ text := fmt.Sprintf("%s\n角色:%s\n状态:%s\n到期:%s\n防共享警告:%d 次",
+ u.Username, u.Role, map[bool]string{true: "正常", false: "已禁用"}[u.IsActive], formatExpiry(u.ExpiredAt), u.ShareWarnings)
+ if protected {
+ return telegramCommandReply{Text: text + "\n\n(受保护账号,不可禁用/删除)", Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回", Data: "adm_users"}}}}
+ }
+ banBtn := telegramInlineButton{Text: "🚫 禁用", Data: "uban:" + u.ID}
+ if !u.IsActive {
+ banBtn = telegramInlineButton{Text: "✅ 解禁", Data: "uunban:" + u.ID}
+ }
+ return telegramCommandReply{
+ Text: text,
+ Buttons: [][]telegramInlineButton{
+ {banBtn, {Text: "⏳ 续期30天", Data: "urenew:" + u.ID + ":30"}},
+ {{Text: "🗑 删除用户", Data: "udel:" + u.ID}},
+ {{Text: "⬅️ 返回", Data: "adm_users"}},
+ },
+ }
+}
+
+func (s *TelegramBotService) replyUserBan(ctx context.Context, userID string, unban bool) telegramCommandReply {
+ if !unban {
+ if reason := s.protectReason(ctx, userID); reason != "" {
+ return telegramCommandReply{Text: reason}
+ }
+ }
+ updates := map[string]any{"is_active": unban}
+ if unban {
+ updates["share_warnings"] = 0
+ updates["last_share_warn_at"] = nil
+ }
+ if err := s.repo.User.UpdateFields(ctx, userID, updates); err != nil {
+ return telegramCommandReply{Text: "操作失败:" + err.Error()}
+ }
+ if unban {
+ _ = s.repo.UserDevice.SetKickedByUser(ctx, userID, false)
+ }
+ return s.replyUserActions(ctx, userID)
+}
+
+func (s *TelegramBotService) replyUserDelete(ctx context.Context, userID string) telegramCommandReply {
+ if reason := s.protectReason(ctx, userID); reason != "" {
+ return telegramCommandReply{Text: reason}
+ }
+ u, _ := s.repo.User.FindByID(ctx, userID)
+ _ = s.repo.UserDevice.DeleteByUser(ctx, userID)
+ if err := s.repo.User.Delete(ctx, userID); err != nil {
+ return telegramCommandReply{Text: "删除失败:" + err.Error()}
+ }
+ name := userID
+ if u != nil {
+ name = u.Username
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已删除用户 %s。", name), Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回", Data: "adm_users"}}}}
+}
+
+func (s *TelegramBotService) replyUserRenew(ctx context.Context, payload string) telegramCommandReply {
+ parts := strings.Split(payload, ":") // :
+ if len(parts) != 2 {
+ return telegramCommandReply{Text: "参数错误。"}
+ }
+ days, _ := strconv.Atoi(parts[1])
+ if err := s.applyRenewal(ctx, parts[0], days); err != nil {
+ return telegramCommandReply{Text: "续期失败:" + err.Error()}
+ }
+ return s.replyUserActions(ctx, parts[0])
+}
+
+func (s *TelegramBotService) cmdUserRenew(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "用法:/renew_user 用户名 天数,天数 0 表示永久。"}
+ }
+ user, _ := s.repo.User.FindByUsername(ctx, args[0])
+ if user == nil {
+ user, _ = s.repo.User.FindByID(ctx, args[0])
+ }
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ days, err := strconv.Atoi(args[1])
+ if err != nil || days < 0 {
+ return telegramCommandReply{Text: "天数必须是非负整数。"}
+ }
+ if err := s.applyRenewal(ctx, user.ID, days); err != nil {
+ return telegramCommandReply{Text: "续期失败:" + err.Error()}
+ }
+ return s.replyUserActions(ctx, user.ID)
+}
+
+func (s *TelegramBotService) cmdUserDelete(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "用法:/delete_user 用户名 confirm\n为避免误删,最后一个参数必须是 confirm。"}
+ }
+ if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") {
+ return telegramCommandReply{Text: "删除用户需要确认:/delete_user 用户名 confirm"}
+ }
+ user, _ := s.repo.User.FindByUsername(ctx, args[0])
+ if user == nil {
+ user, _ = s.repo.User.FindByID(ctx, args[0])
+ }
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ return s.replyUserDelete(ctx, user.ID)
+}
+
+// protectReason returns a non-empty message when a user must not be
+// disabled/deleted (admins, default admin and protected-list users).
+func (s *TelegramBotService) protectReason(ctx context.Context, userID string) string {
+ u, err := s.repo.User.FindByID(ctx, userID)
+ if err != nil || u == nil {
+ return "用户不存在。"
+ }
+ if u.Role == "admin" {
+ return "管理员账号受保护,不可禁用/删除。"
+ }
+ if first, _ := s.repo.User.FirstAdmin(ctx); first != nil && first.ID == u.ID {
+ return "默认管理员账号受保护,不可禁用/删除。"
+ }
+ if _, ok := ProtectedUserIDSet(ctx, s.repo)[u.ID]; ok {
+ return "该账号在 Bot 保护名单中,不可禁用/删除。"
+ }
+ if s.device != nil && s.device.UserRecentlyActive(ctx, u.ID, realtimeSessionTTL) {
+ return "该账号最近仍有实时活跃会话,为避免误删/误禁用,请先确认用户已下线。"
+ }
+ return ""
+}
+
+func (s *TelegramBotService) cmdUserBan(ctx context.Context, args []string, unban bool) telegramCommandReply {
+ if len(args) == 0 {
+ if unban {
+ return telegramCommandReply{Text: "用法:/unban 用户名"}
+ }
+ return telegramCommandReply{Text: "用法:/ban 用户名"}
+ }
+ user, _ := s.repo.User.FindByUsername(ctx, args[0])
+ if user == nil {
+ user, _ = s.repo.User.FindByID(ctx, args[0])
+ }
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ return s.replyUserBan(ctx, user.ID, unban)
+}
diff --git a/internal/service/telegram_api.go b/internal/service/telegram_api.go
index ef4226f..f30306f 100644
--- a/internal/service/telegram_api.go
+++ b/internal/service/telegram_api.go
@@ -55,6 +55,7 @@ func telegramHTTPClient(timeout time.Duration, cfg map[string]string) *http.Clie
func telegramHTTPClients(timeout time.Duration, cfg map[string]string) []*http.Client {
clients := []*http.Client{}
seen := map[string]bool{}
+ customAPIBase := telegramUsesCustomAPIBase(cfg)
for _, proxyRaw := range telegramProxyCandidates(cfg) {
proxyURL, err := normalizeProxyURL(proxyRaw, "http")
if err != nil || proxyURL == nil {
@@ -70,6 +71,9 @@ func telegramHTTPClients(timeout time.Duration, cfg map[string]string) []*http.C
clients = append(clients, &http.Client{Timeout: timeout, Transport: transport})
}
transport := NewExternalTransport()
+ if customAPIBase {
+ transport = NewInternalTransport()
+ }
clients = append(clients, &http.Client{Timeout: timeout, Transport: transport})
return clients
}
@@ -87,6 +91,9 @@ func telegramProxyCandidates(cfg map[string]string) []string {
if len(out) > 0 {
return out
}
+ if telegramUsesCustomAPIBase(cfg) {
+ return out
+ }
for _, value := range []string{
"http://127.0.0.1:10808",
"http://127.0.0.1:10809",
@@ -102,6 +109,10 @@ func telegramProxyCandidates(cfg map[string]string) []string {
return out
}
+func telegramUsesCustomAPIBase(cfg map[string]string) bool {
+ return telegramAPIBaseURL(cfg) != defaultTelegramAPIBaseURL
+}
+
func telegramPostForm(ctx context.Context, cfg map[string]string, method string, form url.Values, timeout time.Duration) error {
apiURL, err := telegramMethodURL(cfg, cfg["bot_token"], method)
if err != nil {
diff --git a/internal/service/telegram_api_test.go b/internal/service/telegram_api_test.go
index 70aa189..6e10856 100644
--- a/internal/service/telegram_api_test.go
+++ b/internal/service/telegram_api_test.go
@@ -249,6 +249,17 @@ func TestTelegramProxyCandidatesDefaultLocalFallbacks(t *testing.T) {
}
}
+func TestTelegramHTTPClientsCustomAPIBaseSkipsDefaultProxyFallback(t *testing.T) {
+ clients := telegramHTTPClients(time.Second, map[string]string{
+ "api_base_url": "http://127.0.0.1:18080",
+ })
+ if len(clients) != 1 {
+ t.Fatalf("clients = %d, want direct client only", len(clients))
+ }
+ if got := telegramClientProxyString(t, clients[0]); got != "" {
+ t.Fatalf("custom api_base_url proxy = %q, want direct", got)
+ }
+}
func TestTelegramHTTPClientsPreferConfiguredProxy(t *testing.T) {
clients := telegramHTTPClients(time.Second, map[string]string{
"proxy_url": "http://proxy.example:7890",
diff --git a/internal/service/telegram_binding.go b/internal/service/telegram_binding.go
new file mode 100644
index 0000000..1130b6a
--- /dev/null
+++ b/internal/service/telegram_binding.go
@@ -0,0 +1,470 @@
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "strconv"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// findChannelByChatID 根据 chat_id 查找已配置的通知渠道。
+func (s *TelegramBotService) findChannelByChatID(ctx context.Context, chatID int) *model.NotifyChannel {
+ channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
+ if err != nil {
+ return nil
+ }
+ target := strconv.Itoa(chatID)
+ for _, ch := range channels {
+ if !ch.Enabled {
+ continue
+ }
+ configStr := ch.Config
+ if s.crypto != nil && configStr != "" {
+ configStr = s.crypto.Decrypt(configStr)
+ }
+ var cfg map[string]string
+ if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
+ continue
+ }
+ if cfg["chat_id"] == target || cfg["command_chat_id"] == target ||
+ cfg["group_chat_id"] == target || cfg["channel_chat_id"] == target {
+ return &ch
+ }
+ }
+ if len(channels) == 1 && channels[0].Enabled {
+ return &channels[0]
+ }
+ return nil
+}
+
+func (s *TelegramBotService) findChannelForMessage(ctx context.Context, msg *TelegramMessage) *model.NotifyChannel {
+ if msg == nil {
+ return nil
+ }
+ if msg.Chat.Type != "" && msg.Chat.Type != "private" {
+ return s.findChannelByChatID(ctx, msg.Chat.ID)
+ }
+ channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
+ if err != nil {
+ return nil
+ }
+ var first *model.NotifyChannel
+ for i := range channels {
+ ch := channels[i]
+ if !ch.Enabled {
+ continue
+ }
+ if first == nil {
+ first = &ch
+ }
+ if s.telegramUserIsAdmin(ctx, &ch, msg.From.ID) || s.telegramUserCanBind(ctx, &ch, msg.From.ID) {
+ return &ch
+ }
+ }
+ return first
+}
+
+func (s *TelegramBotService) channelForMessage(ctx context.Context, msg *TelegramMessage, hint *model.NotifyChannel) *model.NotifyChannel {
+ if hint == nil {
+ return s.findChannelForMessage(ctx, msg)
+ }
+ if msg == nil {
+ return hint
+ }
+ if msg.Chat.Type != "" && msg.Chat.Type != "private" && !s.telegramChatAllowed(hint, msg.Chat.ID) {
+ return nil
+ }
+ return hint
+}
+
+func (s *TelegramBotService) handleCallback(ctx context.Context, cb *TelegramCallbackQuery, channelHint *model.NotifyChannel) error {
+ if cb == nil || cb.Message == nil {
+ return nil
+ }
+ msg := *cb.Message
+ msg.From = cb.From
+ channel := s.channelForMessage(ctx, &msg, channelHint)
+ if channel == nil {
+ channel = s.findChannelByChatID(ctx, cb.Message.Chat.ID)
+ }
+ // 立即应答回调,关闭按钮上的加载状态,避免客户端长时间转圈。
+ if telegramIsGroupChat(cb.Message.Chat.Type) {
+ s.answerCallbackWithText(ctx, channel, cb.ID, "为了隐私,群组内按钮面板已禁用。请私聊 Bot 或在群里发送 /menu,我会把面板私聊给你。", true)
+ s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
+ return nil
+ }
+ if cb.Message.Chat.Type == "private" && cb.Message.Chat.ID != cb.From.ID {
+ s.answerCallbackWithText(ctx, channel, cb.ID, "这个面板不属于你,请发送 /menu 打开自己的面板。", true)
+ return nil
+ }
+ s.answerCallback(ctx, channel, cb.ID)
+ data := strings.TrimSpace(cb.Data)
+ if data == "adult_toggle" {
+ reply := s.cmdHideAdult(ctx, &msg, nil)
+ if reply.Text != "" {
+ err := s.reply(ctx, channel, cb.Message.Chat.ID, reply)
+ s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
+ return err
+ }
+ return nil
+ }
+ if reply, handled := s.handleMenuCallback(ctx, channel, &msg, data); handled {
+ if reply.Text != "" {
+ err := s.reply(ctx, channel, cb.Message.Chat.ID, reply)
+ s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
+ return err
+ }
+ }
+ return nil
+}
+
+// answerCallback 应答 Telegram 回调查询,关闭按钮上的加载提示。
+func (s *TelegramBotService) answerCallback(ctx context.Context, channel *model.NotifyChannel, callbackID string) {
+ s.answerCallbackWithText(ctx, channel, callbackID, "", false)
+}
+
+func (s *TelegramBotService) answerCallbackWithText(ctx context.Context, channel *model.NotifyChannel, callbackID, text string, showAlert bool) {
+ if channel == nil || strings.TrimSpace(callbackID) == "" {
+ return
+ }
+ cfg := s.telegramChannelConfig(channel)
+ if strings.TrimSpace(cfg["bot_token"]) == "" {
+ return
+ }
+ payload := map[string]interface{}{
+ "callback_query_id": callbackID,
+ }
+ if strings.TrimSpace(text) != "" {
+ payload["text"] = text
+ payload["show_alert"] = showAlert
+ }
+ if err := telegramPostJSON(ctx, cfg, "answerCallbackQuery", payload, 8*time.Second); err != nil {
+ s.log.Debug("telegram answerCallbackQuery failed", zap.Error(sanitizeTelegramError(err)))
+ }
+}
+
+func (s *TelegramBotService) telegramBinding(ctx context.Context, telegramUserID int) *model.TelegramBinding {
+ if telegramUserID == 0 {
+ return nil
+ }
+ var binding model.TelegramBinding
+ err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(telegramUserID)).First(&binding).Error
+ if err != nil {
+ return nil
+ }
+ return &binding
+}
+
+func (s *TelegramBotService) unbindTelegramUser(ctx context.Context, telegramUserID int) error {
+ if s == nil || s.repo == nil || s.repo.DB == nil || telegramUserID == 0 {
+ return nil
+ }
+ return s.repo.DB.WithContext(ctx).Unscoped().
+ Where("telegram_user_id = ?", int64(telegramUserID)).
+ Delete(&model.TelegramBinding{}).Error
+}
+
+func (s *TelegramBotService) telegramUserIsAdmin(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool {
+ if s.telegramUserIDConfigured(channel, telegramUserID) {
+ return true
+ }
+ binding := s.telegramBinding(ctx, telegramUserID)
+ if binding == nil {
+ return false
+ }
+ user, err := s.repo.User.FindByID(ctx, binding.UserID)
+ return err == nil && user != nil && user.Role == "admin" && user.IsActive
+}
+
+func (s *TelegramBotService) telegramChatAllowed(channel *model.NotifyChannel, chatID int) bool {
+ if channel == nil {
+ return false
+ }
+ configStr := channel.Config
+ if s.crypto != nil && configStr != "" {
+ configStr = s.crypto.Decrypt(configStr)
+ }
+ var cfg map[string]string
+ if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
+ return false
+ }
+ target := strconv.Itoa(chatID)
+ for _, key := range []string{"group_chat_id", "channel_chat_id", "command_chat_id"} {
+ if configured := strings.TrimSpace(cfg[key]); configured != "" && configured == target {
+ return true
+ }
+ }
+ if strings.TrimSpace(cfg["group_chat_id"]) != "" || strings.TrimSpace(cfg["channel_chat_id"]) != "" || strings.TrimSpace(cfg["command_chat_id"]) != "" {
+ return false
+ }
+ return strings.TrimSpace(cfg["chat_id"]) == target
+}
+
+// telegramBindDecision 表示成员资格校验的三态结果:通过 / 明确不通过 /
+// 无法验证(getChatMember 出错,如 Bot 不在群、群 ID 失效、网络或代理不可达)。
+// 区分「明确不是成员」和「查不了」,是为了避免把验证失败误报成「你不在群」。
+type telegramBindDecision int
+
+const (
+ bindDenied telegramBindDecision = iota // 已查实:不在任何绑定群组/频道
+ bindAllowed // 管理员,或查实是某绑定群组/频道成员
+ bindUnverifiable // 配了群组/频道但 getChatMember 全部失败
+)
+
+// telegramMembership 表示单个 chat 的成员资格三态。
+type telegramMembership int
+
+const (
+ membershipNo telegramMembership = iota // 查实不是成员(left/kicked 等)
+ membershipYes // 查实是成员
+ membershipUnknown // getChatMember 出错,无法判定
+)
+
+func (s *TelegramBotService) telegramUserBindDecision(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) telegramBindDecision {
+ if telegramUserID == 0 || channel == nil {
+ return bindDenied
+ }
+ if s.telegramUserIDConfigured(channel, telegramUserID) {
+ return bindAllowed
+ }
+ chatIDs := s.telegramMembershipChatIDs(channel)
+ if len(chatIDs) == 0 {
+ return bindDenied
+ }
+ sawUnknown := false
+ for _, chatID := range chatIDs {
+ switch s.telegramChatMembership(ctx, channel, chatID, telegramUserID) {
+ case membershipYes:
+ return bindAllowed
+ case membershipUnknown:
+ sawUnknown = true
+ }
+ }
+ if sawUnknown {
+ return bindUnverifiable
+ }
+ return bindDenied
+}
+
+// telegramUserCanBind 是 telegramUserBindDecision 的布尔包装,供尽力而为的场景
+// 使用(如私聊时挑选可用渠道):只有查实通过才返回 true。
+func (s *TelegramBotService) telegramUserCanBind(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool {
+ return s.telegramUserBindDecision(ctx, channel, telegramUserID) == bindAllowed
+}
+
+// telegramBindRejectText 根据三态结果生成面向用户的提示。action 形如「兑换注册账号」
+// 「绑定媒体中心账号」。bindUnverifiable 时不再误导用户「你不在群」,而是提示
+// 管理员检查 Bot 权限与群组 ID。
+func telegramBindRejectText(decision telegramBindDecision, action string) string {
+ if decision == bindUnverifiable {
+ return fmt.Sprintf("暂时无法验证你的群组/频道成员身份,%s未成功。这通常是因为 Bot 未加入绑定群组、在频道中不是管理员,或群组 ID 配置有误(如超级群需带 -100 前缀)。请联系管理员检查 Bot 权限与「绑定群组/频道 ID」。", action)
+ }
+ return fmt.Sprintf("当前 Telegram 账号不在管理员配置的绑定群组/频道中,无法%s。请先加入管理员配置的群组或频道;如果尚未配置,请联系管理员。", action)
+}
+
+func (s *TelegramBotService) telegramChatMembership(ctx context.Context, channel *model.NotifyChannel, chatID string, telegramUserID int) telegramMembership {
+ cfg := s.telegramChannelConfig(channel)
+ if strings.TrimSpace(cfg["bot_token"]) == "" || chatID == "" || telegramUserID == 0 {
+ return membershipUnknown
+ }
+ payload := map[string]interface{}{
+ "chat_id": chatID,
+ "user_id": telegramUserID,
+ }
+ var result struct {
+ OK bool `json:"ok"`
+ Result struct {
+ Status string `json:"status"`
+ } `json:"result"`
+ }
+ if err := telegramPostJSONDecode(ctx, cfg, "getChatMember", payload, 15*time.Second, &result); err != nil {
+ s.log.Warn("telegram getChatMember failed", zap.String("chat_id", chatID), zap.Int("telegram_user_id", telegramUserID), zap.Error(sanitizeTelegramError(err)))
+ return membershipUnknown
+ }
+ if !result.OK {
+ return membershipUnknown
+ }
+ switch strings.ToLower(result.Result.Status) {
+ case "creator", "administrator", "member", "restricted":
+ return membershipYes
+ default:
+ return membershipNo
+ }
+}
+
+// telegramUserIsChatMember 是 telegramChatMembership 的布尔包装,仅在查实是成员时
+// 返回 true(查不了也视为非成员,供尽力而为的场景使用)。
+func (s *TelegramBotService) telegramUserIsChatMember(ctx context.Context, channel *model.NotifyChannel, chatID string, telegramUserID int) bool {
+ return s.telegramChatMembership(ctx, channel, chatID, telegramUserID) == membershipYes
+}
+
+func (s *TelegramBotService) telegramUserIDConfigured(channel *model.NotifyChannel, telegramUserID int) bool {
+ if channel == nil || telegramUserID == 0 {
+ return false
+ }
+ cfg := s.telegramChannelConfig(channel)
+ target := strconv.Itoa(telegramUserID)
+ for _, value := range telegramConfiguredUserIDs(cfg["admin_user_ids"]) {
+ if value == target {
+ return true
+ }
+ }
+ if strings.TrimSpace(cfg["admin_user_ids"]) == "" && strings.TrimSpace(cfg["chat_id"]) == target {
+ return true
+ }
+ return false
+}
+
+func (s *TelegramBotService) telegramChannelConfig(channel *model.NotifyChannel) map[string]string {
+ return telegramConfigFromChannel(s.crypto, channel)
+}
+
+func telegramConfigFromChannel(crypto *CryptoService, channel *model.NotifyChannel) map[string]string {
+ if channel == nil {
+ return map[string]string{}
+ }
+ configStr := channel.Config
+ if crypto != nil && configStr != "" {
+ configStr = crypto.Decrypt(configStr)
+ }
+ var cfg map[string]string
+ if err := json.Unmarshal([]byte(configStr), &cfg); err != nil || cfg == nil {
+ return map[string]string{}
+ }
+ normalizeTelegramConfig(cfg)
+ return cfg
+}
+
+func normalizeTelegramConfig(cfg map[string]string) {
+ if cfg == nil {
+ return
+ }
+ chatID := strings.TrimSpace(cfg["chat_id"])
+ if chatID == "" {
+ return
+ }
+ if strings.HasPrefix(chatID, "-") {
+ if strings.TrimSpace(cfg["group_chat_id"]) == "" && strings.TrimSpace(cfg["channel_chat_id"]) == "" && strings.TrimSpace(cfg["command_chat_id"]) == "" {
+ cfg["group_chat_id"] = chatID
+ }
+ return
+ }
+ if strings.TrimSpace(cfg["admin_user_ids"]) == "" {
+ cfg["admin_user_ids"] = chatID
+ }
+}
+
+func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *TelegramMessage, userID string) error {
+ name := strings.TrimSpace(msg.From.FirstName)
+ if msg.From.Username != "" {
+ name = "@" + strings.TrimSpace(msg.From.Username)
+ }
+ telegramUserID := int64(msg.From.ID)
+ return s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
+ var existing model.TelegramBinding
+ err := tx.Where("telegram_user_id = ?", telegramUserID).First(&existing).Error
+ if err == nil {
+ if err := s.replaceTelegramAccountBindingTx(ctx, tx, userID, telegramUserID); err != nil {
+ return err
+ }
+ if err := tx.Model(&existing).Updates(map[string]any{
+ "telegram_name": name,
+ "chat_id": telegramBindingChatIDForMessage(msg, &existing),
+ "user_id": userID,
+ }).Error; telegramBindingUniqueErr(err) {
+ return errTelegramAccountAlreadyBound
+ } else if err != nil {
+ return err
+ }
+ return nil
+ }
+ if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
+ return err
+ }
+ if err := tx.Unscoped().Where("telegram_user_id = ?", telegramUserID).Delete(&model.TelegramBinding{}).Error; err != nil {
+ return err
+ }
+ if err := s.replaceTelegramAccountBindingTx(ctx, tx, userID, telegramUserID); err != nil {
+ return err
+ }
+ err = tx.Create(&model.TelegramBinding{
+ TelegramUserID: telegramUserID,
+ TelegramName: name,
+ ChatID: telegramBindingChatIDForMessage(msg, nil),
+ UserID: userID,
+ }).Error
+ if telegramBindingUniqueErr(err) {
+ return errTelegramAccountAlreadyBound
+ }
+ return err
+ })
+}
+
+func telegramBindingChatIDForMessage(msg *TelegramMessage, existing *model.TelegramBinding) int64 {
+ if msg == nil {
+ if existing != nil {
+ return existing.ChatID
+ }
+ return 0
+ }
+ if msg.Chat.Type == "" || msg.Chat.Type == "private" {
+ return int64(msg.Chat.ID)
+ }
+ if existing != nil && existing.ChatID > 0 {
+ return existing.ChatID
+ }
+ return int64(msg.From.ID)
+}
+
+func telegramPrivateChatIDFromBinding(binding model.TelegramBinding) int64 {
+ if binding.ChatID > 0 {
+ return binding.ChatID
+ }
+ return binding.TelegramUserID
+}
+
+func (s *TelegramBotService) replaceTelegramAccountBindingTx(ctx context.Context, tx *gorm.DB, userID string, telegramUserID int64) error {
+ return tx.WithContext(ctx).Unscoped().
+ Where("user_id = ? AND telegram_user_id <> ?", userID, telegramUserID).
+ Delete(&model.TelegramBinding{}).Error
+}
+
+func telegramBindingUniqueErr(err error) bool {
+ if err == nil {
+ return false
+ }
+ msg := strings.ToLower(err.Error())
+ return strings.Contains(msg, "idx_telegram_bindings_user_id_active") ||
+ strings.Contains(msg, "telegram_bindings.user_id") ||
+ (strings.Contains(msg, "unique") && strings.Contains(msg, "telegram_bindings"))
+}
+
+func parseStartCredentials(args []string) (string, string) {
+ if len(args) >= 2 {
+ return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " "))
+ }
+ if len(args) == 1 {
+ raw := strings.TrimSpace(args[0])
+ for _, sep := range []string{"-", ":", ":"} {
+ if parts := strings.SplitN(raw, sep, 2); len(parts) == 2 {
+ return strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1])
+ }
+ }
+ }
+ return "", ""
+}
+
+func userNameOrFallback(user *model.User) string {
+ if user == nil || strings.TrimSpace(user.Username) == "" {
+ return "未知用户"
+ }
+ return user.Username
+}
diff --git a/internal/service/telegram_bot.go b/internal/service/telegram_bot.go
index 41e349f..3faf93b 100644
--- a/internal/service/telegram_bot.go
+++ b/internal/service/telegram_bot.go
@@ -9,8 +9,6 @@ import (
"encoding/json"
"errors"
"fmt"
- "io"
- "net/http"
"strconv"
"strings"
"sync"
@@ -18,7 +16,6 @@ import (
"go.uber.org/zap"
"golang.org/x/crypto/bcrypt"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -86,17 +83,6 @@ type TelegramBotService struct {
pending map[int64]pendingInput // telegram_user_id -> awaited text input
}
-// TelegramPollingStartResult describes what happened when local long polling
-// was requested. The admin UI uses it to avoid a silent "started" toast when
-// no Telegram channel can actually poll.
-type TelegramPollingStartResult struct {
- Message string `json:"message"`
- Started int `json:"started"`
- AlreadyRunning int `json:"already_running"`
- Skipped int `json:"skipped"`
- Errors []string `json:"errors,omitempty"`
-}
-
// pendingInput tracks a button-initiated action that awaits the user's next
// text message (e.g. tapping「注册」then sending "用户名 密码").
type pendingInput struct {
@@ -602,954 +588,6 @@ func (s *TelegramBotService) cmdHideAdult(ctx context.Context, msg *TelegramMess
}
}
-// cmdStatus 处理 /status 命令。
-func (s *TelegramBotService) cmdStatus(ctx context.Context) (telegramCommandReply, error) {
- libraryIDs, err := s.activeTelegramStatsLibraryIDs(ctx)
- if err != nil {
- return telegramCommandReply{}, err
- }
- var mediaCount int64
- s.mediaStatsQuery(libraryIDs).Count(&mediaCount)
-
- var totalSize int64
- if err := s.mediaStatsQuery(libraryIDs).Select("COALESCE(SUM(size_bytes), 0)").Row().Scan(&totalSize); err != nil {
- return telegramCommandReply{}, err
- }
- totalSizeGB := float64(totalSize) / 1024 / 1024 / 1024
-
- return telegramCommandReply{Text: fmt.Sprintf(
- "系统运行状态\n\n"+
- "🎬 媒体总数: %d\n"+
- "💾 存储占用: %.1f GB",
- mediaCount, totalSizeGB,
- )}, nil
-}
-
-// cmdSearch 处理 /search 命令。
-func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (telegramCommandReply, error) {
- if len(args) == 0 {
- return telegramCommandReply{Text: "请提供搜索关键词\n例: /search 哥斯拉"}, nil
- }
-
- keyword := strings.Join(args, " ")
- var results []model.Media
- err := s.repo.DB.Where("title LIKE ?", "%"+keyword+"%").
- Order("year DESC").Limit(8).
- Find(&results).Error
- if err != nil {
- return telegramCommandReply{}, err
- }
-
- if len(results) == 0 {
- return telegramCommandReply{Text: fmt.Sprintf("未找到与 %s 相关的媒体", keyword)}, nil
- }
-
- var sb strings.Builder
- sb.WriteString(fmt.Sprintf("搜索: %s\n\n", keyword))
- for i, m := range results {
- year := ""
- if m.Year > 0 {
- year = fmt.Sprintf(" (%d)", m.Year)
- }
- ep := ""
- if m.SeasonNum > 0 && m.EpisodeNum > 0 {
- ep = fmt.Sprintf(" S%02dE%02d", m.SeasonNum, m.EpisodeNum)
- }
- sb.WriteString(fmt.Sprintf("%d. %s%s%s — %s\n", i+1, m.Title, year, ep, formatSize(m.SizeBytes)))
- }
-
- return telegramCommandReply{Text: sb.String()}, nil
-}
-
-// cmdDownloads 处理 /downloads 命令。
-func (s *TelegramBotService) cmdDownloads(ctx context.Context) (telegramCommandReply, error) {
- type Row struct {
- Title string
- Status string
- }
- var rows []Row
- if err := s.repo.DB.Raw(
- "SELECT COALESCE(NULLIF(title,''),'下载任务') as title, COALESCE(status,'unknown') as status FROM download_tasks ORDER BY created_at DESC LIMIT 8",
- ).Scan(&rows).Error; err != nil {
- return telegramCommandReply{}, err
- }
-
- if len(rows) == 0 {
- return telegramCommandReply{Text: "当前没有下载任务。"}, nil
- }
-
- var sb strings.Builder
- sb.WriteString(fmt.Sprintf("下载任务 (%d)\n\n", len(rows)))
- for _, r := range rows {
- icon := "⏳"
- switch r.Status {
- case "completed":
- icon = "✅"
- case "downloading":
- icon = "📥"
- case "error":
- icon = "❌"
- }
- name := strings.TrimSpace(r.Title)
- if name == "" {
- name = "下载任务"
- }
- if len(name) > 60 {
- name = name[:57] + "..."
- }
- sb.WriteString(fmt.Sprintf("%s %s\n", icon, name))
- }
-
- return telegramCommandReply{Text: sb.String()}, nil
-}
-
-// cmdStats 处理 /stats 命令。
-func (s *TelegramBotService) cmdStats(ctx context.Context) (telegramCommandReply, error) {
- libs, err := s.activeTelegramStatsLibraries(ctx)
- if err != nil {
- return telegramCommandReply{}, err
- }
- libraryIDs := make([]string, 0, len(libs))
- for _, lib := range libs {
- libraryIDs = append(libraryIDs, lib.ID)
- }
- var totalMedia int64
- s.mediaStatsQuery(libraryIDs).Count(&totalMedia)
-
- var totalSize int64
- if err := s.mediaStatsQuery(libraryIDs).Select("COALESCE(SUM(size_bytes), 0)").Row().Scan(&totalSize); err != nil {
- return telegramCommandReply{}, err
- }
-
- type LibStat struct {
- Name string
- Type string
- Count int64
- }
- stats := make([]LibStat, 0, len(libs))
- for _, lib := range libs {
- var count int64
- if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("library_id = ?", lib.ID).Count(&count).Error; err != nil {
- return telegramCommandReply{}, err
- }
- stats = append(stats, LibStat{Name: lib.Name, Type: lib.Type, Count: count})
- }
-
- var sb strings.Builder
- sb.WriteString("媒体库统计\n\n")
- sb.WriteString(fmt.Sprintf("📚 总数: %d\n", totalMedia))
- sb.WriteString(fmt.Sprintf("💾 大小: %s\n", formatSize(totalSize)))
-
- if len(stats) > 0 {
- sb.WriteString("\n各库分布:\n")
- for _, l := range stats {
- icon := "🎬"
- switch l.Type {
- case "tv":
- icon = "📺"
- case "anime":
- icon = "🍥"
- case "music":
- icon = "🎵"
- }
- sb.WriteString(fmt.Sprintf("%s %s: %d\n", icon, l.Name, l.Count))
- }
- }
-
- return telegramCommandReply{Text: sb.String()}, nil
-}
-
-func (s *TelegramBotService) activeTelegramStatsLibraries(ctx context.Context) ([]model.Library, error) {
- if s == nil || s.repo == nil || s.repo.Library == nil {
- return nil, nil
- }
- libs, err := s.repo.Library.List(ctx)
- if err != nil {
- return nil, err
- }
- libs = FilterDisplayCloudLibraries(ctx, s.repo, libs)
- out := libs[:0]
- for _, lib := range libs {
- if lib.Enabled {
- out = append(out, lib)
- }
- }
- return out, nil
-}
-
-func (s *TelegramBotService) activeTelegramStatsLibraryIDs(ctx context.Context) ([]string, error) {
- libs, err := s.activeTelegramStatsLibraries(ctx)
- if err != nil {
- return nil, err
- }
- ids := make([]string, 0, len(libs))
- for _, lib := range libs {
- ids = append(ids, lib.ID)
- }
- return ids, nil
-}
-
-func (s *TelegramBotService) mediaStatsQuery(libraryIDs []string) *gorm.DB {
- q := s.repo.DB.Model(&model.Media{})
- if len(libraryIDs) == 0 {
- return q.Where("1 = 0")
- }
- return q.Where("library_id IN ?", libraryIDs)
-}
-
-// ── Polling ──
-
-// StartPolling 为所有已启用的 Telegram 通知渠道启动长轮询。
-func (s *TelegramBotService) StartPolling(ctx context.Context) TelegramPollingStartResult {
- result := TelegramPollingStartResult{Message: "telegram polling started"}
- channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
- if err != nil {
- s.log.Error("failed to list telegram channels for polling", zap.Error(err))
- result.Message = "failed to list telegram channels"
- result.Errors = append(result.Errors, err.Error())
- return result
- }
- if len(channels) == 0 {
- result.Message = "no telegram channels configured"
- result.Errors = append(result.Errors, "没有配置 Telegram 通知渠道")
- return result
- }
-
- for _, ch := range channels {
- if !ch.Enabled {
- result.Skipped++
- result.Errors = append(result.Errors, ch.Name+": 通知渠道未启用")
- continue
- }
- configStr := ch.Config
- if s.crypto != nil && configStr != "" {
- configStr = s.crypto.Decrypt(configStr)
- }
- var rawCfg map[string]any
- if err := json.Unmarshal([]byte(configStr), &rawCfg); err != nil {
- result.Skipped++
- result.Errors = append(result.Errors, ch.Name+": Telegram 配置解析失败: "+err.Error())
- continue
- }
- cfg := telegramStringConfigFromAny(rawCfg)
- botToken := cfg["bot_token"]
- if botToken == "" {
- result.Skipped++
- result.Errors = append(result.Errors, ch.Name+": Telegram Bot Token 为空")
- continue
- }
- s.pollingMu.Lock()
- if _, running := s.pollingCancel[botToken]; running {
- s.pollingMu.Unlock()
- result.AlreadyRunning++
- continue
- }
- s.pollingMu.Unlock()
-
- if err := registerTelegramBotCommands(ctx, cfg); err != nil && s.log != nil {
- s.log.Warn("telegram setMyCommands failed", zap.Error(sanitizeTelegramError(err)))
- }
- if err := deleteTelegramWebhook(ctx, cfg); err != nil {
- result.Skipped++
- result.Errors = append(result.Errors, ch.Name+": "+sanitizeTelegramError(err).Error())
- continue
- }
-
- s.pollingMu.Lock()
- if _, running := s.pollingCancel[botToken]; running {
- s.pollingMu.Unlock()
- result.AlreadyRunning++
- continue
- }
- pollCtx, cancel := context.WithCancel(context.Background())
- s.pollingCancel[botToken] = cancel
- s.pollingMu.Unlock()
-
- channel := ch
- go s.pollLoop(pollCtx, cfg, &channel)
- result.Started++
- s.log.Info("started telegram polling", zap.String("channel", ch.Name))
- }
- if result.Started == 0 && result.AlreadyRunning == 0 {
- result.Message = "no enabled telegram channels started"
- }
- return result
-}
-
-// StopPolling 停止所有 Telegram 长轮询。
-func (s *TelegramBotService) StopPolling() int {
- s.pollingMu.Lock()
- defer s.pollingMu.Unlock()
- stopped := 0
- for token, cancel := range s.pollingCancel {
- cancel()
- delete(s.pollingCancel, token)
- stopped++
- }
- s.log.Info("telegram polling stopped")
- return stopped
-}
-
-// pollLoop 对单个 Bot Token 执行长轮询。
-func (s *TelegramBotService) pollLoop(ctx context.Context, cfg map[string]string, channel *model.NotifyChannel) {
- var offset int64 = 0
- pollURL, err := telegramMethodURL(cfg, cfg["bot_token"], "getUpdates")
- if err != nil {
- s.log.Warn("telegram polling config invalid", zap.Error(err))
- return
- }
- clients := telegramHTTPClients(45*time.Second, cfg)
-
- for {
- select {
- case <-ctx.Done():
- return
- default:
- }
-
- reqBody, _ := json.Marshal(map[string]interface{}{
- "offset": offset,
- "timeout": 30,
- })
- respBody, err := telegramPollingRequest(ctx, clients, pollURL, string(reqBody))
- if err != nil {
- s.log.Debug("telegram polling failed", zap.Error(err))
- time.Sleep(5 * time.Second)
- continue
- }
-
- var result struct {
- OK bool `json:"ok"`
- Result []TelegramUpdate `json:"result"`
- }
- if err := json.Unmarshal(respBody, &result); err != nil || !result.OK {
- time.Sleep(3 * time.Second)
- continue
- }
-
- for _, upd := range result.Result {
- if upd.UpdateID >= int(offset) {
- offset = int64(upd.UpdateID) + 1
- }
- if !telegramUpdateActionable(upd) {
- continue
- }
- go func(u TelegramUpdate) {
- handlerCtx, cancel := context.WithTimeout(ctx, 2*time.Minute)
- defer cancel()
- _ = s.handleTelegramUpdate(handlerCtx, u, channel)
- }(upd)
- }
- }
-}
-
-// telegramUpdateActionable 判断一条 update 是否需要分发处理。
-// 长轮询默认会返回 message 与 callback_query 两类更新;命令消息需有文本,
-// 而内联按钮回调(callback_query)必须被分发,否则成人目录显隐开关会失效。
-func telegramUpdateActionable(upd TelegramUpdate) bool {
- if upd.CallbackQuery != nil {
- return true
- }
- return upd.Message != nil && upd.Message.Text != ""
-}
-
-func telegramPollingRequest(ctx context.Context, clients []*http.Client, pollURL, body string) ([]byte, error) {
- var lastErr error
- for _, client := range clients {
- req, err := http.NewRequestWithContext(ctx, http.MethodPost, pollURL, strings.NewReader(body))
- if err != nil {
- return nil, err
- }
- req.Header.Set("Content-Type", "application/json")
- resp, err := client.Do(req)
- if err != nil {
- lastErr = sanitizeTelegramError(err)
- continue
- }
- respBody, _ := io.ReadAll(resp.Body)
- _ = resp.Body.Close()
- if resp.StatusCode >= 400 {
- lastErr = fmt.Errorf("telegram api error %d: %s", resp.StatusCode, sanitizeTelegramText(string(respBody)))
- continue
- }
- return respBody, nil
- }
- if lastErr != nil {
- return nil, lastErr
- }
- return nil, errors.New("telegram polling failed")
-}
-
-// ── Message Sending ──
-
-const defaultTelegramMessageDeleteDelay = 120 * time.Second
-
-type telegramSendMessageResponse struct {
- OK bool `json:"ok"`
- Result struct {
- MessageID int `json:"message_id"`
- } `json:"result"`
-}
-
-// reply 通过 Telegram Bot API 发送回复消息。
-func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error {
- cfg := s.telegramChannelConfig(channel)
- if strings.TrimSpace(cfg["bot_token"]) == "" {
- return fmt.Errorf("bot_token not configured")
- }
-
- payload := map[string]interface{}{
- "chat_id": strconv.Itoa(chatID),
- "text": reply.Text,
- "parse_mode": "HTML",
- }
- if len(reply.Buttons) > 0 {
- keyboard := make([][]map[string]string, 0, len(reply.Buttons))
- for _, row := range reply.Buttons {
- buttons := make([]map[string]string, 0, len(row))
- for _, button := range row {
- buttons = append(buttons, map[string]string{
- "text": button.Text,
- "callback_data": button.Data,
- })
- }
- keyboard = append(keyboard, buttons)
- }
- payload["reply_markup"] = map[string]interface{}{"inline_keyboard": keyboard}
- }
- var sent telegramSendMessageResponse
- if err := telegramPostJSONDecode(ctx, cfg, "sendMessage", payload, 15*time.Second, &sent); err != nil {
- return err
- }
- if sent.Result.MessageID > 0 {
- s.scheduleTelegramMessageDelete(cfg, chatID, sent.Result.MessageID)
- }
- return nil
-}
-
-func (s *TelegramBotService) replyForMessage(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, reply telegramCommandReply) error {
- if msg == nil {
- return nil
- }
- if strings.TrimSpace(reply.Text) == "" {
- return nil
- }
- return s.reply(ctx, channel, msg.Chat.ID, reply)
-}
-
-func (s *TelegramBotService) deleteTelegramSourceMessage(channel *model.NotifyChannel, chatID, messageID int) {
- if messageID <= 0 {
- return
- }
- s.scheduleTelegramMessageDelete(s.telegramChannelConfig(channel), chatID, messageID)
-}
-
-func (s *TelegramBotService) scheduleTelegramMessageDelete(cfg map[string]string, chatID, messageID int) {
- if chatID == 0 || messageID <= 0 || strings.TrimSpace(cfg["bot_token"]) == "" {
- return
- }
- delay := telegramMessageDeleteDelay(cfg)
- if delay < 0 {
- return
- }
- cfgCopy := make(map[string]string, len(cfg))
- for k, v := range cfg {
- cfgCopy[k] = v
- }
- go func() {
- if delay > 0 {
- timer := time.NewTimer(delay)
- defer timer.Stop()
- <-timer.C
- }
- deleteCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
- defer cancel()
- err := telegramPostJSON(deleteCtx, cfgCopy, "deleteMessage", map[string]interface{}{
- "chat_id": strconv.Itoa(chatID),
- "message_id": messageID,
- }, 10*time.Second)
- if err != nil && s.log != nil {
- s.log.Debug("telegram deleteMessage failed",
- zap.Int("chat_id", chatID),
- zap.Int("message_id", messageID),
- zap.Error(sanitizeTelegramError(err)),
- )
- }
- }()
-}
-
-func telegramMessageDeleteDelay(cfg map[string]string) time.Duration {
- for _, key := range []string{"auto_delete_seconds", "message_delete_seconds", "delete_after_seconds"} {
- raw := strings.TrimSpace(cfg[key])
- if raw == "" {
- continue
- }
- seconds, err := strconv.Atoi(raw)
- if err != nil {
- continue
- }
- if seconds < 0 {
- return -1
- }
- return time.Duration(seconds) * time.Second
- }
- return defaultTelegramMessageDeleteDelay
-}
-
-// findChannelByChatID 根据 chat_id 查找已配置的通知渠道。
-func (s *TelegramBotService) findChannelByChatID(ctx context.Context, chatID int) *model.NotifyChannel {
- channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
- if err != nil {
- return nil
- }
- target := strconv.Itoa(chatID)
- for _, ch := range channels {
- if !ch.Enabled {
- continue
- }
- configStr := ch.Config
- if s.crypto != nil && configStr != "" {
- configStr = s.crypto.Decrypt(configStr)
- }
- var cfg map[string]string
- if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
- continue
- }
- if cfg["chat_id"] == target || cfg["command_chat_id"] == target ||
- cfg["group_chat_id"] == target || cfg["channel_chat_id"] == target {
- return &ch
- }
- }
- if len(channels) == 1 && channels[0].Enabled {
- return &channels[0]
- }
- return nil
-}
-
-func (s *TelegramBotService) findChannelForMessage(ctx context.Context, msg *TelegramMessage) *model.NotifyChannel {
- if msg == nil {
- return nil
- }
- if msg.Chat.Type != "" && msg.Chat.Type != "private" {
- return s.findChannelByChatID(ctx, msg.Chat.ID)
- }
- channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
- if err != nil {
- return nil
- }
- var first *model.NotifyChannel
- for i := range channels {
- ch := channels[i]
- if !ch.Enabled {
- continue
- }
- if first == nil {
- first = &ch
- }
- if s.telegramUserIsAdmin(ctx, &ch, msg.From.ID) || s.telegramUserCanBind(ctx, &ch, msg.From.ID) {
- return &ch
- }
- }
- return first
-}
-
-func (s *TelegramBotService) channelForMessage(ctx context.Context, msg *TelegramMessage, hint *model.NotifyChannel) *model.NotifyChannel {
- if hint == nil {
- return s.findChannelForMessage(ctx, msg)
- }
- if msg == nil {
- return hint
- }
- if msg.Chat.Type != "" && msg.Chat.Type != "private" && !s.telegramChatAllowed(hint, msg.Chat.ID) {
- return nil
- }
- return hint
-}
-
-func (s *TelegramBotService) handleCallback(ctx context.Context, cb *TelegramCallbackQuery, channelHint *model.NotifyChannel) error {
- if cb == nil || cb.Message == nil {
- return nil
- }
- msg := *cb.Message
- msg.From = cb.From
- channel := s.channelForMessage(ctx, &msg, channelHint)
- if channel == nil {
- channel = s.findChannelByChatID(ctx, cb.Message.Chat.ID)
- }
- // 立即应答回调,关闭按钮上的加载状态,避免客户端长时间转圈。
- if telegramIsGroupChat(cb.Message.Chat.Type) {
- s.answerCallbackWithText(ctx, channel, cb.ID, "为了隐私,群组内按钮面板已禁用。请私聊 Bot 或在群里发送 /menu,我会把面板私聊给你。", true)
- s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
- return nil
- }
- if cb.Message.Chat.Type == "private" && cb.Message.Chat.ID != cb.From.ID {
- s.answerCallbackWithText(ctx, channel, cb.ID, "这个面板不属于你,请发送 /menu 打开自己的面板。", true)
- return nil
- }
- s.answerCallback(ctx, channel, cb.ID)
- data := strings.TrimSpace(cb.Data)
- if data == "adult_toggle" {
- reply := s.cmdHideAdult(ctx, &msg, nil)
- if reply.Text != "" {
- err := s.reply(ctx, channel, cb.Message.Chat.ID, reply)
- s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
- return err
- }
- return nil
- }
- if reply, handled := s.handleMenuCallback(ctx, channel, &msg, data); handled {
- if reply.Text != "" {
- err := s.reply(ctx, channel, cb.Message.Chat.ID, reply)
- s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID)
- return err
- }
- }
- return nil
-}
-
-// answerCallback 应答 Telegram 回调查询,关闭按钮上的加载提示。
-func (s *TelegramBotService) answerCallback(ctx context.Context, channel *model.NotifyChannel, callbackID string) {
- s.answerCallbackWithText(ctx, channel, callbackID, "", false)
-}
-
-func (s *TelegramBotService) answerCallbackWithText(ctx context.Context, channel *model.NotifyChannel, callbackID, text string, showAlert bool) {
- if channel == nil || strings.TrimSpace(callbackID) == "" {
- return
- }
- cfg := s.telegramChannelConfig(channel)
- if strings.TrimSpace(cfg["bot_token"]) == "" {
- return
- }
- payload := map[string]interface{}{
- "callback_query_id": callbackID,
- }
- if strings.TrimSpace(text) != "" {
- payload["text"] = text
- payload["show_alert"] = showAlert
- }
- if err := telegramPostJSON(ctx, cfg, "answerCallbackQuery", payload, 8*time.Second); err != nil {
- s.log.Debug("telegram answerCallbackQuery failed", zap.Error(sanitizeTelegramError(err)))
- }
-}
-
-func (s *TelegramBotService) telegramBinding(ctx context.Context, telegramUserID int) *model.TelegramBinding {
- if telegramUserID == 0 {
- return nil
- }
- var binding model.TelegramBinding
- err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(telegramUserID)).First(&binding).Error
- if err != nil {
- return nil
- }
- return &binding
-}
-
-func (s *TelegramBotService) unbindTelegramUser(ctx context.Context, telegramUserID int) error {
- if s == nil || s.repo == nil || s.repo.DB == nil || telegramUserID == 0 {
- return nil
- }
- return s.repo.DB.WithContext(ctx).Unscoped().
- Where("telegram_user_id = ?", int64(telegramUserID)).
- Delete(&model.TelegramBinding{}).Error
-}
-
-func (s *TelegramBotService) telegramUserIsAdmin(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool {
- if s.telegramUserIDConfigured(channel, telegramUserID) {
- return true
- }
- binding := s.telegramBinding(ctx, telegramUserID)
- if binding == nil {
- return false
- }
- user, err := s.repo.User.FindByID(ctx, binding.UserID)
- return err == nil && user != nil && user.Role == "admin" && user.IsActive
-}
-
-func (s *TelegramBotService) telegramChatAllowed(channel *model.NotifyChannel, chatID int) bool {
- if channel == nil {
- return false
- }
- configStr := channel.Config
- if s.crypto != nil && configStr != "" {
- configStr = s.crypto.Decrypt(configStr)
- }
- var cfg map[string]string
- if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
- return false
- }
- target := strconv.Itoa(chatID)
- for _, key := range []string{"group_chat_id", "channel_chat_id", "command_chat_id"} {
- if configured := strings.TrimSpace(cfg[key]); configured != "" && configured == target {
- return true
- }
- }
- if strings.TrimSpace(cfg["group_chat_id"]) != "" || strings.TrimSpace(cfg["channel_chat_id"]) != "" || strings.TrimSpace(cfg["command_chat_id"]) != "" {
- return false
- }
- return strings.TrimSpace(cfg["chat_id"]) == target
-}
-
-// telegramBindDecision 表示成员资格校验的三态结果:通过 / 明确不通过 /
-// 无法验证(getChatMember 出错,如 Bot 不在群、群 ID 失效、网络或代理不可达)。
-// 区分「明确不是成员」和「查不了」,是为了避免把验证失败误报成「你不在群」。
-type telegramBindDecision int
-
-const (
- bindDenied telegramBindDecision = iota // 已查实:不在任何绑定群组/频道
- bindAllowed // 管理员,或查实是某绑定群组/频道成员
- bindUnverifiable // 配了群组/频道但 getChatMember 全部失败
-)
-
-// telegramMembership 表示单个 chat 的成员资格三态。
-type telegramMembership int
-
-const (
- membershipNo telegramMembership = iota // 查实不是成员(left/kicked 等)
- membershipYes // 查实是成员
- membershipUnknown // getChatMember 出错,无法判定
-)
-
-func (s *TelegramBotService) telegramUserBindDecision(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) telegramBindDecision {
- if telegramUserID == 0 || channel == nil {
- return bindDenied
- }
- if s.telegramUserIDConfigured(channel, telegramUserID) {
- return bindAllowed
- }
- chatIDs := s.telegramMembershipChatIDs(channel)
- if len(chatIDs) == 0 {
- return bindDenied
- }
- sawUnknown := false
- for _, chatID := range chatIDs {
- switch s.telegramChatMembership(ctx, channel, chatID, telegramUserID) {
- case membershipYes:
- return bindAllowed
- case membershipUnknown:
- sawUnknown = true
- }
- }
- if sawUnknown {
- return bindUnverifiable
- }
- return bindDenied
-}
-
-// telegramUserCanBind 是 telegramUserBindDecision 的布尔包装,供尽力而为的场景
-// 使用(如私聊时挑选可用渠道):只有查实通过才返回 true。
-func (s *TelegramBotService) telegramUserCanBind(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool {
- return s.telegramUserBindDecision(ctx, channel, telegramUserID) == bindAllowed
-}
-
-// telegramBindRejectText 根据三态结果生成面向用户的提示。action 形如「兑换注册账号」
-// 「绑定媒体中心账号」。bindUnverifiable 时不再误导用户「你不在群」,而是提示
-// 管理员检查 Bot 权限与群组 ID。
-func telegramBindRejectText(decision telegramBindDecision, action string) string {
- if decision == bindUnverifiable {
- return fmt.Sprintf("暂时无法验证你的群组/频道成员身份,%s未成功。这通常是因为 Bot 未加入绑定群组、在频道中不是管理员,或群组 ID 配置有误(如超级群需带 -100 前缀)。请联系管理员检查 Bot 权限与「绑定群组/频道 ID」。", action)
- }
- return fmt.Sprintf("当前 Telegram 账号不在管理员配置的绑定群组/频道中,无法%s。请先加入管理员配置的群组或频道;如果尚未配置,请联系管理员。", action)
-}
-
-func (s *TelegramBotService) telegramChatMembership(ctx context.Context, channel *model.NotifyChannel, chatID string, telegramUserID int) telegramMembership {
- cfg := s.telegramChannelConfig(channel)
- if strings.TrimSpace(cfg["bot_token"]) == "" || chatID == "" || telegramUserID == 0 {
- return membershipUnknown
- }
- payload := map[string]interface{}{
- "chat_id": chatID,
- "user_id": telegramUserID,
- }
- var result struct {
- OK bool `json:"ok"`
- Result struct {
- Status string `json:"status"`
- } `json:"result"`
- }
- if err := telegramPostJSONDecode(ctx, cfg, "getChatMember", payload, 15*time.Second, &result); err != nil {
- s.log.Warn("telegram getChatMember failed", zap.String("chat_id", chatID), zap.Int("telegram_user_id", telegramUserID), zap.Error(sanitizeTelegramError(err)))
- return membershipUnknown
- }
- if !result.OK {
- return membershipUnknown
- }
- switch strings.ToLower(result.Result.Status) {
- case "creator", "administrator", "member", "restricted":
- return membershipYes
- default:
- return membershipNo
- }
-}
-
-// telegramUserIsChatMember 是 telegramChatMembership 的布尔包装,仅在查实是成员时
-// 返回 true(查不了也视为非成员,供尽力而为的场景使用)。
-func (s *TelegramBotService) telegramUserIsChatMember(ctx context.Context, channel *model.NotifyChannel, chatID string, telegramUserID int) bool {
- return s.telegramChatMembership(ctx, channel, chatID, telegramUserID) == membershipYes
-}
-
-func (s *TelegramBotService) telegramUserIDConfigured(channel *model.NotifyChannel, telegramUserID int) bool {
- if channel == nil || telegramUserID == 0 {
- return false
- }
- cfg := s.telegramChannelConfig(channel)
- target := strconv.Itoa(telegramUserID)
- for _, value := range telegramConfiguredUserIDs(cfg["admin_user_ids"]) {
- if value == target {
- return true
- }
- }
- if strings.TrimSpace(cfg["admin_user_ids"]) == "" && strings.TrimSpace(cfg["chat_id"]) == target {
- return true
- }
- return false
-}
-
-func (s *TelegramBotService) telegramChannelConfig(channel *model.NotifyChannel) map[string]string {
- return telegramConfigFromChannel(s.crypto, channel)
-}
-
-func telegramConfigFromChannel(crypto *CryptoService, channel *model.NotifyChannel) map[string]string {
- if channel == nil {
- return map[string]string{}
- }
- configStr := channel.Config
- if crypto != nil && configStr != "" {
- configStr = crypto.Decrypt(configStr)
- }
- var cfg map[string]string
- if err := json.Unmarshal([]byte(configStr), &cfg); err != nil || cfg == nil {
- return map[string]string{}
- }
- normalizeTelegramConfig(cfg)
- return cfg
-}
-
-func normalizeTelegramConfig(cfg map[string]string) {
- if cfg == nil {
- return
- }
- chatID := strings.TrimSpace(cfg["chat_id"])
- if chatID == "" {
- return
- }
- if strings.HasPrefix(chatID, "-") {
- if strings.TrimSpace(cfg["group_chat_id"]) == "" && strings.TrimSpace(cfg["channel_chat_id"]) == "" && strings.TrimSpace(cfg["command_chat_id"]) == "" {
- cfg["group_chat_id"] = chatID
- }
- return
- }
- if strings.TrimSpace(cfg["admin_user_ids"]) == "" {
- cfg["admin_user_ids"] = chatID
- }
-}
-
-func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *TelegramMessage, userID string) error {
- name := strings.TrimSpace(msg.From.FirstName)
- if msg.From.Username != "" {
- name = "@" + strings.TrimSpace(msg.From.Username)
- }
- telegramUserID := int64(msg.From.ID)
- return s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
- var existing model.TelegramBinding
- err := tx.Where("telegram_user_id = ?", telegramUserID).First(&existing).Error
- if err == nil {
- if err := s.replaceTelegramAccountBindingTx(ctx, tx, userID, telegramUserID); err != nil {
- return err
- }
- if err := tx.Model(&existing).Updates(map[string]any{
- "telegram_name": name,
- "chat_id": telegramBindingChatIDForMessage(msg, &existing),
- "user_id": userID,
- }).Error; telegramBindingUniqueErr(err) {
- return errTelegramAccountAlreadyBound
- } else if err != nil {
- return err
- }
- return nil
- }
- if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
- return err
- }
- if err := tx.Unscoped().Where("telegram_user_id = ?", telegramUserID).Delete(&model.TelegramBinding{}).Error; err != nil {
- return err
- }
- if err := s.replaceTelegramAccountBindingTx(ctx, tx, userID, telegramUserID); err != nil {
- return err
- }
- err = tx.Create(&model.TelegramBinding{
- TelegramUserID: telegramUserID,
- TelegramName: name,
- ChatID: telegramBindingChatIDForMessage(msg, nil),
- UserID: userID,
- }).Error
- if telegramBindingUniqueErr(err) {
- return errTelegramAccountAlreadyBound
- }
- return err
- })
-}
-
-func telegramBindingChatIDForMessage(msg *TelegramMessage, existing *model.TelegramBinding) int64 {
- if msg == nil {
- if existing != nil {
- return existing.ChatID
- }
- return 0
- }
- if msg.Chat.Type == "" || msg.Chat.Type == "private" {
- return int64(msg.Chat.ID)
- }
- if existing != nil && existing.ChatID > 0 {
- return existing.ChatID
- }
- return int64(msg.From.ID)
-}
-
-func telegramPrivateChatIDFromBinding(binding model.TelegramBinding) int64 {
- if binding.ChatID > 0 {
- return binding.ChatID
- }
- return binding.TelegramUserID
-}
-
-func (s *TelegramBotService) replaceTelegramAccountBindingTx(ctx context.Context, tx *gorm.DB, userID string, telegramUserID int64) error {
- return tx.WithContext(ctx).Unscoped().
- Where("user_id = ? AND telegram_user_id <> ?", userID, telegramUserID).
- Delete(&model.TelegramBinding{}).Error
-}
-
-func telegramBindingUniqueErr(err error) bool {
- if err == nil {
- return false
- }
- msg := strings.ToLower(err.Error())
- return strings.Contains(msg, "idx_telegram_bindings_user_id_active") ||
- strings.Contains(msg, "telegram_bindings.user_id") ||
- (strings.Contains(msg, "unique") && strings.Contains(msg, "telegram_bindings"))
-}
-
-func parseStartCredentials(args []string) (string, string) {
- if len(args) >= 2 {
- return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " "))
- }
- if len(args) == 1 {
- raw := strings.TrimSpace(args[0])
- for _, sep := range []string{"-", ":", ":"} {
- if parts := strings.SplitN(raw, sep, 2); len(parts) == 2 {
- return strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1])
- }
- }
- }
- return "", ""
-}
-
-func userNameOrFallback(user *model.User) string {
- if user == nil || strings.TrimSpace(user.Username) == "" {
- return "未知用户"
- }
- return user.Username
-}
-
// ── Webhook Management ──
// SetWebhook 注册 Telegram Bot Webhook URL。
@@ -1574,21 +612,3 @@ func (s *TelegramBotService) GetWebhookInfo(ctx context.Context, botToken string
}
return result, nil
}
-
-// formatSize 格式化字节数为可读字符串。
-func formatSize(bytes int64) string {
- if bytes <= 0 {
- return "0 B"
- }
- units := []string{"B", "KB", "MB", "GB", "TB"}
- v := float64(bytes)
- i := 0
- for v >= 1024 && i < len(units)-1 {
- v /= 1024
- i++
- }
- if i == 0 {
- return fmt.Sprintf("%.0f %s", v, units[i])
- }
- return fmt.Sprintf("%.1f %s", v, units[i])
-}
diff --git a/internal/service/telegram_cleanup_rules.go b/internal/service/telegram_cleanup_rules.go
new file mode 100644
index 0000000..6103b83
--- /dev/null
+++ b/internal/service/telegram_cleanup_rules.go
@@ -0,0 +1,220 @@
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "fmt"
+ "strconv"
+ "strings"
+)
+
+func (s *TelegramBotService) currentCleanupRules(ctx context.Context) []accountCleanupRule {
+ cfg := loadBotConfig(ctx, s.repo)
+ return cfg.AccountCleanupRules
+}
+
+func (s *TelegramBotService) saveCleanupRules(ctx context.Context, rules []accountCleanupRule) error {
+ raw, err := json.Marshal(normalizeCleanupRules(rules))
+ if err != nil {
+ return err
+ }
+ return s.repo.Setting.Set(ctx, SettingAccountCleanupRules, string(raw))
+}
+
+func parseCommandBool(value string) (bool, bool) {
+ switch strings.ToLower(strings.TrimSpace(value)) {
+ case "on", "true", "1", "yes", "enable", "enabled", "开启", "开":
+ return true, true
+ case "off", "false", "0", "no", "disable", "disabled", "关闭", "关":
+ return false, true
+ default:
+ return false, false
+ }
+}
+
+func parseCleanupRuleCommand(args []string) (accountCleanupRule, error) {
+ if len(args) < 2 {
+ return accountCleanupRule{}, fmt.Errorf("新增规则参数不足")
+ }
+ rule := accountCleanupRule{
+ Type: strings.ToLower(strings.TrimSpace(args[0])),
+ ID: strings.TrimSpace(args[1]),
+ Enabled: true,
+ WindowDaysMin: 3,
+ WindowDaysMax: 5,
+ MinHours: 6,
+ MinCount: 1,
+ }
+ switch rule.Type {
+ case "watch_hours":
+ name, values := cleanupRuleNameAndValues(args[2:], 3)
+ rule.Name = name
+ if len(values) >= 3 {
+ rule.WindowDaysMin, _ = strconv.Atoi(values[0])
+ rule.WindowDaysMax, _ = strconv.Atoi(values[1])
+ rule.MinHours, _ = strconv.ParseFloat(values[2], 64)
+ if rule.Name == "" {
+ rule.Name = fmt.Sprintf("%d~%d 天观看满 %s 小时", rule.WindowDaysMin, rule.WindowDaysMax, formatRuleHours(rule.MinHours))
+ }
+ }
+ case "recent_login":
+ name, values := cleanupRuleNameAndValues(args[2:], 1)
+ rule.Name = name
+ if len(values) >= 1 {
+ rule.WindowDaysMax, _ = strconv.Atoi(values[0])
+ if rule.Name == "" {
+ rule.Name = fmt.Sprintf("%d 天内登录", rule.WindowDaysMax)
+ }
+ }
+ case "signin_streak", "account_age_grace":
+ name, values := cleanupRuleNameAndValues(args[2:], 1)
+ rule.Name = name
+ if len(values) >= 1 {
+ rule.MinCount, _ = strconv.Atoi(values[0])
+ if rule.Name == "" {
+ if rule.Type == "signin_streak" {
+ rule.Name = fmt.Sprintf("连续签到 %d 天", rule.MinCount)
+ } else {
+ rule.Name = fmt.Sprintf("新号宽限 %d 天", rule.MinCount)
+ }
+ }
+ }
+ default:
+ return accountCleanupRule{}, fmt.Errorf("不支持的规则类型:%s", rule.Type)
+ }
+ normalized := normalizeCleanupRules([]accountCleanupRule{rule})
+ if len(normalized) == 0 {
+ return accountCleanupRule{}, fmt.Errorf("规则无效")
+ }
+ return normalized[0], nil
+}
+
+func cleanupRuleNameAndValues(args []string, numericCount int) (string, []string) {
+ if len(args) == 0 {
+ return "", nil
+ }
+ if len(args) >= numericCount && cleanupRuleValuesAreNumeric(args[:numericCount]) {
+ return "", args
+ }
+ return strings.TrimSpace(args[0]), args[1:]
+}
+
+func cleanupRuleValuesAreNumeric(values []string) bool {
+ for _, value := range values {
+ if _, err := strconv.ParseFloat(strings.TrimSpace(value), 64); err != nil {
+ return false
+ }
+ }
+ return true
+}
+
+func formatCleanupRules(rules []accountCleanupRule) string {
+ if len(rules) == 0 {
+ return "保号规则\n\n暂无规则。"
+ }
+ var sb strings.Builder
+ sb.WriteString("保号规则\n")
+ for i, r := range rules {
+ state := map[bool]string{true: "启用", false: "停用"}[r.Enabled]
+ detail := cleanupRuleDetail(r)
+ parts := []string{
+ fmt.Sprintf("\n%d. %s", i+1, r.ID),
+ }
+ if shouldShowCleanupRuleName(r, detail) {
+ parts = append(parts, r.Name)
+ }
+ parts = append(parts, cleanupRuleTypeLabel(r.Type), state)
+ if detail != "" {
+ parts = append(parts, detail)
+ }
+ sb.WriteString(strings.Join(parts, " · "))
+ }
+ return sb.String()
+}
+
+func shouldShowCleanupRuleName(r accountCleanupRule, detail string) bool {
+ name := strings.TrimSpace(r.Name)
+ if name == "" || strings.EqualFold(name, r.ID) {
+ return false
+ }
+ if detail != "" && strings.EqualFold(name, detail) {
+ return false
+ }
+ return true
+}
+
+func cleanupRuleDetail(r accountCleanupRule) string {
+ switch r.Type {
+ case "watch_hours":
+ return fmt.Sprintf("%d~%d 天 %s 小时", r.WindowDaysMin, r.WindowDaysMax, formatRuleHours(r.MinHours))
+ case "recent_login":
+ return fmt.Sprintf("%d 天内登录", r.WindowDaysMax)
+ case "signin_streak":
+ return fmt.Sprintf("连续签到 %d 天", r.MinCount)
+ case "account_age_grace":
+ return fmt.Sprintf("新号宽限 %d 天", r.MinCount)
+ default:
+ return ""
+ }
+}
+
+func formatRuleHours(hours float64) string {
+ if hours == float64(int(hours)) {
+ return strconv.Itoa(int(hours))
+ }
+ return fmt.Sprintf("%.1f", hours)
+}
+
+func cleanupRuleTypeLabel(t string) string {
+ switch t {
+ case "watch_hours":
+ return "观看时长"
+ case "recent_login":
+ return "最近登录"
+ case "signin_streak":
+ return "连续签到"
+ case "account_age_grace":
+ return "新号宽限"
+ default:
+ return t
+ }
+}
+
+func cleanupRuleHelp() string {
+ return "Mgo 保号规则命令\n\n" +
+ "/cleanup_rule list — 查看规则\n" +
+ "/cleanup_rule add watch_hours watch_3_5d_6h 观看3到5天满6小时 3 5 6\n" +
+ "/cleanup_rule add recent_login login_7d 七天内登录 7\n" +
+ "/cleanup_rule add signin_streak sign_3 连续签到3天 3\n" +
+ "/cleanup_rule add account_age_grace new_7d 新号宽限7天 7\n" +
+ "/cleanup_rule edit 规则类型 规则ID 名称 参数... — 修改同 ID 规则\n" +
+ "/cleanup_rule 修改 规则类型 规则ID 名称 参数... — 中文修改入口\n" +
+ "/cleanup_rule enable 规则ID / disable 规则ID\n" +
+ "/cleanup_rule del 规则ID\n\n" +
+ "保号模式固定为:满足任意一条启用规则即保留;全部不满足才会清理。"
+}
+
+func onOff(b bool) string {
+ return map[bool]string{true: "已开启", false: "已关闭"}[b]
+}
+
+func toggleLabel(name string, enabled bool) string {
+ if enabled {
+ return "关闭" + name
+ }
+ return "开启" + name
+}
+
+func cleanupModeLabel(mode string) string {
+ return "满足任意一条"
+}
+
+func countEnabledCleanupRules(rules []accountCleanupRule) int {
+ n := 0
+ for _, r := range rules {
+ if r.Enabled {
+ n++
+ }
+ }
+ return n
+}
diff --git a/internal/service/telegram_cleanup_rules_test.go b/internal/service/telegram_cleanup_rules_test.go
new file mode 100644
index 0000000..47f5fdf
--- /dev/null
+++ b/internal/service/telegram_cleanup_rules_test.go
@@ -0,0 +1,34 @@
+package service
+
+import (
+ "strings"
+ "testing"
+)
+
+func TestParseCleanupRuleCommandWithNamedWatchHours(t *testing.T) {
+ rule, err := parseCleanupRuleCommand([]string{"watch_hours", "watch_3_5d_6h", "观看3到5天满6小时", "3", "5", "6"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if rule.Type != "watch_hours" || rule.ID != "watch_3_5d_6h" || rule.Name != "观看3到5天满6小时" {
+ t.Fatalf("unexpected rule identity: %+v", rule)
+ }
+ if !rule.Enabled || rule.WindowDaysMin != 3 || rule.WindowDaysMax != 5 || rule.MinHours != 6 {
+ t.Fatalf("unexpected watch-hours rule values: %+v", rule)
+ }
+}
+
+func TestFormatCleanupRulesShowsUsefulDetails(t *testing.T) {
+ text := formatCleanupRules([]accountCleanupRule{{
+ ID: "login_7d",
+ Type: "recent_login",
+ Name: "七天内登录",
+ Enabled: true,
+ WindowDaysMax: 7,
+ }})
+ for _, want := range []string{"保号规则", "login_7d", "七天内登录", "最近登录", "启用"} {
+ if !strings.Contains(text, want) {
+ t.Fatalf("formatCleanupRules() missing %q in %q", want, text)
+ }
+ }
+}
diff --git a/internal/service/telegram_commands.go b/internal/service/telegram_commands.go
index 21fe1f3..411d880 100644
--- a/internal/service/telegram_commands.go
+++ b/internal/service/telegram_commands.go
@@ -22,6 +22,17 @@ type telegramCommandDefinition struct {
func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage) []telegramCommandDefinition {
adminOnly := "此命令仅管理员可用。"
+ defs := s.telegramCoreCommandDefinitions(ctx, channel, msg)
+ defs = append(defs, s.telegramSelfServiceCommandDefinitions(ctx, channel, msg)...)
+ defs = append(defs, s.telegramAdminCoreCommandDefinitions(ctx, msg, adminOnly)...)
+ defs = append(defs, s.telegramMgoUserCommandDefinitions(ctx, adminOnly)...)
+ defs = append(defs, s.telegramMgoAuditCommandDefinitions(ctx, adminOnly)...)
+ defs = append(defs, s.telegramMgoMaintenanceCommandDefinitions(ctx, channel, adminOnly)...)
+ defs = append(defs, s.telegramMgoPolicyCommandDefinitions(ctx, channel, adminOnly)...)
+ return defs
+}
+
+func (s *TelegramBotService) telegramCoreCommandDefinitions(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage) []telegramCommandDefinition {
return []telegramCommandDefinition{
{Aliases: []string{"/start"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) {
if len(args) == 0 {
@@ -39,6 +50,11 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
{Aliases: []string{"/help"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) {
return telegramCommandReply{Text: s.cmdHelp(ctx, msg)}, nil
}},
+ }
+}
+
+func (s *TelegramBotService) telegramSelfServiceCommandDefinitions(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/hideadult", "/hide_adult", "/adult"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdHideAdult(ctx, msg, args), nil }},
{Aliases: []string{"/account", "/me", "/myinfo"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) { return s.replyAccount(ctx, msg), nil }},
{Aliases: []string{"/count"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdStats(ctx) }},
@@ -53,7 +69,11 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
}},
{Aliases: []string{"/redeem_renew"}, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdRedeemRenew(ctx, msg, args), nil }},
{Aliases: []string{"/register", "/reg", "/signup"}, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdRegister(ctx, channel, msg, args), nil }},
+ }
+}
+func (s *TelegramBotService) telegramAdminCoreCommandDefinitions(ctx context.Context, msg *TelegramMessage, adminOnly string) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/registration", "/reg_switch", "/openreg"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdRegistrationToggle(ctx, args), nil }},
{Aliases: []string{"/capacity"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.replyCapacity(ctx), nil }},
{Aliases: []string{"/users", "/kk"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.replyUserList(ctx), nil }},
@@ -75,11 +95,21 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
{Aliases: []string{"/downloads"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdDownloads(ctx) }},
{Aliases: []string{"/stats"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdStats(ctx) }},
{Aliases: []string{"/renew"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdUserRenew(ctx, args), nil }},
+ }
+}
+
+func (s *TelegramBotService) telegramMgoUserCommandDefinitions(ctx context.Context, adminOnly string) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/ucr"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoCreateUser(ctx, args), nil }},
{Aliases: []string{"/uinfo"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoUserInfo(ctx, args), nil }},
{Aliases: []string{"/rmemby", "/urm", "/only_rm_emby"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoDeleteUser(ctx, args), nil }},
{Aliases: []string{"/only_rm_record"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoOnlyRemoveRecord(ctx, args), nil }},
{Aliases: []string{"/userip"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoUserIP(ctx, args), nil }},
+ }
+}
+
+func (s *TelegramBotService) telegramMgoAuditCommandDefinitions(ctx context.Context, adminOnly string) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/udeviceid"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) {
return s.cmdMgoAuditDevices(ctx, "udeviceid", args), nil
}},
@@ -92,6 +122,11 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
{Aliases: []string{"/auditclient"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) {
return s.cmdMgoAuditDevices(ctx, "auditclient", args), nil
}},
+ }
+}
+
+func (s *TelegramBotService) telegramMgoMaintenanceCommandDefinitions(ctx context.Context, channel *model.NotifyChannel, adminOnly string) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/renewall"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoRenewAll(ctx, args), nil }},
{Aliases: []string{"/callall"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoCallAll(ctx, channel, args), nil }},
{Aliases: []string{"/syncunbound"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoSyncUnbound(ctx, args), nil }},
@@ -111,6 +146,11 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
{Aliases: []string{"/week_ranks"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) {
return s.cmdMgoRanks(ctx, 7*24*time.Hour, false), nil
}},
+ }
+}
+
+func (s *TelegramBotService) telegramMgoPolicyCommandDefinitions(ctx context.Context, channel *model.NotifyChannel, adminOnly string) []telegramCommandDefinition {
+ return []telegramCommandDefinition{
{Aliases: []string{"/embyadmin"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoAdminRole(ctx, args), nil }},
{Aliases: []string{"/unbanall"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoBanAll(ctx, true, args), nil }},
{Aliases: []string{"/banall"}, AdminOnly: true, AdminOnlyText: adminOnly, Handle: func(args []string) (telegramCommandReply, error) { return s.cmdMgoBanAll(ctx, false, args), nil }},
diff --git a/internal/service/telegram_device_policy.go b/internal/service/telegram_device_policy.go
new file mode 100644
index 0000000..79eeb57
--- /dev/null
+++ b/internal/service/telegram_device_policy.go
@@ -0,0 +1,252 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+)
+
+func (s *TelegramBotService) replyDevicePolicy(ctx context.Context) telegramCommandReply {
+ cfg := loadBotConfig(ctx, s.repo)
+ text := fmt.Sprintf(
+ "设备策略\n\n① 防共享:%s\n 并发播放终端上限 %d / 登录终端上限 %d;同一终端多个 App 只算 1 台,App 作为登录渠道记录。\n 设备指纹异常警告 %d 次后禁用账号。\n\n② Mgo 保号规则:%s\n 保号模式:%s;启用规则 %d 条。\n\n命令:\n/antishare on play=3 login=3 warn=2\n/cleanup run 预览候选\n/cleanup run confirm 确认清理\n/cleanup on|off\n/cleanup_rule list|add|edit|修改|del|enable|disable\n\n策略默认关闭;清理前会先预览候选;满足任意一条保号规则即保留;管理员/受保护账号永不自动处理。",
+ onOff(cfg.AntiShareEnabled), cfg.MaxConcurrentPlay, cfg.MaxLoggedClients, cfg.WarnThreshold,
+ onOff(cfg.AccountCleanupEnabled), cleanupModeLabel(cfg.AccountCleanupKeepMode), countEnabledCleanupRules(cfg.AccountCleanupRules))
+ return telegramCommandReply{
+ Text: text,
+ Buttons: [][]telegramInlineButton{
+ {{Text: toggleLabel("防共享", cfg.AntiShareEnabled), Data: "dp_toggle:antishare"}},
+ {{Text: toggleLabel("保号规则", cfg.AccountCleanupEnabled), Data: "dp_toggle:cleanup"}},
+ {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
+ },
+ }
+}
+
+func (s *TelegramBotService) cmdDevicePolicy(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 || strings.EqualFold(args[0], "status") {
+ return s.replyDevicePolicy(ctx)
+ }
+ switch strings.ToLower(strings.TrimSpace(args[0])) {
+ case "run", "sweep":
+ return s.cmdCleanup(ctx, []string{"run"})
+ default:
+ return telegramCommandReply{Text: "用法:/devicepolicy 查看策略,或使用 /antishare、/cleanup、/cleanup_rule 管理。"}
+ }
+}
+
+func (s *TelegramBotService) cmdAntiShare(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 || strings.EqualFold(args[0], "status") {
+ return s.replyDevicePolicy(ctx)
+ }
+ enabled, ok := parseCommandBool(args[0])
+ if !ok {
+ return telegramCommandReply{Text: "用法:/antishare on|off [play=3] [login=3] [warn=2],login 表示登录终端设备上限,同一终端多个 App 不重复计数。"}
+ }
+ if err := s.repo.Setting.Set(ctx, SettingAntiShareEnabled, strconv.FormatBool(enabled)); err != nil {
+ return telegramCommandReply{Text: "更新失败:" + err.Error()}
+ }
+ for _, arg := range args[1:] {
+ key, value, ok := strings.Cut(arg, "=")
+ if !ok {
+ continue
+ }
+ n, err := strconv.Atoi(strings.TrimSpace(value))
+ if err != nil || n < 1 {
+ continue
+ }
+ switch strings.ToLower(strings.TrimSpace(key)) {
+ case "play", "maxplay", "播放":
+ _ = s.repo.Setting.Set(ctx, SettingMaxConcurrentPlay, strconv.Itoa(n))
+ case "login", "client", "clients", "登录":
+ _ = s.repo.Setting.Set(ctx, SettingMaxLoggedClients, strconv.Itoa(n))
+ case "warn", "warnings", "警告":
+ _ = s.repo.Setting.Set(ctx, SettingWarnThreshold, strconv.Itoa(n))
+ }
+ }
+ return s.replyDevicePolicy(ctx)
+}
+
+func (s *TelegramBotService) cmdCleanup(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 || strings.EqualFold(args[0], "status") {
+ return s.replyDevicePolicy(ctx)
+ }
+ switch strings.ToLower(strings.TrimSpace(args[0])) {
+ case "on", "true", "1", "开启", "enable":
+ if err := s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
+ return telegramCommandReply{Text: "开启失败:" + err.Error()}
+ }
+ return s.replyDevicePolicy(ctx)
+ case "off", "false", "0", "关闭", "disable":
+ if err := s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, "false"); err != nil {
+ return telegramCommandReply{Text: "关闭失败:" + err.Error()}
+ }
+ return s.replyDevicePolicy(ctx)
+ case "run", "sweep", "巡检", "preview", "预览":
+ device := s.device
+ if device == nil {
+ device = NewDeviceService(s.log, s.repo)
+ }
+ if len(args) > 1 && isCleanupConfirmArg(args[1]) {
+ cfg := loadBotConfig(ctx, s.repo)
+ if !cfg.AccountCleanupEnabled {
+ return telegramCommandReply{Text: "保号规则未开启,不会清理账号。"}
+ }
+ if countEnabledCleanupRules(cfg.AccountCleanupRules) == 0 {
+ return telegramCommandReply{Text: "没有启用的保号规则,不会清理账号。"}
+ }
+ removed, err := device.SweepAccountCleanup(ctx)
+ if err != nil {
+ return telegramCommandReply{Text: "确认清理失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("保号规则确认清理完成,已清理 %d 个账号。", removed)}
+ }
+ candidates, err := device.PreviewAccountCleanup(ctx)
+ if err != nil {
+ return telegramCommandReply{Text: "巡检预览失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: s.formatCleanupPreview(ctx, candidates)}
+ default:
+ return telegramCommandReply{Text: "用法:/cleanup on|off、/cleanup run 预览、/cleanup run confirm 确认清理"}
+ }
+}
+
+func isCleanupConfirmArg(arg string) bool {
+ switch strings.ToLower(strings.TrimSpace(arg)) {
+ case "confirm", "yes", "delete", "确认", "清理", "删除":
+ return true
+ default:
+ return false
+ }
+}
+
+func (s *TelegramBotService) formatCleanupPreview(ctx context.Context, candidates []accountCleanupCandidate) string {
+ cfg := loadBotConfig(ctx, s.repo)
+ if !cfg.AccountCleanupEnabled {
+ return "保号规则未开启,不会清理账号。"
+ }
+ if countEnabledCleanupRules(cfg.AccountCleanupRules) == 0 {
+ return "没有启用的保号规则,不会清理账号。"
+ }
+ if len(candidates) == 0 {
+ return "保号规则预览完成:没有需要清理的账号。"
+ }
+ var sb strings.Builder
+ sb.WriteString(fmt.Sprintf("保号规则预览\n\n将清理候选:%d 个账号。\n当前只是预览,未删除任何账号。\n\n", len(candidates)))
+ limit := len(candidates)
+ if limit > 10 {
+ limit = 10
+ }
+ for i := 0; i < limit; i++ {
+ candidate := candidates[i]
+ sb.WriteString(fmt.Sprintf("%d. %s\n%s\n", i+1, escapeHTML(candidate.Username), escapeHTML(candidate.Details)))
+ }
+ if len(candidates) > limit {
+ sb.WriteString(fmt.Sprintf("……另有 %d 个候选未展示。\n", len(candidates)-limit))
+ }
+ sb.WriteString("\n确认无误后再执行:/cleanup run confirm")
+ return sb.String()
+}
+
+func (s *TelegramBotService) cmdCleanupMode(ctx context.Context, args []string) telegramCommandReply {
+ if err := s.repo.Setting.Set(ctx, SettingAccountCleanupKeepMode, "any"); err != nil {
+ return telegramCommandReply{Text: "更新失败:" + err.Error()}
+ }
+ if err := s.repo.Setting.Set(ctx, SettingAccountCleanupRequiredCount, "1"); err != nil {
+ return telegramCommandReply{Text: "更新失败:" + err.Error()}
+ }
+ reply := s.replyDevicePolicy(ctx)
+ reply.Text = "Mgo 保号模式固定为:满足任意一条启用规则即保留;只有全部规则都不满足才进入清理候选。\n\n" + reply.Text
+ return reply
+}
+
+func (s *TelegramBotService) cmdCleanupRule(ctx context.Context, args []string) telegramCommandReply {
+ rules := s.currentCleanupRules(ctx)
+ if len(args) == 0 {
+ return telegramCommandReply{Text: formatCleanupRules(rules)}
+ }
+ action := strings.ToLower(strings.TrimSpace(args[0]))
+ switch action {
+ case "list", "ls", "status":
+ return telegramCommandReply{Text: formatCleanupRules(rules)}
+ case "help", "?", "帮助":
+ return telegramCommandReply{Text: cleanupRuleHelp()}
+ case "del", "delete", "rm":
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "用法:/cleanup_rule del 规则ID"}
+ }
+ next := make([]accountCleanupRule, 0, len(rules))
+ removed := false
+ for _, r := range rules {
+ if r.ID == args[1] {
+ removed = true
+ continue
+ }
+ next = append(next, r)
+ }
+ if !removed {
+ return telegramCommandReply{Text: "未找到该规则。"}
+ }
+ if err := s.saveCleanupRules(ctx, next); err != nil {
+ return telegramCommandReply{Text: "保存失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: "已删除规则。\n\n" + formatCleanupRules(next)}
+ case "enable", "on", "disable", "off":
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "用法:/cleanup_rule enable|disable 规则ID"}
+ }
+ enable := action == "enable" || action == "on"
+ changed := false
+ for i := range rules {
+ if rules[i].ID == args[1] {
+ rules[i].Enabled = enable
+ changed = true
+ }
+ }
+ if !changed {
+ return telegramCommandReply{Text: "未找到该规则。"}
+ }
+ if err := s.saveCleanupRules(ctx, rules); err != nil {
+ return telegramCommandReply{Text: "保存失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: "已更新规则状态。\n\n" + formatCleanupRules(rules)}
+ case "add", "set", "edit", "update", "修改", "更新", "改":
+ rule, err := parseCleanupRuleCommand(args[1:])
+ if err != nil {
+ return telegramCommandReply{Text: err.Error() + "\n\n" + cleanupRuleHelp()}
+ }
+ updated := false
+ for i := range rules {
+ if rules[i].ID == rule.ID {
+ rules[i] = rule
+ updated = true
+ break
+ }
+ }
+ if !updated {
+ rules = append(rules, rule)
+ }
+ rules = normalizeCleanupRules(rules)
+ if err := s.saveCleanupRules(ctx, rules); err != nil {
+ return telegramCommandReply{Text: "保存失败:" + err.Error()}
+ }
+ actionText := "已新增规则。"
+ if updated {
+ actionText = "已更新规则。"
+ }
+ return telegramCommandReply{Text: actionText + "\n\n" + formatCleanupRules(rules)}
+ default:
+ return telegramCommandReply{Text: cleanupRuleHelp()}
+ }
+}
+
+func (s *TelegramBotService) replyDevicePolicyToggle(ctx context.Context, which string) telegramCommandReply {
+ cfg := loadBotConfig(ctx, s.repo)
+ switch which {
+ case "antishare":
+ _ = s.repo.Setting.Set(ctx, SettingAntiShareEnabled, strconv.FormatBool(!cfg.AntiShareEnabled))
+ case "cleanup":
+ _ = s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, strconv.FormatBool(!cfg.AccountCleanupEnabled))
+ }
+ return s.replyDevicePolicy(ctx)
+}
diff --git a/internal/service/telegram_menu.go b/internal/service/telegram_menu.go
index 0a254f8..064d543 100644
--- a/internal/service/telegram_menu.go
+++ b/internal/service/telegram_menu.go
@@ -2,25 +2,17 @@ package service
import (
"context"
- "encoding/json"
- "errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
- "gorm.io/gorm"
)
// pendingTTL bounds how long a button-initiated text prompt stays valid.
const pendingTTL = 5 * time.Minute
-var (
- errRegistrationCodeAlreadyUsed = errors.New("registration code already used")
- errRegistrationCodeExpired = errors.New("registration code expired")
-)
-
func (s *TelegramBotService) setPending(userID int64, kind string) {
s.pendingMu.Lock()
s.pending[userID] = pendingInput{Kind: kind, CreatedAt: time.Now()}
@@ -50,121 +42,45 @@ func (s *TelegramBotService) boundUser(ctx context.Context, telegramUserID int)
return u
}
-// mainMenu builds the button-based menu, tailored to the user's binding and
-// admin status. Ordinary users only see self-service actions; admins get an
-// extra management section.
-func (s *TelegramBotService) mainMenu(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage) telegramCommandReply {
- isAdmin := s.telegramUserIsAdmin(ctx, channel, msg.From.ID)
- isGroup := telegramIsGroupChat(msg.Chat.Type)
- user := s.boundUser(ctx, msg.From.ID)
-
- var rows [][]telegramInlineButton
- var header string
-
- if isGroup {
- if user == nil {
- header = "MediaStationGo 群组自助菜单\n\n你还没有绑定媒体中心账号。绑定、注册、兑换等包含敏感信息的操作请私聊 Bot。"
- } else {
- adult := map[bool]string{true: "已隐藏", false: "已显示"}[user.HideAdult]
- header = fmt.Sprintf("MediaStationGo 群组自助菜单\n\n账号:%s\n到期:%s\n成人目录:%s",
- user.Username, formatExpiry(user.ExpiredAt), adult)
- rows = append(rows,
- []telegramInlineButton{
- {Text: "👤 我的账号", Data: "act_account"},
- {Text: "📅 签到", Data: "act_signin"},
- },
- []telegramInlineButton{
- {Text: "📱 我的设备", Data: "act_devices"},
- {Text: map[bool]string{true: "🔞 显示成人目录", false: "🔞 隐藏成人目录"}[user.HideAdult], Data: "adult_toggle"},
- },
- )
- }
- if isAdmin {
- header += "\n\n管理员入口"
- rows = append(rows,
- []telegramInlineButton{{Text: "—— 管理员 ——", Data: "noop"}},
- []telegramInlineButton{
- {Text: "📊 容量/状态", Data: "adm_capacity"},
- {Text: "👥 用户管理", Data: "adm_users"},
- },
- []telegramInlineButton{
- {Text: "🔓 开注设置", Data: "adm_openreg"},
- {Text: "🎟 生成兑换码", Data: "adm_gencode"},
- },
- []telegramInlineButton{
- {Text: "⚙️ 设备策略", Data: "adm_devicepolicy"},
- {Text: "🛠 管理命令", Data: "adm_mgo_commands"},
- },
- )
- }
- return telegramCommandReply{Text: header, Buttons: rows}
- }
-
- if user == nil {
- header = "MediaStationGo\n\n你还没有绑定媒体中心账号。"
- rows = append(rows, []telegramInlineButton{{Text: "🔗 绑定账号", Data: "act_bind"}})
- if s.openRegEnabled(ctx) {
- rows = append(rows, []telegramInlineButton{{Text: "📝 注册新账号", Data: "act_register"}})
- }
- rows = append(rows, []telegramInlineButton{{Text: "🎟 兑换码注册", Data: "act_redeem_register"}})
- } else {
- adult := map[bool]string{true: "已隐藏", false: "已显示"}[user.HideAdult]
- header = fmt.Sprintf("MediaStationGo\n\n账号:%s\n到期:%s\n成人目录:%s",
- user.Username, formatExpiry(user.ExpiredAt), adult)
- rows = append(rows,
- []telegramInlineButton{
- {Text: "👤 我的账号", Data: "act_account"},
- {Text: "📅 签到", Data: "act_signin"},
- },
- []telegramInlineButton{
- {Text: "📱 我的设备", Data: "act_devices"},
- {Text: map[bool]string{true: "🔞 显示成人目录", false: "🔞 隐藏成人目录"}[user.HideAdult], Data: "adult_toggle"},
- },
- []telegramInlineButton{
- {Text: "✏️ 改用户名", Data: "act_setname"},
- {Text: "🔑 改密码", Data: "act_setpass"},
- },
- []telegramInlineButton{{Text: "🎟 兑换码续期", Data: "act_redeem_renew"}},
- )
- }
-
- if isAdmin {
- rows = append(rows,
- []telegramInlineButton{{Text: "—— 管理员 ——", Data: "noop"}},
- []telegramInlineButton{
- {Text: "📊 容量/状态", Data: "adm_capacity"},
- {Text: "👥 用户管理", Data: "adm_users"},
- },
- []telegramInlineButton{
- {Text: "🔓 开注设置", Data: "adm_openreg"},
- {Text: "🎟 生成兑换码", Data: "adm_gencode"},
- },
- []telegramInlineButton{
- {Text: "⚙️ 设备策略", Data: "adm_devicepolicy"},
- {Text: "🛠 管理命令", Data: "adm_mgo_commands"},
- },
- )
- }
-
- return telegramCommandReply{Text: header, Buttons: rows}
-}
-
// handleMenuCallback routes inline-button taps. Returns (reply, handled).
func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, data string) (telegramCommandReply, bool) {
isAdmin := s.telegramUserIsAdmin(ctx, channel, msg.From.ID)
isGroup := telegramIsGroupChat(msg.Chat.Type)
+ if reply, handled := s.handleUserMenuCallback(ctx, channel, msg, data, isGroup); handled {
+ return reply, true
+ }
+ if !isAdmin {
+ if isGroup {
+ return telegramCommandReply{}, true
+ }
+ return telegramCommandReply{Text: "此功能仅管理员可用。"}, true
+ }
+ return s.handleAdminMenuCallback(ctx, msg, data)
+}
+func (s *TelegramBotService) handleUserMenuCallback(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, data string, isGroup bool) (telegramCommandReply, bool) {
switch {
case data == "noop":
return telegramCommandReply{}, true
case data == "menu_main":
return s.mainMenu(ctx, channel, msg), true
- case data == "act_bind":
- if isGroup {
- return telegramCommandReply{Text: telegramGroupPrivateUserHint("绑定账号")}, true
- }
- return telegramCommandReply{Text: "请发送:/start 用户名 密码 绑定已有账号。"}, true
- case data == "act_register":
+ case data == "act_account":
+ return s.replyAccount(ctx, msg), true
+ case data == "act_signin":
+ return s.replySignIn(ctx, msg), true
+ case data == "act_devices":
+ return s.replyDevices(ctx, msg), true
+ case strings.HasPrefix(data, "kick:"):
+ return s.replyKick(ctx, msg, strings.TrimPrefix(data, "kick:")), true
+ }
+ return s.handlePrivatePromptMenuCallback(ctx, msg, data, isGroup)
+}
+
+func (s *TelegramBotService) handlePrivatePromptMenuCallback(ctx context.Context, msg *TelegramMessage, data string, isGroup bool) (telegramCommandReply, bool) {
+ switch data {
+ case "act_bind":
+ return telegramPrivateOnlyMenuReply(isGroup, "绑定账号", "请发送:/start 用户名 密码 绑定已有账号。"), true
+ case "act_register":
if isGroup {
return telegramCommandReply{Text: telegramGroupPrivateUserHint("注册账号")}, true
}
@@ -173,50 +89,63 @@ func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *mo
}
s.setPending(int64(msg.From.ID), "register")
return telegramCommandReply{Text: "请发送新账号的 用户名 密码(空格分隔),例如:alice mypass123"}, true
- case data == "act_redeem_register":
+ case "act_redeem_register":
if isGroup {
return telegramCommandReply{Text: telegramGroupPrivateUserHint("兑换码注册")}, true
}
s.setPending(int64(msg.From.ID), "redeem_register")
return telegramCommandReply{Text: "请发送你的注册兑换码,例如:ABCD2345EFGH\n(兑换后会要求设置用户名密码)"}, true
- case data == "act_redeem_renew":
+ case "act_redeem_renew":
if isGroup {
return telegramCommandReply{Text: telegramGroupPrivateUserHint("兑换码续期")}, true
}
s.setPending(int64(msg.From.ID), "redeem_renew")
return telegramCommandReply{Text: "请发送你的续期兑换码,将为当前绑定账号续期。"}, true
- case data == "act_account":
- return s.replyAccount(ctx, msg), true
- case data == "act_signin":
- return s.replySignIn(ctx, msg), true
- case data == "act_devices":
- return s.replyDevices(ctx, msg), true
- case data == "act_setname":
- if isGroup {
- return telegramCommandReply{Text: telegramGroupPrivateUserHint("修改用户名")}, true
- }
- s.setPending(int64(msg.From.ID), "setname")
- return telegramCommandReply{Text: "请发送:当前密码 新用户名。"}, true
- case data == "act_setpass":
- if isGroup {
- return telegramCommandReply{Text: telegramGroupPrivateUserHint("修改密码")}, true
- }
- s.setPending(int64(msg.From.ID), "setpass")
- return telegramCommandReply{Text: "请发送:当前密码 新密码(新密码至少 6 位)。"}, true
- case strings.HasPrefix(data, "kick:"):
- return s.replyKick(ctx, msg, strings.TrimPrefix(data, "kick:")), true
+ case "act_setname":
+ return s.setPendingPrivatePrompt(msg, isGroup, "修改用户名", "setname", "请发送:当前密码 新用户名。"), true
+ case "act_setpass":
+ return s.setPendingPrivatePrompt(msg, isGroup, "修改密码", "setpass", "请发送:当前密码 新密码(新密码至少 6 位)。"), true
}
+ return telegramCommandReply{}, false
+}
- // ── 管理员专属 ──
- if isGroup && !isAdmin {
- return telegramCommandReply{}, true
+func telegramPrivateOnlyMenuReply(isGroup bool, action, privateText string) telegramCommandReply {
+ if isGroup {
+ return telegramCommandReply{Text: telegramGroupPrivateUserHint(action)}
}
- if !isAdmin {
- return telegramCommandReply{Text: "此功能仅管理员可用。"}, true
+ return telegramCommandReply{Text: privateText}
+}
+
+func (s *TelegramBotService) setPendingPrivatePrompt(msg *TelegramMessage, isGroup bool, action, kind, text string) telegramCommandReply {
+ if isGroup {
+ return telegramCommandReply{Text: telegramGroupPrivateUserHint(action)}
+ }
+ s.setPending(int64(msg.From.ID), kind)
+ return telegramCommandReply{Text: text}
+}
+
+func (s *TelegramBotService) handleAdminMenuCallback(ctx context.Context, msg *TelegramMessage, data string) (telegramCommandReply, bool) {
+ if reply, handled := s.handleAdminRegistrationCallback(ctx, msg, data); handled {
+ return reply, true
+ }
+ if reply, handled := s.handleAdminUserCallback(ctx, data); handled {
+ return reply, true
}
switch {
case data == "adm_capacity":
return s.replyCapacity(ctx), true
+ case data == "adm_devicepolicy":
+ return s.replyDevicePolicy(ctx), true
+ case data == "adm_mgo_commands":
+ return telegramCommandReply{Text: telegramMgoAdminCommandHelp(), Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}, true
+ case strings.HasPrefix(data, "dp_toggle:"):
+ return s.replyDevicePolicyToggle(ctx, strings.TrimPrefix(data, "dp_toggle:")), true
+ }
+ return telegramCommandReply{}, false
+}
+
+func (s *TelegramBotService) handleAdminRegistrationCallback(ctx context.Context, msg *TelegramMessage, data string) (telegramCommandReply, bool) {
+ switch {
case data == "adm_openreg":
return s.replyOpenRegMenu(ctx), true
case data == "adm_openreg_close":
@@ -236,6 +165,12 @@ func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *mo
return s.replyGenCodeMenu(), true
case strings.HasPrefix(data, "gc:"):
return s.replyGenCode(ctx, msg, data), true
+ }
+ return telegramCommandReply{}, false
+}
+
+func (s *TelegramBotService) handleAdminUserCallback(ctx context.Context, data string) (telegramCommandReply, bool) {
+ switch {
case data == "adm_users":
return s.replyUserList(ctx), true
case strings.HasPrefix(data, "usr:"):
@@ -248,12 +183,6 @@ func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *mo
return s.replyUserDelete(ctx, strings.TrimPrefix(data, "udel:")), true
case strings.HasPrefix(data, "urenew:"):
return s.replyUserRenew(ctx, strings.TrimPrefix(data, "urenew:")), true
- case data == "adm_devicepolicy":
- return s.replyDevicePolicy(ctx), true
- case data == "adm_mgo_commands":
- return telegramCommandReply{Text: telegramMgoAdminCommandHelp(), Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}, true
- case strings.HasPrefix(data, "dp_toggle:"):
- return s.replyDevicePolicyToggle(ctx, strings.TrimPrefix(data, "dp_toggle:")), true
}
return telegramCommandReply{}, false
}
@@ -288,1343 +217,3 @@ func (s *TelegramBotService) handlePendingText(ctx context.Context, channel *mod
}
return telegramCommandReply{}, false
}
-
-// ── 用户自助 ──────────────────────────────────────────────────────────────
-
-func (s *TelegramBotService) cmdKick(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号:/start 用户名 密码"}
- }
- if len(args) == 0 {
- return telegramCommandReply{Text: "请指定要踢下线的设备:/kick all 或 /kick 设备编号。先用 /devices 查看编号。"}
- }
- target := strings.TrimSpace(args[0])
- if strings.EqualFold(target, "all") || target == "全部" {
- if s.device != nil {
- if err := s.device.KickAllDevices(ctx, user.ID); err != nil {
- return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
- }
- } else if err := s.repo.UserDevice.SetKickedByUser(ctx, user.ID, true); err != nil {
- return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
- }
- return telegramCommandReply{Text: "已踢下线此账号的全部设备。"}
- }
- devices, _ := s.repo.UserDevice.ListByUser(ctx, user.ID)
- if len(devices) == 0 {
- return telegramCommandReply{Text: "当前没有记录到登录设备。"}
- }
- var chosen *model.UserDevice
- if n, err := strconv.Atoi(target); err == nil && n >= 1 && n <= len(devices) {
- chosen = &devices[n-1]
- } else {
- for i := range devices {
- if devices[i].ID == target || devices[i].DeviceID == target {
- chosen = &devices[i]
- break
- }
- }
- }
- if chosen == nil {
- return telegramCommandReply{Text: "未找到该设备。请用 /devices 查看设备编号后重试。"}
- }
- if err := s.repo.UserDevice.SetKicked(ctx, chosen.ID, true); err != nil {
- return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
- }
- return telegramCommandReply{Text: fmt.Sprintf("已踢下线:%s。", deviceLabel(chosen.DeviceName, chosen.Client))}
-}
-
-func (s *TelegramBotService) cmdSetName(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) < 2 {
- return telegramCommandReply{Text: "请发送:/setname 当前密码 新用户名"}
- }
- return s.selfSetName(ctx, msg, strings.Join(args, " "))
-}
-
-func (s *TelegramBotService) cmdSetPass(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) < 2 {
- return telegramCommandReply{Text: "请发送:/setpass 当前密码 新密码"}
- }
- return s.selfSetPass(ctx, msg, strings.Join(args, " "))
-}
-
-func (s *TelegramBotService) cmdRedeem(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "请发送:/redeem 兑换码\n未绑定账号时自动尝试注册码;已绑定账号时自动尝试续期码。"}
- }
- code := strings.Join(args, " ")
- if s.boundUser(ctx, msg.From.ID) == nil {
- return s.redeemRegisterFlow(ctx, channel, msg, code)
- }
- return s.redeemRenewFlow(ctx, msg, code)
-}
-
-func (s *TelegramBotService) cmdRedeemRegister(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "请发送:/redeem_register 注册兑换码"}
- }
- return s.redeemRegisterFlow(ctx, channel, msg, strings.Join(args, " "))
-}
-
-func (s *TelegramBotService) cmdRedeemRenew(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "请发送:/redeem_renew 续期兑换码"}
- }
- return s.redeemRenewFlow(ctx, msg, strings.Join(args, " "))
-}
-
-func (s *TelegramBotService) replyAccount(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号:/start 用户名 密码"}
- }
- streak := 0
- if rec, _ := s.repo.SignIn.Get(ctx, user.ID); rec != nil {
- streak = rec.StreakDays
- }
- devices, _ := s.repo.UserDevice.ListByUser(ctx, user.ID)
- text := fmt.Sprintf("我的账号\n\n用户名:%s\n状态:%s\n到期:%s\n连续签到:%d 天\n登录设备:%d 台",
- user.Username,
- map[bool]string{true: "正常", false: "已禁用"}[user.IsActive],
- formatExpiry(user.ExpiredAt), streak, len(devices))
- return telegramCommandReply{Text: text, Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}
-}
-
-func (s *TelegramBotService) replySignIn(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号后再签到。"}
- }
- res, err := s.signIn(ctx, user.ID)
- if err != nil {
- return telegramCommandReply{Text: "签到失败:" + err.Error()}
- }
- if res.AlreadySigned {
- return telegramCommandReply{Text: fmt.Sprintf("今天已经签到过啦~\n连续签到 %d 天,累计 %d 天。", res.Streak, res.Total)}
- }
- return telegramCommandReply{Text: fmt.Sprintf("签到成功 ✅\n连续签到 %d 天,累计 %d 天。", res.Streak, res.Total)}
-}
-
-func (s *TelegramBotService) replyDevices(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号。"}
- }
- devices, _ := s.repo.UserDevice.ListByUser(ctx, user.ID)
- if len(devices) == 0 {
- return telegramCommandReply{Text: "当前没有记录到登录设备。"}
- }
- var sb strings.Builder
- sb.WriteString("我的登录设备\n点击下方按钮可一键踢下线:\n")
- var rows [][]telegramInlineButton
- for i, d := range devices {
- status := ""
- if d.Kicked {
- status = "(已踢下线)"
- }
- sb.WriteString(fmt.Sprintf("\n%d. %s%s\n 最近活跃:%s", i+1, deviceLabel(d.DeviceName, d.Client), status, d.LastSeenAt.Format("01-02 15:04")))
- if !d.Kicked {
- rows = append(rows, []telegramInlineButton{{Text: "🚫 踢下线:" + deviceLabel(d.DeviceName, d.Client), Data: "kick:" + d.ID}})
- }
- }
- rows = append(rows, []telegramInlineButton{{Text: "⬅️ 返回菜单", Data: "menu_main"}})
- return telegramCommandReply{Text: sb.String(), Buttons: rows}
-}
-
-func (s *TelegramBotService) replyKick(ctx context.Context, msg *TelegramMessage, deviceRowID string) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号。"}
- }
- // Verify the device belongs to this user before kicking.
- var d model.UserDevice
- if err := s.repo.DB.WithContext(ctx).Where("id = ? AND user_id = ?", deviceRowID, user.ID).First(&d).Error; err != nil {
- return telegramCommandReply{Text: "未找到该设备。"}
- }
- if err := s.repo.UserDevice.SetKicked(ctx, d.ID, true); err != nil {
- return telegramCommandReply{Text: "操作失败:" + err.Error()}
- }
- return s.replyDevices(ctx, msg)
-}
-
-func (s *TelegramBotService) selfSetName(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号。"}
- }
- currentPassword, newName := splitCurrentPasswordAndValue(input)
- if currentPassword == "" || newName == "" {
- return telegramCommandReply{Text: "请发送:当前密码 新用户名。"}
- }
- newName = strings.TrimSpace(newName)
- if len(newName) < 2 || strings.ContainsAny(newName, " \t\n") {
- return telegramCommandReply{Text: "用户名至少 2 位且不能含空格,请重试。"}
- }
- if reply, ok := s.verifyTelegramSelfPassword(ctx, msg, user, currentPassword); !ok {
- return reply
- }
- if existing, _ := s.repo.User.FindByUsername(ctx, newName); existing != nil && existing.ID != user.ID {
- return telegramCommandReply{Text: "该用户名已被占用,请换一个。"}
- }
- if err := s.repo.User.UpdateFields(ctx, user.ID, map[string]any{"username": newName}); err != nil {
- return telegramCommandReply{Text: "修改失败:" + err.Error()}
- }
- return telegramCommandReply{Text: fmt.Sprintf("用户名已修改为 %s。请用新用户名登录。", newName)}
-}
-
-func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号。"}
- }
- currentPassword, newPass := splitCurrentPasswordAndValue(input)
- if currentPassword == "" || newPass == "" {
- return telegramCommandReply{Text: "请发送:当前密码 新密码。"}
- }
- newPass = strings.TrimSpace(newPass)
- if s.auth == nil {
- return telegramCommandReply{Text: "服务暂不可用。"}
- }
- if err := s.auth.ChangePassword(ctx, user.ID, currentPassword, newPass); err != nil {
- if errors.Is(err, ErrInvalidCredentials) {
- _ = s.unbindTelegramUser(ctx, msg.From.ID)
- return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"}
- }
- return telegramCommandReply{Text: "修改失败:" + err.Error()}
- }
- if s.device != nil {
- _ = s.device.KickAllDevices(ctx, user.ID)
- }
- return telegramCommandReply{Text: "密码已修改,请用新密码重新登录第三方客户端。"}
-}
-
-func splitCurrentPasswordAndValue(input string) (string, string) {
- fields := strings.Fields(strings.TrimSpace(input))
- if len(fields) < 2 {
- return "", ""
- }
- return fields[0], strings.TrimSpace(strings.Join(fields[1:], " "))
-}
-
-func (s *TelegramBotService) verifyTelegramSelfPassword(ctx context.Context, msg *TelegramMessage, user *model.User, currentPassword string) (telegramCommandReply, bool) {
- if s.auth == nil {
- return telegramCommandReply{Text: "服务暂不可用。"}, false
- }
- if err := s.auth.VerifyPassword(ctx, user.ID, currentPassword); err != nil {
- if errors.Is(err, ErrInvalidCredentials) {
- _ = s.unbindTelegramUser(ctx, msg.From.ID)
- return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"}, false
- }
- return telegramCommandReply{Text: "验证失败:" + err.Error()}, false
- }
- return telegramCommandReply{}, true
-}
-
-// ── 兑换码流程 ───────────────────────────────────────────────────────────────
-
-func (s *TelegramBotService) redeemRegisterFlow(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, raw string) telegramCommandReply {
- if channel == nil {
- channel = s.findChannelForMessage(ctx, msg)
- }
- if dec := s.telegramUserBindDecision(ctx, channel, msg.From.ID); dec != bindAllowed {
- return telegramCommandReply{Text: telegramBindRejectText(dec, "兑换注册账号")}
- }
- rc, errMsg := s.lookupRedeemableCode(ctx, raw, model.RegistrationCodeRegister)
- if rc == nil {
- return telegramCommandReply{Text: errMsg}
- }
- if s.auth == nil {
- return telegramCommandReply{Text: "注册服务暂不可用。"}
- }
- if binding := s.telegramBinding(ctx, msg.From.ID); binding != nil {
- if u, _ := s.repo.User.FindByID(ctx, binding.UserID); u != nil {
- return telegramCommandReply{Text: fmt.Sprintf("当前 Telegram 已绑定账号 %s,无需再用注册码。", u.Username)}
- }
- }
- user, password, claimedCode, err := s.createUserFromRegistrationCode(ctx, rc.Code)
- if err != nil {
- if errors.Is(err, errRegistrationCodeAlreadyUsed) {
- return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
- }
- if errors.Is(err, errRegistrationCodeExpired) {
- return telegramCommandReply{Text: "兑换码已过期。"}
- }
- if errors.Is(err, ErrUserLimitReached) {
- return telegramCommandReply{Text: "注册失败:用户数量已达授权上限。"}
- }
- return telegramCommandReply{Text: "注册失败:" + err.Error()}
- }
- if claimedCode == nil {
- return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
- }
- _ = s.upsertTelegramBinding(ctx, msg, user.ID)
- return telegramCommandReply{
- Text: fmt.Sprintf("兑换成功并已创建账号:\n用户名:%s\n密码:%s\n到期:%s\n\n请尽快用「改用户名/改密码」修改为你自己的凭据。",
- user.Username, password, formatExpiry(s.userExpiry(ctx, user.ID))),
- Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}},
- }
-}
-
-func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context, rawCode string) (*model.User, string, *model.RegistrationCode, error) {
- code := normalizeRedemptionCode(rawCode)
- if code == "" {
- return nil, "", nil, errRegistrationCodeAlreadyUsed
- }
- password := randomCode(10)
- var created model.User
- var claimed model.RegistrationCode
- err := s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
- if err := tx.Where("code = ? AND kind = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", code, model.RegistrationCodeRegister).
- First(&claimed).Error; err != nil {
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return errRegistrationCodeAlreadyUsed
- }
- return err
- }
- if claimed.IsExpired() {
- return errRegistrationCodeExpired
- }
- var count int64
- if err := tx.Model(&model.User{}).Count(&count).Error; err != nil {
- return err
- }
- if count >= LicensedMaxUsers(ctx, s.repo) {
- return ErrUserLimitReached
- }
- hash, err := hashPassword(password)
- if err != nil {
- return err
- }
- codePrefix := strings.ToLower(claimed.Code)
- if len(codePrefix) > 8 {
- codePrefix = codePrefix[:8]
- }
- created = model.User{
- Username: "u" + codePrefix,
- PasswordHash: hash,
- Role: "user",
- Tier: "free",
- HideAdult: true,
- ExpiredAt: renewExpiry(nil, claimed.DurationDays),
- }
- if err := tx.Create(&created).Error; err != nil {
- return err
- }
- if err := tx.Create(DefaultPermissions(created.ID)).Error; err != nil {
- return err
- }
- now := time.Now()
- res := tx.Model(&model.RegistrationCode{}).
- Where("id = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", claimed.ID).
- Updates(map[string]any{
- "used_by_user_id": created.ID,
- "used_count": gorm.Expr("used_count + 1"),
- "used_at": gorm.Expr("CASE WHEN used_count + 1 >= CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END THEN ? ELSE used_at END", now),
- })
- if res.Error != nil {
- return res.Error
- }
- if res.RowsAffected == 0 {
- return errRegistrationCodeAlreadyUsed
- }
- claimed.UsedByUserID = created.ID
- claimed.UsedCount++
- if claimed.UsedCount >= claimed.EffectiveMaxUses() {
- claimed.UsedAt = &now
- }
- return nil
- })
- if err != nil {
- return nil, "", nil, err
- }
- return &created, password, &claimed, nil
-}
-
-func (s *TelegramBotService) redeemRenewFlow(ctx context.Context, msg *TelegramMessage, raw string) telegramCommandReply {
- user := s.boundUser(ctx, msg.From.ID)
- if user == nil {
- return telegramCommandReply{Text: "请先绑定账号再续期。"}
- }
- rc, errMsg := s.lookupRedeemableCode(ctx, raw, model.RegistrationCodeRenew)
- if rc == nil {
- return telegramCommandReply{Text: errMsg}
- }
- if err := s.repo.RegCode.MarkUsed(ctx, rc.ID, user.ID); err != nil {
- return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
- }
- if err := s.applyRenewal(ctx, user.ID, rc.DurationDays); err != nil {
- return telegramCommandReply{Text: "续期失败:" + err.Error()}
- }
- return telegramCommandReply{Text: fmt.Sprintf("续期成功 ✅ 当前到期:%s", formatExpiry(s.userExpiry(ctx, user.ID)))}
-}
-
-func (s *TelegramBotService) userExpiry(ctx context.Context, userID string) *time.Time {
- if u, _ := s.repo.User.FindByID(ctx, userID); u != nil {
- return u.ExpiredAt
- }
- return nil
-}
-
-// ── 管理员:容量 / 开注 / 兑换码 / 用户管理 / 设备策略 ─────────────────────────
-
-func (s *TelegramBotService) replyCapacity(ctx context.Context) telegramCommandReply {
- c := s.loadCapacity(ctx)
- quota := "未开放"
- if c.OpenRegOn {
- if c.OpenRegLimit > 0 {
- quota = fmt.Sprintf("已开放(%d/%d 名额)", c.OpenRegUsed, c.OpenRegLimit)
- } else {
- quota = "已开放(不限名额,受授权上限约束)"
- }
- }
- text := fmt.Sprintf("容量 / 状态\n\n授权上限:%d 人(随凭证授权实时变化)\n已用:%d 人\n剩余可注册:%d 人\n开注状态:%s",
- c.MaxUsers, c.UsedUsers, c.Remaining(), quota)
- return telegramCommandReply{Text: text, Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}
-}
-
-func (s *TelegramBotService) replyOpenRegMenu(ctx context.Context) telegramCommandReply {
- c := s.loadCapacity(ctx)
- state := "未开放"
- if c.OpenRegOn {
- state = fmt.Sprintf("已开放(%d/%d)", c.OpenRegUsed, c.OpenRegLimit)
- }
- return telegramCommandReply{
- Text: "开注设置\n当前:" + state + "\n选择要开放的名额:",
- Buttons: [][]telegramInlineButton{
- {{Text: "5 个", Data: "adm_openreg_set:5"}, {Text: "10 个", Data: "adm_openreg_set:10"}, {Text: "20 个", Data: "adm_openreg_set:20"}},
- {{Text: "不限名额", Data: "adm_openreg_set:0"}, {Text: "关闭注册", Data: "adm_openreg_close"}},
- {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
- },
- }
-}
-
-func (s *TelegramBotService) replyGenCodeMenu() telegramCommandReply {
- return telegramCommandReply{
- Text: "生成兑换码\n选择类型与时长:",
- Buttons: [][]telegramInlineButton{
- {{Text: "注册码·30天", Data: "gc:register:30"}, {Text: "注册码·永久", Data: "gc:register:0"}},
- {{Text: "续期码·30天", Data: "gc:renew:30"}, {Text: "续期码·90天", Data: "gc:renew:90"}},
- {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
- },
- }
-}
-
-func (s *TelegramBotService) replyGenCode(ctx context.Context, msg *TelegramMessage, data string) telegramCommandReply {
- parts := strings.Split(data, ":") // gc::
- if len(parts) != 3 {
- return telegramCommandReply{Text: "参数错误。"}
- }
- kind := parts[1]
- days, _ := strconv.Atoi(parts[2])
- createdBy := ""
- if u := s.boundUser(ctx, msg.From.ID); u != nil {
- createdBy = u.ID
- }
- code, err := s.generateCode(ctx, kind, days, 0, createdBy)
- if err != nil {
- return telegramCommandReply{Text: "生成失败:" + err.Error()}
- }
- kindLabel := map[string]string{model.RegistrationCodeRegister: "注册码", model.RegistrationCodeRenew: "续期码"}[code.Kind]
- dur := "永久"
- if days > 0 {
- dur = fmt.Sprintf("%d 天", days)
- }
- return telegramCommandReply{
- Text: fmt.Sprintf("已生成%s(%s):\n\n%s\n\n发给用户在 Bot 中兑换即可。", kindLabel, dur, code.Code),
- Buttons: [][]telegramInlineButton{{{Text: "再生成一个", Data: "adm_gencode"}, {Text: "⬅️ 返回菜单", Data: "menu_main"}}},
- }
-}
-
-func (s *TelegramBotService) cmdGenCode(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
- if len(args) < 2 {
- return telegramCommandReply{Text: "用法:/gencode register|renew 天数 [有效天数] [可用次数]\n示例:/gencode register 30、/gencode renew 90 7 5"}
- }
- kind := strings.ToLower(strings.TrimSpace(args[0]))
- switch kind {
- case "reg", "register", "注册码":
- kind = model.RegistrationCodeRegister
- case "renew", "续期", "续期码":
- kind = model.RegistrationCodeRenew
- default:
- return telegramCommandReply{Text: "类型无效,只支持 register / renew。"}
- }
- days, err := strconv.Atoi(args[1])
- if err != nil || days < 0 {
- return telegramCommandReply{Text: "天数必须是非负整数,0 表示永久。"}
- }
- validDays := 0
- if len(args) > 2 {
- validDays, err = strconv.Atoi(args[2])
- if err != nil || validDays < 0 {
- return telegramCommandReply{Text: "有效天数必须是非负整数。"}
- }
- }
- maxUses := 1
- if len(args) > 3 {
- maxUses, err = strconv.Atoi(args[3])
- if err != nil || maxUses <= 0 {
- return telegramCommandReply{Text: "可用次数必须是正整数。"}
- }
- }
- createdBy := ""
- if u := s.boundUser(ctx, msg.From.ID); u != nil {
- createdBy = u.ID
- }
- code, err := s.generateCodeWithUses(ctx, kind, days, validDays, maxUses, createdBy)
- if err != nil {
- return telegramCommandReply{Text: "生成失败:" + err.Error()}
- }
- kindLabel := map[string]string{model.RegistrationCodeRegister: "注册码", model.RegistrationCodeRenew: "续期码"}[code.Kind]
- dur := "永久"
- if days > 0 {
- dur = fmt.Sprintf("%d 天", days)
- }
- valid := "长期有效"
- if validDays > 0 && code.ExpiresAt != nil {
- valid = "有效至 " + code.ExpiresAt.Format("2006-01-02 15:04")
- }
- uses := "单次使用"
- if code.EffectiveMaxUses() > 1 {
- uses = fmt.Sprintf("最多 %d 次", code.EffectiveMaxUses())
- }
- return telegramCommandReply{Text: fmt.Sprintf("已生成%s(%s,%s,%s):\n\n%s", kindLabel, dur, valid, uses, code.Code)}
-}
-
-func (s *TelegramBotService) replyUserList(ctx context.Context) telegramCommandReply {
- users, err := s.repo.User.List(ctx)
- if err != nil {
- return telegramCommandReply{Text: "读取用户失败:" + err.Error()}
- }
- if len(users) == 0 {
- return telegramCommandReply{Text: "暂无用户。"}
- }
- var rows [][]telegramInlineButton
- limit := len(users)
- if limit > 12 {
- limit = 12
- }
- for i := 0; i < limit; i++ {
- u := users[i]
- flag := ""
- if !u.IsActive {
- flag = "🚫"
- }
- if u.Role == "admin" {
- flag = "👑"
- }
- rows = append(rows, []telegramInlineButton{{Text: flag + " " + u.Username, Data: "usr:" + u.ID}})
- }
- rows = append(rows, []telegramInlineButton{{Text: "⬅️ 返回菜单", Data: "menu_main"}})
- return telegramCommandReply{Text: fmt.Sprintf("用户管理(共 %d 人,显示前 %d)\n点击用户进行操作:", len(users), limit), Buttons: rows}
-}
-
-func (s *TelegramBotService) replyUserActions(ctx context.Context, userID string) telegramCommandReply {
- u, err := s.repo.User.FindByID(ctx, userID)
- if err != nil || u == nil {
- return telegramCommandReply{Text: "用户不存在。"}
- }
- protected := UserIsProtectedAccount(ctx, s.repo, u)
- text := fmt.Sprintf("%s\n角色:%s\n状态:%s\n到期:%s\n防共享警告:%d 次",
- u.Username, u.Role, map[bool]string{true: "正常", false: "已禁用"}[u.IsActive], formatExpiry(u.ExpiredAt), u.ShareWarnings)
- if protected {
- return telegramCommandReply{Text: text + "\n\n(受保护账号,不可禁用/删除)", Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回", Data: "adm_users"}}}}
- }
- banBtn := telegramInlineButton{Text: "🚫 禁用", Data: "uban:" + u.ID}
- if !u.IsActive {
- banBtn = telegramInlineButton{Text: "✅ 解禁", Data: "uunban:" + u.ID}
- }
- return telegramCommandReply{
- Text: text,
- Buttons: [][]telegramInlineButton{
- {banBtn, {Text: "⏳ 续期30天", Data: "urenew:" + u.ID + ":30"}},
- {{Text: "🗑 删除用户", Data: "udel:" + u.ID}},
- {{Text: "⬅️ 返回", Data: "adm_users"}},
- },
- }
-}
-
-func (s *TelegramBotService) replyUserBan(ctx context.Context, userID string, unban bool) telegramCommandReply {
- if !unban {
- if reason := s.protectReason(ctx, userID); reason != "" {
- return telegramCommandReply{Text: reason}
- }
- }
- updates := map[string]any{"is_active": unban}
- if unban {
- updates["share_warnings"] = 0
- updates["last_share_warn_at"] = nil
- }
- if err := s.repo.User.UpdateFields(ctx, userID, updates); err != nil {
- return telegramCommandReply{Text: "操作失败:" + err.Error()}
- }
- if unban {
- _ = s.repo.UserDevice.SetKickedByUser(ctx, userID, false)
- }
- return s.replyUserActions(ctx, userID)
-}
-
-func (s *TelegramBotService) replyUserDelete(ctx context.Context, userID string) telegramCommandReply {
- if reason := s.protectReason(ctx, userID); reason != "" {
- return telegramCommandReply{Text: reason}
- }
- u, _ := s.repo.User.FindByID(ctx, userID)
- _ = s.repo.UserDevice.DeleteByUser(ctx, userID)
- if err := s.repo.User.Delete(ctx, userID); err != nil {
- return telegramCommandReply{Text: "删除失败:" + err.Error()}
- }
- name := userID
- if u != nil {
- name = u.Username
- }
- return telegramCommandReply{Text: fmt.Sprintf("已删除用户 %s。", name), Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回", Data: "adm_users"}}}}
-}
-
-func (s *TelegramBotService) replyUserRenew(ctx context.Context, payload string) telegramCommandReply {
- parts := strings.Split(payload, ":") // :
- if len(parts) != 2 {
- return telegramCommandReply{Text: "参数错误。"}
- }
- days, _ := strconv.Atoi(parts[1])
- if err := s.applyRenewal(ctx, parts[0], days); err != nil {
- return telegramCommandReply{Text: "续期失败:" + err.Error()}
- }
- return s.replyUserActions(ctx, parts[0])
-}
-
-func (s *TelegramBotService) cmdUserRenew(ctx context.Context, args []string) telegramCommandReply {
- if len(args) < 2 {
- return telegramCommandReply{Text: "用法:/renew_user 用户名 天数,天数 0 表示永久。"}
- }
- user, _ := s.repo.User.FindByUsername(ctx, args[0])
- if user == nil {
- user, _ = s.repo.User.FindByID(ctx, args[0])
- }
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- days, err := strconv.Atoi(args[1])
- if err != nil || days < 0 {
- return telegramCommandReply{Text: "天数必须是非负整数。"}
- }
- if err := s.applyRenewal(ctx, user.ID, days); err != nil {
- return telegramCommandReply{Text: "续期失败:" + err.Error()}
- }
- return s.replyUserActions(ctx, user.ID)
-}
-
-func (s *TelegramBotService) cmdUserDelete(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "用法:/delete_user 用户名 confirm\n为避免误删,最后一个参数必须是 confirm。"}
- }
- if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") {
- return telegramCommandReply{Text: "删除用户需要确认:/delete_user 用户名 confirm"}
- }
- user, _ := s.repo.User.FindByUsername(ctx, args[0])
- if user == nil {
- user, _ = s.repo.User.FindByID(ctx, args[0])
- }
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- return s.replyUserDelete(ctx, user.ID)
-}
-
-func (s *TelegramBotService) cmdUnbind(ctx context.Context, args []string) telegramCommandReply {
- targets := parseTelegramUnbindTargets(args)
- if len(targets) == 0 {
- return telegramCommandReply{Text: "用法:/unbind 用户名1 用户名2\n也支持逗号分隔,或使用 tg:TelegramID 按 Telegram ID 解绑。此命令只解绑 Bot,不删除媒体账号。"}
- }
- var removed int64
- var done []string
- var skipped []string
- var missing []string
- for _, target := range targets {
- if tgIDRaw, ok := strings.CutPrefix(strings.ToLower(target), "tg:"); ok {
- tgID, err := strconv.ParseInt(tgIDRaw, 10, 64)
- if err != nil || tgID == 0 {
- missing = append(missing, target)
- continue
- }
- n, err := s.deleteTelegramBindings(ctx, "telegram_user_id = ?", tgID)
- if err != nil {
- return telegramCommandReply{Text: "解绑失败:" + err.Error()}
- }
- if n == 0 {
- missing = append(missing, target)
- continue
- }
- removed += n
- done = append(done, target)
- continue
- }
-
- user, _ := s.repo.User.FindByUsername(ctx, target)
- if user == nil {
- user, _ = s.repo.User.FindByID(ctx, target)
- }
- if user == nil {
- missing = append(missing, target)
- continue
- }
- if user.Role == "admin" {
- skipped = append(skipped, user.Username+"(管理员)")
- continue
- }
- n, err := s.deleteTelegramBindings(ctx, "user_id = ?", user.ID)
- if err != nil {
- return telegramCommandReply{Text: "解绑失败:" + err.Error()}
- }
- if n == 0 {
- missing = append(missing, user.Username+"(未绑定)")
- continue
- }
- removed += n
- done = append(done, user.Username)
- }
- return formatUnbindResult("批量解绑完成", removed, done, skipped, missing)
-}
-
-func (s *TelegramBotService) cmdUnbindDuplicates(ctx context.Context) telegramCommandReply {
- if s == nil || s.repo == nil || s.repo.DB == nil {
- return telegramCommandReply{Text: "仓库不可用。"}
- }
- var bindings []model.TelegramBinding
- if err := s.repo.DB.WithContext(ctx).Order("updated_at desc, created_at desc").Find(&bindings).Error; err != nil {
- return telegramCommandReply{Text: "读取绑定失败:" + err.Error()}
- }
- seenTelegram := make(map[int64]string)
- seenUser := make(map[string]string)
- var removeIDs []string
- var removedLabels []string
- for _, binding := range bindings {
- remove := false
- if binding.UserID == "" || binding.TelegramUserID == 0 {
- remove = true
- } else if user, _ := s.repo.User.FindByID(ctx, binding.UserID); user == nil {
- remove = true
- } else if _, ok := seenTelegram[binding.TelegramUserID]; ok {
- remove = true
- } else if _, ok := seenUser[binding.UserID]; ok {
- remove = true
- }
- if remove {
- removeIDs = append(removeIDs, binding.ID)
- removedLabels = append(removedLabels, fmt.Sprintf("tg:%d", binding.TelegramUserID))
- continue
- }
- seenTelegram[binding.TelegramUserID] = binding.ID
- seenUser[binding.UserID] = binding.ID
- }
- if len(removeIDs) == 0 {
- return telegramCommandReply{Text: "未发现重复或无效绑定。"}
- }
- n, err := s.deleteTelegramBindings(ctx, "id IN ?", removeIDs)
- if err != nil {
- return telegramCommandReply{Text: "清理失败:" + err.Error()}
- }
- return formatUnbindResult("重复/无效绑定清理完成", n, removedLabels, nil, nil)
-}
-
-func (s *TelegramBotService) cmdUnbindInactive(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "用法:/unbind_inactive 天数\n例如 /unbind_inactive 30 会解绑 30 天未登录的普通用户 Bot 绑定,不删除账号。"}
- }
- days, err := strconv.Atoi(strings.TrimSpace(args[0]))
- if err != nil || days < 1 {
- return telegramCommandReply{Text: "天数必须是大于 0 的整数。"}
- }
- users, err := s.repo.User.List(ctx)
- if err != nil {
- return telegramCommandReply{Text: "读取用户失败:" + err.Error()}
- }
- cutoff := time.Now().Add(-time.Duration(days) * 24 * time.Hour)
- var userIDs []string
- var done []string
- for _, user := range users {
- if user.Role == "admin" {
- continue
- }
- lastActive := user.CreatedAt
- if user.LastLoginAt != nil {
- lastActive = *user.LastLoginAt
- }
- if lastActive.IsZero() || lastActive.After(cutoff) {
- continue
- }
- var count int64
- _ = s.repo.DB.WithContext(ctx).Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error
- if count == 0 {
- continue
- }
- userIDs = append(userIDs, user.ID)
- done = append(done, user.Username)
- }
- if len(userIDs) == 0 {
- return telegramCommandReply{Text: fmt.Sprintf("未发现 %d 天未登录且已绑定 Bot 的普通用户。", days)}
- }
- n, err := s.deleteTelegramBindings(ctx, "user_id IN ?", userIDs)
- if err != nil {
- return telegramCommandReply{Text: "解绑失败:" + err.Error()}
- }
- return formatUnbindResult(fmt.Sprintf("已解绑 %d 天未登录用户", days), n, done, nil, nil)
-}
-
-func parseTelegramUnbindTargets(args []string) []string {
- seen := make(map[string]struct{})
- var targets []string
- for _, arg := range args {
- for _, part := range strings.FieldsFunc(arg, func(r rune) bool {
- return r == ',' || r == ',' || r == ';' || r == ';' || r == '\n' || r == '\t'
- }) {
- part = strings.TrimSpace(part)
- if part == "" {
- continue
- }
- key := strings.ToLower(part)
- if _, ok := seen[key]; ok {
- continue
- }
- seen[key] = struct{}{}
- targets = append(targets, part)
- }
- }
- return targets
-}
-
-func (s *TelegramBotService) deleteTelegramBindings(ctx context.Context, query string, args ...interface{}) (int64, error) {
- if s == nil || s.repo == nil || s.repo.DB == nil {
- return 0, nil
- }
- tx := s.repo.DB.WithContext(ctx).Unscoped().Where(query, args...).Delete(&model.TelegramBinding{})
- return tx.RowsAffected, tx.Error
-}
-
-func formatUnbindResult(title string, removed int64, done, skipped, missing []string) telegramCommandReply {
- var sb strings.Builder
- sb.WriteString("")
- sb.WriteString(title)
- sb.WriteString("\n\n")
- sb.WriteString(fmt.Sprintf("已解绑:%d 条绑定", removed))
- if len(done) > 0 {
- sb.WriteString("\n目标:")
- sb.WriteString(formatShortList(done, 12))
- }
- if len(skipped) > 0 {
- sb.WriteString("\n跳过:")
- sb.WriteString(formatShortList(skipped, 8))
- }
- if len(missing) > 0 {
- sb.WriteString("\n未找到/未绑定:")
- sb.WriteString(formatShortList(missing, 8))
- }
- return telegramCommandReply{Text: sb.String()}
-}
-
-func formatShortList(items []string, limit int) string {
- if len(items) == 0 {
- return ""
- }
- if limit < 1 {
- limit = 1
- }
- out := items
- if len(out) > limit {
- out = out[:limit]
- }
- text := "" + strings.Join(out, "、") + ""
- if len(items) > limit {
- text += fmt.Sprintf(" 等 %d 项", len(items))
- }
- return text
-}
-
-// protectReason returns a non-empty message when a user must not be
-// disabled/deleted (admins, default admin and protected-list users).
-func (s *TelegramBotService) protectReason(ctx context.Context, userID string) string {
- u, err := s.repo.User.FindByID(ctx, userID)
- if err != nil || u == nil {
- return "用户不存在。"
- }
- if u.Role == "admin" {
- return "管理员账号受保护,不可禁用/删除。"
- }
- if first, _ := s.repo.User.FirstAdmin(ctx); first != nil && first.ID == u.ID {
- return "默认管理员账号受保护,不可禁用/删除。"
- }
- if _, ok := ProtectedUserIDSet(ctx, s.repo)[u.ID]; ok {
- return "该账号在 Bot 保护名单中,不可禁用/删除。"
- }
- return ""
-}
-
-func (s *TelegramBotService) replyDevicePolicy(ctx context.Context) telegramCommandReply {
- cfg := loadBotConfig(ctx, s.repo)
- text := fmt.Sprintf(
- "设备策略\n\n① 防共享:%s\n 并发播放终端上限 %d / 登录终端上限 %d;同一终端多个 App 只算 1 台,App 作为登录渠道记录。\n 设备指纹异常警告 %d 次后禁用账号。\n\n② Mgo 保号规则:%s\n 保号模式:%s;启用规则 %d 条。\n\n命令:\n/antishare on play=3 login=3 warn=2\n/cleanup run 预览候选\n/cleanup run confirm 确认清理\n/cleanup on|off\n/cleanup_rule list|add|edit|修改|del|enable|disable\n\n策略默认关闭;清理前会先预览候选;满足任意一条保号规则即保留;管理员/受保护账号永不自动处理。",
- onOff(cfg.AntiShareEnabled), cfg.MaxConcurrentPlay, cfg.MaxLoggedClients, cfg.WarnThreshold,
- onOff(cfg.AccountCleanupEnabled), cleanupModeLabel(cfg.AccountCleanupKeepMode), countEnabledCleanupRules(cfg.AccountCleanupRules))
- return telegramCommandReply{
- Text: text,
- Buttons: [][]telegramInlineButton{
- {{Text: toggleLabel("防共享", cfg.AntiShareEnabled), Data: "dp_toggle:antishare"}},
- {{Text: toggleLabel("保号规则", cfg.AccountCleanupEnabled), Data: "dp_toggle:cleanup"}},
- {{Text: "⬅️ 返回菜单", Data: "menu_main"}},
- },
- }
-}
-
-func (s *TelegramBotService) cmdDevicePolicy(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 || strings.EqualFold(args[0], "status") {
- return s.replyDevicePolicy(ctx)
- }
- switch strings.ToLower(strings.TrimSpace(args[0])) {
- case "run", "sweep":
- return s.cmdCleanup(ctx, []string{"run"})
- default:
- return telegramCommandReply{Text: "用法:/devicepolicy 查看策略,或使用 /antishare、/cleanup、/cleanup_rule 管理。"}
- }
-}
-
-func (s *TelegramBotService) cmdAntiShare(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 || strings.EqualFold(args[0], "status") {
- return s.replyDevicePolicy(ctx)
- }
- enabled, ok := parseCommandBool(args[0])
- if !ok {
- return telegramCommandReply{Text: "用法:/antishare on|off [play=3] [login=3] [warn=2],login 表示登录终端设备上限,同一终端多个 App 不重复计数。"}
- }
- if err := s.repo.Setting.Set(ctx, SettingAntiShareEnabled, strconv.FormatBool(enabled)); err != nil {
- return telegramCommandReply{Text: "更新失败:" + err.Error()}
- }
- for _, arg := range args[1:] {
- key, value, ok := strings.Cut(arg, "=")
- if !ok {
- continue
- }
- n, err := strconv.Atoi(strings.TrimSpace(value))
- if err != nil || n < 1 {
- continue
- }
- switch strings.ToLower(strings.TrimSpace(key)) {
- case "play", "maxplay", "播放":
- _ = s.repo.Setting.Set(ctx, SettingMaxConcurrentPlay, strconv.Itoa(n))
- case "login", "client", "clients", "登录":
- _ = s.repo.Setting.Set(ctx, SettingMaxLoggedClients, strconv.Itoa(n))
- case "warn", "warnings", "警告":
- _ = s.repo.Setting.Set(ctx, SettingWarnThreshold, strconv.Itoa(n))
- }
- }
- return s.replyDevicePolicy(ctx)
-}
-
-func (s *TelegramBotService) cmdCleanup(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 || strings.EqualFold(args[0], "status") {
- return s.replyDevicePolicy(ctx)
- }
- switch strings.ToLower(strings.TrimSpace(args[0])) {
- case "on", "true", "1", "开启", "enable":
- if err := s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
- return telegramCommandReply{Text: "开启失败:" + err.Error()}
- }
- return s.replyDevicePolicy(ctx)
- case "off", "false", "0", "关闭", "disable":
- if err := s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, "false"); err != nil {
- return telegramCommandReply{Text: "关闭失败:" + err.Error()}
- }
- return s.replyDevicePolicy(ctx)
- case "run", "sweep", "巡检", "preview", "预览":
- device := s.device
- if device == nil {
- device = NewDeviceService(s.log, s.repo)
- }
- if len(args) > 1 && isCleanupConfirmArg(args[1]) {
- cfg := loadBotConfig(ctx, s.repo)
- if !cfg.AccountCleanupEnabled {
- return telegramCommandReply{Text: "保号规则未开启,不会清理账号。"}
- }
- if countEnabledCleanupRules(cfg.AccountCleanupRules) == 0 {
- return telegramCommandReply{Text: "没有启用的保号规则,不会清理账号。"}
- }
- removed, err := device.SweepAccountCleanup(ctx)
- if err != nil {
- return telegramCommandReply{Text: "确认清理失败:" + err.Error()}
- }
- return telegramCommandReply{Text: fmt.Sprintf("保号规则确认清理完成,已清理 %d 个账号。", removed)}
- }
- candidates, err := device.PreviewAccountCleanup(ctx)
- if err != nil {
- return telegramCommandReply{Text: "巡检预览失败:" + err.Error()}
- }
- return telegramCommandReply{Text: s.formatCleanupPreview(ctx, candidates)}
- default:
- return telegramCommandReply{Text: "用法:/cleanup on|off、/cleanup run 预览、/cleanup run confirm 确认清理"}
- }
-}
-
-func isCleanupConfirmArg(arg string) bool {
- switch strings.ToLower(strings.TrimSpace(arg)) {
- case "confirm", "yes", "delete", "确认", "清理", "删除":
- return true
- default:
- return false
- }
-}
-
-func (s *TelegramBotService) formatCleanupPreview(ctx context.Context, candidates []accountCleanupCandidate) string {
- cfg := loadBotConfig(ctx, s.repo)
- if !cfg.AccountCleanupEnabled {
- return "保号规则未开启,不会清理账号。"
- }
- if countEnabledCleanupRules(cfg.AccountCleanupRules) == 0 {
- return "没有启用的保号规则,不会清理账号。"
- }
- if len(candidates) == 0 {
- return "保号规则预览完成:没有需要清理的账号。"
- }
- var sb strings.Builder
- sb.WriteString(fmt.Sprintf("保号规则预览\n\n将清理候选:%d 个账号。\n当前只是预览,未删除任何账号。\n\n", len(candidates)))
- limit := len(candidates)
- if limit > 10 {
- limit = 10
- }
- for i := 0; i < limit; i++ {
- candidate := candidates[i]
- sb.WriteString(fmt.Sprintf("%d. %s\n%s\n", i+1, escapeHTML(candidate.Username), escapeHTML(candidate.Details)))
- }
- if len(candidates) > limit {
- sb.WriteString(fmt.Sprintf("……另有 %d 个候选未展示。\n", len(candidates)-limit))
- }
- sb.WriteString("\n确认无误后再执行:/cleanup run confirm")
- return sb.String()
-}
-
-func (s *TelegramBotService) cmdCleanupMode(ctx context.Context, args []string) telegramCommandReply {
- if err := s.repo.Setting.Set(ctx, SettingAccountCleanupKeepMode, "any"); err != nil {
- return telegramCommandReply{Text: "更新失败:" + err.Error()}
- }
- if err := s.repo.Setting.Set(ctx, SettingAccountCleanupRequiredCount, "1"); err != nil {
- return telegramCommandReply{Text: "更新失败:" + err.Error()}
- }
- reply := s.replyDevicePolicy(ctx)
- reply.Text = "Mgo 保号模式固定为:满足任意一条启用规则即保留;只有全部规则都不满足才进入清理候选。\n\n" + reply.Text
- return reply
-}
-
-func (s *TelegramBotService) cmdCleanupRule(ctx context.Context, args []string) telegramCommandReply {
- rules := s.currentCleanupRules(ctx)
- if len(args) == 0 {
- return telegramCommandReply{Text: formatCleanupRules(rules)}
- }
- action := strings.ToLower(strings.TrimSpace(args[0]))
- switch action {
- case "list", "ls", "status":
- return telegramCommandReply{Text: formatCleanupRules(rules)}
- case "help", "?", "帮助":
- return telegramCommandReply{Text: cleanupRuleHelp()}
- case "del", "delete", "rm":
- if len(args) < 2 {
- return telegramCommandReply{Text: "用法:/cleanup_rule del 规则ID"}
- }
- next := make([]accountCleanupRule, 0, len(rules))
- removed := false
- for _, r := range rules {
- if r.ID == args[1] {
- removed = true
- continue
- }
- next = append(next, r)
- }
- if !removed {
- return telegramCommandReply{Text: "未找到该规则。"}
- }
- if err := s.saveCleanupRules(ctx, next); err != nil {
- return telegramCommandReply{Text: "保存失败:" + err.Error()}
- }
- return telegramCommandReply{Text: "已删除规则。\n\n" + formatCleanupRules(next)}
- case "enable", "on", "disable", "off":
- if len(args) < 2 {
- return telegramCommandReply{Text: "用法:/cleanup_rule enable|disable 规则ID"}
- }
- enable := action == "enable" || action == "on"
- changed := false
- for i := range rules {
- if rules[i].ID == args[1] {
- rules[i].Enabled = enable
- changed = true
- }
- }
- if !changed {
- return telegramCommandReply{Text: "未找到该规则。"}
- }
- if err := s.saveCleanupRules(ctx, rules); err != nil {
- return telegramCommandReply{Text: "保存失败:" + err.Error()}
- }
- return telegramCommandReply{Text: "已更新规则状态。\n\n" + formatCleanupRules(rules)}
- case "add", "set", "edit", "update", "修改", "更新", "改":
- rule, err := parseCleanupRuleCommand(args[1:])
- if err != nil {
- return telegramCommandReply{Text: err.Error() + "\n\n" + cleanupRuleHelp()}
- }
- updated := false
- for i := range rules {
- if rules[i].ID == rule.ID {
- rules[i] = rule
- updated = true
- break
- }
- }
- if !updated {
- rules = append(rules, rule)
- }
- rules = normalizeCleanupRules(rules)
- if err := s.saveCleanupRules(ctx, rules); err != nil {
- return telegramCommandReply{Text: "保存失败:" + err.Error()}
- }
- actionText := "已新增规则。"
- if updated {
- actionText = "已更新规则。"
- }
- return telegramCommandReply{Text: actionText + "\n\n" + formatCleanupRules(rules)}
- default:
- return telegramCommandReply{Text: cleanupRuleHelp()}
- }
-}
-
-func (s *TelegramBotService) replyDevicePolicyToggle(ctx context.Context, which string) telegramCommandReply {
- cfg := loadBotConfig(ctx, s.repo)
- switch which {
- case "antishare":
- _ = s.repo.Setting.Set(ctx, SettingAntiShareEnabled, strconv.FormatBool(!cfg.AntiShareEnabled))
- case "cleanup":
- _ = s.repo.Setting.Set(ctx, SettingAccountCleanupEnabled, strconv.FormatBool(!cfg.AccountCleanupEnabled))
- }
- return s.replyDevicePolicy(ctx)
-}
-
-func (s *TelegramBotService) cmdUserBan(ctx context.Context, args []string, unban bool) telegramCommandReply {
- if len(args) == 0 {
- if unban {
- return telegramCommandReply{Text: "用法:/unban 用户名"}
- }
- return telegramCommandReply{Text: "用法:/ban 用户名"}
- }
- user, _ := s.repo.User.FindByUsername(ctx, args[0])
- if user == nil {
- user, _ = s.repo.User.FindByID(ctx, args[0])
- }
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- return s.replyUserBan(ctx, user.ID, unban)
-}
-
-func (s *TelegramBotService) currentCleanupRules(ctx context.Context) []accountCleanupRule {
- cfg := loadBotConfig(ctx, s.repo)
- return cfg.AccountCleanupRules
-}
-
-func (s *TelegramBotService) saveCleanupRules(ctx context.Context, rules []accountCleanupRule) error {
- raw, err := json.Marshal(normalizeCleanupRules(rules))
- if err != nil {
- return err
- }
- return s.repo.Setting.Set(ctx, SettingAccountCleanupRules, string(raw))
-}
-
-func parseCommandBool(value string) (bool, bool) {
- switch strings.ToLower(strings.TrimSpace(value)) {
- case "on", "true", "1", "yes", "enable", "enabled", "开启", "开":
- return true, true
- case "off", "false", "0", "no", "disable", "disabled", "关闭", "关":
- return false, true
- default:
- return false, false
- }
-}
-
-func parseCleanupRuleCommand(args []string) (accountCleanupRule, error) {
- if len(args) < 2 {
- return accountCleanupRule{}, fmt.Errorf("新增规则参数不足")
- }
- rule := accountCleanupRule{
- Type: strings.ToLower(strings.TrimSpace(args[0])),
- ID: strings.TrimSpace(args[1]),
- Enabled: true,
- WindowDaysMin: 3,
- WindowDaysMax: 5,
- MinHours: 6,
- MinCount: 1,
- }
- switch rule.Type {
- case "watch_hours":
- name, values := cleanupRuleNameAndValues(args[2:], 3)
- rule.Name = name
- if len(values) >= 3 {
- rule.WindowDaysMin, _ = strconv.Atoi(values[0])
- rule.WindowDaysMax, _ = strconv.Atoi(values[1])
- rule.MinHours, _ = strconv.ParseFloat(values[2], 64)
- if rule.Name == "" {
- rule.Name = fmt.Sprintf("%d~%d 天观看满 %s 小时", rule.WindowDaysMin, rule.WindowDaysMax, formatRuleHours(rule.MinHours))
- }
- }
- case "recent_login":
- name, values := cleanupRuleNameAndValues(args[2:], 1)
- rule.Name = name
- if len(values) >= 1 {
- rule.WindowDaysMax, _ = strconv.Atoi(values[0])
- if rule.Name == "" {
- rule.Name = fmt.Sprintf("%d 天内登录", rule.WindowDaysMax)
- }
- }
- case "signin_streak", "account_age_grace":
- name, values := cleanupRuleNameAndValues(args[2:], 1)
- rule.Name = name
- if len(values) >= 1 {
- rule.MinCount, _ = strconv.Atoi(values[0])
- if rule.Name == "" {
- if rule.Type == "signin_streak" {
- rule.Name = fmt.Sprintf("连续签到 %d 天", rule.MinCount)
- } else {
- rule.Name = fmt.Sprintf("新号宽限 %d 天", rule.MinCount)
- }
- }
- }
- default:
- return accountCleanupRule{}, fmt.Errorf("不支持的规则类型:%s", rule.Type)
- }
- normalized := normalizeCleanupRules([]accountCleanupRule{rule})
- if len(normalized) == 0 {
- return accountCleanupRule{}, fmt.Errorf("规则无效")
- }
- return normalized[0], nil
-}
-
-func cleanupRuleNameAndValues(args []string, numericCount int) (string, []string) {
- if len(args) == 0 {
- return "", nil
- }
- if len(args) >= numericCount && cleanupRuleValuesAreNumeric(args[:numericCount]) {
- return "", args
- }
- return strings.TrimSpace(args[0]), args[1:]
-}
-
-func cleanupRuleValuesAreNumeric(values []string) bool {
- for _, value := range values {
- if _, err := strconv.ParseFloat(strings.TrimSpace(value), 64); err != nil {
- return false
- }
- }
- return true
-}
-
-func formatCleanupRules(rules []accountCleanupRule) string {
- if len(rules) == 0 {
- return "保号规则\n\n暂无规则。"
- }
- var sb strings.Builder
- sb.WriteString("保号规则\n")
- for i, r := range rules {
- state := map[bool]string{true: "启用", false: "停用"}[r.Enabled]
- detail := cleanupRuleDetail(r)
- parts := []string{
- fmt.Sprintf("\n%d. %s", i+1, r.ID),
- }
- if shouldShowCleanupRuleName(r, detail) {
- parts = append(parts, r.Name)
- }
- parts = append(parts, cleanupRuleTypeLabel(r.Type), state)
- if detail != "" {
- parts = append(parts, detail)
- }
- sb.WriteString(strings.Join(parts, " · "))
- }
- return sb.String()
-}
-
-func shouldShowCleanupRuleName(r accountCleanupRule, detail string) bool {
- name := strings.TrimSpace(r.Name)
- if name == "" || strings.EqualFold(name, r.ID) {
- return false
- }
- if detail != "" && strings.EqualFold(name, detail) {
- return false
- }
- return true
-}
-
-func cleanupRuleDetail(r accountCleanupRule) string {
- switch r.Type {
- case "watch_hours":
- return fmt.Sprintf("%d~%d 天 %s 小时", r.WindowDaysMin, r.WindowDaysMax, formatRuleHours(r.MinHours))
- case "recent_login":
- return fmt.Sprintf("%d 天内登录", r.WindowDaysMax)
- case "signin_streak":
- return fmt.Sprintf("连续签到 %d 天", r.MinCount)
- case "account_age_grace":
- return fmt.Sprintf("新号宽限 %d 天", r.MinCount)
- default:
- return ""
- }
-}
-
-func formatRuleHours(hours float64) string {
- if hours == float64(int(hours)) {
- return strconv.Itoa(int(hours))
- }
- return fmt.Sprintf("%.1f", hours)
-}
-
-func cleanupRuleTypeLabel(t string) string {
- switch t {
- case "watch_hours":
- return "观看时长"
- case "recent_login":
- return "最近登录"
- case "signin_streak":
- return "连续签到"
- case "account_age_grace":
- return "新号宽限"
- default:
- return t
- }
-}
-
-func cleanupRuleHelp() string {
- return "Mgo 保号规则命令\n\n" +
- "/cleanup_rule list — 查看规则\n" +
- "/cleanup_rule add watch_hours watch_3_5d_6h 观看3到5天满6小时 3 5 6\n" +
- "/cleanup_rule add recent_login login_7d 七天内登录 7\n" +
- "/cleanup_rule add signin_streak sign_3 连续签到3天 3\n" +
- "/cleanup_rule add account_age_grace new_7d 新号宽限7天 7\n" +
- "/cleanup_rule edit 规则类型 规则ID 名称 参数... — 修改同 ID 规则\n" +
- "/cleanup_rule 修改 规则类型 规则ID 名称 参数... — 中文修改入口\n" +
- "/cleanup_rule enable 规则ID / disable 规则ID\n" +
- "/cleanup_rule del 规则ID\n\n" +
- "保号模式固定为:满足任意一条启用规则即保留;全部不满足才会清理。"
-}
-
-func onOff(b bool) string {
- return map[bool]string{true: "已开启", false: "已关闭"}[b]
-}
-
-func toggleLabel(name string, enabled bool) string {
- if enabled {
- return "关闭" + name
- }
- return "开启" + name
-}
-
-func cleanupModeLabel(mode string) string {
- return "满足任意一条"
-}
-
-func countEnabledCleanupRules(rules []accountCleanupRule) int {
- n := 0
- for _, r := range rules {
- if r.Enabled {
- n++
- }
- }
- return n
-}
diff --git a/internal/service/telegram_menu_layout.go b/internal/service/telegram_menu_layout.go
new file mode 100644
index 0000000..851ca67
--- /dev/null
+++ b/internal/service/telegram_menu_layout.go
@@ -0,0 +1,115 @@
+package service
+
+import (
+ "context"
+ "fmt"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// mainMenu builds the button-based menu, tailored to the user's binding and
+// admin status. Ordinary users only see self-service actions; admins get an
+// extra management section.
+func (s *TelegramBotService) mainMenu(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage) telegramCommandReply {
+ isAdmin := s.telegramUserIsAdmin(ctx, channel, msg.From.ID)
+ user := s.boundUser(ctx, msg.From.ID)
+ if telegramIsGroupChat(msg.Chat.Type) {
+ return s.groupMainMenu(isAdmin, user)
+ }
+ return s.privateMainMenu(ctx, isAdmin, user)
+}
+
+func (s *TelegramBotService) groupMainMenu(isAdmin bool, user *model.User) telegramCommandReply {
+ header := "MediaStationGo 群组自助菜单\n\n你还没有绑定媒体中心账号。绑定、注册、兑换等包含敏感信息的操作请私聊 Bot。"
+ var rows [][]telegramInlineButton
+ if user != nil {
+ header = telegramUserMenuHeader("MediaStationGo 群组自助菜单", user)
+ rows = telegramBoundUserMenuRows(user, false)
+ }
+ if isAdmin {
+ header += "\n\n管理员入口"
+ rows = append(rows, telegramAdminMenuRows()...)
+ }
+ return telegramCommandReply{Text: header, Buttons: rows}
+}
+
+func (s *TelegramBotService) privateMainMenu(ctx context.Context, isAdmin bool, user *model.User) telegramCommandReply {
+ header := "MediaStationGo\n\n你还没有绑定媒体中心账号。"
+ rows := s.privateUnboundMenuRows(ctx)
+ if user != nil {
+ header = telegramUserMenuHeader("MediaStationGo", user)
+ rows = telegramBoundUserMenuRows(user, true)
+ }
+ if isAdmin {
+ rows = append(rows, telegramAdminMenuRows()...)
+ }
+ return telegramCommandReply{Text: header, Buttons: rows}
+}
+
+func (s *TelegramBotService) privateUnboundMenuRows(ctx context.Context) [][]telegramInlineButton {
+ rows := [][]telegramInlineButton{{{Text: "🔗 绑定账号", Data: "act_bind"}}}
+ if s.openRegEnabled(ctx) {
+ rows = append(rows, []telegramInlineButton{{Text: "📝 注册新账号", Data: "act_register"}})
+ }
+ return append(rows, []telegramInlineButton{{Text: "🎟 兑换码注册", Data: "act_redeem_register"}})
+}
+
+func telegramUserMenuHeader(title string, user *model.User) string {
+ return fmt.Sprintf("%s\n\n账号:%s\n到期:%s\n成人目录:%s",
+ title, user.Username, formatExpiry(user.ExpiredAt), telegramAdultVisibilityLabel(user.HideAdult))
+}
+
+func telegramAdultVisibilityLabel(hidden bool) string {
+ if hidden {
+ return "已隐藏"
+ }
+ return "已显示"
+}
+
+func telegramAdultToggleText(hidden bool) string {
+ if hidden {
+ return "🔞 显示成人目录"
+ }
+ return "🔞 隐藏成人目录"
+}
+
+func telegramBoundUserMenuRows(user *model.User, includePrivateActions bool) [][]telegramInlineButton {
+ rows := [][]telegramInlineButton{
+ {
+ {Text: "👤 我的账号", Data: "act_account"},
+ {Text: "📅 签到", Data: "act_signin"},
+ },
+ {
+ {Text: "📱 我的设备", Data: "act_devices"},
+ {Text: telegramAdultToggleText(user.HideAdult), Data: "adult_toggle"},
+ },
+ }
+ if includePrivateActions {
+ rows = append(rows,
+ []telegramInlineButton{
+ {Text: "✏️ 改用户名", Data: "act_setname"},
+ {Text: "🔑 改密码", Data: "act_setpass"},
+ },
+ []telegramInlineButton{{Text: "🎟 兑换码续期", Data: "act_redeem_renew"}},
+ )
+ }
+ return rows
+}
+
+func telegramAdminMenuRows() [][]telegramInlineButton {
+ return [][]telegramInlineButton{
+ {{Text: "—— 管理员 ——", Data: "noop"}},
+ {
+ {Text: "📊 容量/状态", Data: "adm_capacity"},
+ {Text: "👥 用户管理", Data: "adm_users"},
+ },
+ {
+ {Text: "🔓 开注设置", Data: "adm_openreg"},
+ {Text: "🎟 生成兑换码", Data: "adm_gencode"},
+ },
+ {
+ {Text: "⚙️ 设备策略", Data: "adm_devicepolicy"},
+ {Text: "🛠 管理命令", Data: "adm_mgo_commands"},
+ },
+ }
+}
diff --git a/internal/service/telegram_mgo_compat.go b/internal/service/telegram_mgo_compat.go
index 7c75931..bd5a928 100644
--- a/internal/service/telegram_mgo_compat.go
+++ b/internal/service/telegram_mgo_compat.go
@@ -12,169 +12,6 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/model"
)
-func (s *TelegramBotService) cmdMgoCreateUser(ctx context.Context, args []string) telegramCommandReply {
- if len(args) < 2 {
- return telegramCommandReply{Text: "用法:/ucr 用户名 密码 [天数],天数 0 表示永久。"}
- }
- if s.auth == nil {
- return telegramCommandReply{Text: "注册服务暂不可用。"}
- }
- user, _, err := s.auth.Register(ctx, args[0], args[1])
- if err != nil {
- return telegramCommandReply{Text: "创建失败:" + err.Error()}
- }
- days := 0
- if len(args) >= 3 {
- parsed, err := strconv.Atoi(args[2])
- if err != nil || parsed < 0 {
- return telegramCommandReply{Text: "账号已创建,但天数无效。请用 /renew 用户名 天数 调整。"}
- }
- days = parsed
- if err := s.applyRenewal(ctx, user.ID, days); err != nil {
- return telegramCommandReply{Text: "账号已创建,但续期失败:" + err.Error()}
- }
- }
- return telegramCommandReply{Text: fmt.Sprintf("已创建用户:%s\n到期:%s", user.Username, formatExpiry(s.userExpiry(ctx, user.ID)))}
-}
-
-func (s *TelegramBotService) cmdMgoUserInfo(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "用法:/uinfo 用户名"}
- }
- user := s.findMgoBotUser(ctx, args[0])
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- devices, _ := s.repo.UserDevice.ListByUser(ctx, user.ID)
- var historyCount int64
- _ = s.repo.DB.WithContext(ctx).Model(&model.PlaybackHistory{}).Where("user_id = ?", user.ID).Count(&historyCount).Error
- var binding model.TelegramBinding
- tg := "未绑定"
- if err := s.repo.DB.WithContext(ctx).Where("user_id = ?", user.ID).First(&binding).Error; err == nil {
- tg = fmt.Sprintf("tg:%d", binding.TelegramUserID)
- if binding.TelegramName != "" {
- tg += " " + binding.TelegramName
- }
- }
- return telegramCommandReply{Text: fmt.Sprintf(
- "用户信息\n\n用户名:%s\n角色:%s\n状态:%s\n到期:%s\nTelegram:%s\n设备:%d\n播放记录:%d\n最后登录:%s",
- user.Username, user.Role, activeLabel(user), formatExpiry(user.ExpiredAt), tg, len(devices), historyCount, formatOptionalTime(user.LastLoginAt),
- )}
-}
-
-func (s *TelegramBotService) cmdMgoDeleteUser(ctx context.Context, args []string) telegramCommandReply {
- if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") {
- return telegramCommandReply{Text: "删除用户需要确认:/rmemby 用户名 confirm 或 /urm 用户名 confirm"}
- }
- user := s.findMgoBotUser(ctx, args[0])
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- if reason := s.protectReason(ctx, user.ID); reason != "" {
- return telegramCommandReply{Text: reason}
- }
- _ = s.repo.UserDevice.DeleteByUser(ctx, user.ID)
- if err := s.repo.User.Delete(ctx, user.ID); err != nil {
- return telegramCommandReply{Text: "删除失败:" + err.Error()}
- }
- return telegramCommandReply{Text: fmt.Sprintf("已删除用户 %s。", user.Username)}
-}
-
-func (s *TelegramBotService) cmdMgoOnlyRemoveRecord(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "用法:/only_rm_record tg:123456 或 /only_rm_record 用户名,只删除 Telegram 绑定记录。"}
- }
- target := strings.TrimSpace(args[0])
- var removed int64
- if raw, ok := strings.CutPrefix(strings.ToLower(target), "tg:"); ok {
- tgID, err := strconv.ParseInt(raw, 10, 64)
- if err != nil || tgID == 0 {
- return telegramCommandReply{Text: "Telegram ID 无效。"}
- }
- removed, err = s.deleteTelegramBindings(ctx, "telegram_user_id = ?", tgID)
- if err != nil {
- return telegramCommandReply{Text: "删除绑定失败:" + err.Error()}
- }
- } else {
- user := s.findMgoBotUser(ctx, target)
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- n, err := s.deleteTelegramBindings(ctx, "user_id = ?", user.ID)
- if err != nil {
- return telegramCommandReply{Text: "删除绑定失败:" + err.Error()}
- }
- removed = n
- }
- return telegramCommandReply{Text: fmt.Sprintf("已删除 Telegram 绑定记录:%d 条。", removed)}
-}
-
-func (s *TelegramBotService) cmdMgoUserIP(ctx context.Context, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: "用法:/userip 用户名"}
- }
- user := s.findMgoBotUser(ctx, args[0])
- if user == nil {
- return telegramCommandReply{Text: "未找到用户。"}
- }
- devices, err := s.repo.UserDevice.ListByUser(ctx, user.ID)
- if err != nil {
- return telegramCommandReply{Text: "查询失败:" + err.Error()}
- }
- if len(devices) == 0 {
- return telegramCommandReply{Text: "该用户暂无设备/IP记录。"}
- }
- var out []string
- for i, d := range devices {
- if i >= 20 {
- break
- }
- out = append(out, fmt.Sprintf("%d. %s / %s / %s / %s", i+1, blankDash(d.LastIP), blankDash(d.DeviceName), blankDash(d.Client), d.LastSeenAt.Format("2006-01-02 15:04")))
- }
- return telegramCommandReply{Text: "" + user.Username + " 的设备/IP\n\n" + strings.Join(out, "\n") + ""}
-}
-
-func (s *TelegramBotService) cmdMgoAuditDevices(ctx context.Context, mode string, args []string) telegramCommandReply {
- if len(args) == 0 {
- return telegramCommandReply{Text: fmt.Sprintf("用法:/%s 关键词", mode)}
- }
- keyword := strings.TrimSpace(strings.Join(args, " "))
- var rows []struct {
- Username string
- DeviceID string
- DeviceName string
- Client string
- LastIP string
- LastSeenAt time.Time
- }
- q := s.repo.DB.WithContext(ctx).Table("user_devices").
- Select("users.username, user_devices.device_id, user_devices.device_name, user_devices.client, user_devices.last_ip, user_devices.last_seen_at").
- Joins("JOIN users ON users.id = user_devices.user_id").
- Order("user_devices.last_seen_at desc").
- Limit(20)
- switch mode {
- case "auditip":
- q = q.Where("user_devices.last_ip LIKE ?", "%"+keyword+"%")
- case "auditdevice":
- q = q.Where("user_devices.device_name LIKE ? OR user_devices.device_id LIKE ?", "%"+keyword+"%", "%"+keyword+"%")
- case "auditclient":
- q = q.Where("user_devices.client LIKE ?", "%"+keyword+"%")
- case "udeviceid":
- q = q.Where("user_devices.device_id LIKE ?", "%"+keyword+"%")
- }
- if err := q.Scan(&rows).Error; err != nil {
- return telegramCommandReply{Text: "查询失败:" + err.Error()}
- }
- if len(rows) == 0 {
- return telegramCommandReply{Text: "没有匹配记录。"}
- }
- var out []string
- for i, r := range rows {
- out = append(out, fmt.Sprintf("%d. %s / %s / %s / %s / %s", i+1, r.Username, blankDash(r.LastIP), blankDash(r.DeviceName), blankDash(r.Client), r.LastSeenAt.Format("2006-01-02 15:04")))
- }
- return telegramCommandReply{Text: "审计结果\n\n" + strings.Join(out, "\n") + ""}
-}
-
func (s *TelegramBotService) cmdMgoRenewAll(ctx context.Context, args []string) telegramCommandReply {
if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") {
return telegramCommandReply{Text: "批量续期需要确认:/renewall 天数 confirm"}
diff --git a/internal/service/telegram_mgo_users.go b/internal/service/telegram_mgo_users.go
new file mode 100644
index 0000000..eb4f685
--- /dev/null
+++ b/internal/service/telegram_mgo_users.go
@@ -0,0 +1,174 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *TelegramBotService) cmdMgoCreateUser(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "用法:/ucr 用户名 密码 [天数],天数 0 表示永久。"}
+ }
+ if s.auth == nil {
+ return telegramCommandReply{Text: "注册服务暂不可用。"}
+ }
+ user, _, err := s.auth.Register(ctx, args[0], args[1])
+ if err != nil {
+ return telegramCommandReply{Text: "创建失败:" + err.Error()}
+ }
+ days := 0
+ if len(args) >= 3 {
+ parsed, err := strconv.Atoi(args[2])
+ if err != nil || parsed < 0 {
+ return telegramCommandReply{Text: "账号已创建,但天数无效。请用 /renew 用户名 天数 调整。"}
+ }
+ days = parsed
+ if err := s.applyRenewal(ctx, user.ID, days); err != nil {
+ return telegramCommandReply{Text: "账号已创建,但续期失败:" + err.Error()}
+ }
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已创建用户:%s\n到期:%s", user.Username, formatExpiry(s.userExpiry(ctx, user.ID)))}
+}
+
+func (s *TelegramBotService) cmdMgoUserInfo(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "用法:/uinfo 用户名"}
+ }
+ user := s.findMgoBotUser(ctx, args[0])
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ devices, _ := s.listUserDevices(ctx, user.ID)
+ var historyCount int64
+ _ = s.repo.DB.WithContext(ctx).Model(&model.PlaybackHistory{}).Where("user_id = ?", user.ID).Count(&historyCount).Error
+ var binding model.TelegramBinding
+ tg := "未绑定"
+ if err := s.repo.DB.WithContext(ctx).Where("user_id = ?", user.ID).First(&binding).Error; err == nil {
+ tg = fmt.Sprintf("tg:%d", binding.TelegramUserID)
+ if binding.TelegramName != "" {
+ tg += " " + binding.TelegramName
+ }
+ }
+ return telegramCommandReply{Text: fmt.Sprintf(
+ "用户信息\n\n用户名:%s\n角色:%s\n状态:%s\n到期:%s\nTelegram:%s\n设备:%d\n播放记录:%d\n最后登录:%s",
+ user.Username, user.Role, activeLabel(user), formatExpiry(user.ExpiredAt), tg, len(devices), historyCount, formatOptionalTime(user.LastLoginAt),
+ )}
+}
+
+func (s *TelegramBotService) cmdMgoDeleteUser(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") {
+ return telegramCommandReply{Text: "删除用户需要确认:/rmemby 用户名 confirm 或 /urm 用户名 confirm"}
+ }
+ user := s.findMgoBotUser(ctx, args[0])
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ if reason := s.protectReason(ctx, user.ID); reason != "" {
+ return telegramCommandReply{Text: reason}
+ }
+ _ = s.repo.UserDevice.DeleteByUser(ctx, user.ID)
+ if err := s.repo.User.Delete(ctx, user.ID); err != nil {
+ return telegramCommandReply{Text: "删除失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已删除用户 %s。", user.Username)}
+}
+
+func (s *TelegramBotService) cmdMgoOnlyRemoveRecord(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "用法:/only_rm_record tg:123456 或 /only_rm_record 用户名,只删除 Telegram 绑定记录。"}
+ }
+ target := strings.TrimSpace(args[0])
+ var removed int64
+ if raw, ok := strings.CutPrefix(strings.ToLower(target), "tg:"); ok {
+ tgID, err := strconv.ParseInt(raw, 10, 64)
+ if err != nil || tgID == 0 {
+ return telegramCommandReply{Text: "Telegram ID 无效。"}
+ }
+ removed, err = s.deleteTelegramBindings(ctx, "telegram_user_id = ?", tgID)
+ if err != nil {
+ return telegramCommandReply{Text: "删除绑定失败:" + err.Error()}
+ }
+ } else {
+ user := s.findMgoBotUser(ctx, target)
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ n, err := s.deleteTelegramBindings(ctx, "user_id = ?", user.ID)
+ if err != nil {
+ return telegramCommandReply{Text: "删除绑定失败:" + err.Error()}
+ }
+ removed = n
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已删除 Telegram 绑定记录:%d 条。", removed)}
+}
+
+func (s *TelegramBotService) cmdMgoUserIP(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "用法:/userip 用户名"}
+ }
+ user := s.findMgoBotUser(ctx, args[0])
+ if user == nil {
+ return telegramCommandReply{Text: "未找到用户。"}
+ }
+ devices, err := s.listUserDevices(ctx, user.ID)
+ if err != nil {
+ return telegramCommandReply{Text: "查询失败:" + err.Error()}
+ }
+ if len(devices) == 0 {
+ return telegramCommandReply{Text: "该用户暂无设备/IP记录。"}
+ }
+ var out []string
+ for i, d := range devices {
+ if i >= 20 {
+ break
+ }
+ out = append(out, fmt.Sprintf("%d. %s / %s / %s / %s", i+1, blankDash(d.LastIP), blankDash(d.DeviceName), blankDash(d.Client), d.LastSeenAt.Format("2006-01-02 15:04")))
+ }
+ return telegramCommandReply{Text: "" + user.Username + " 的设备/IP\n\n" + strings.Join(out, "\n") + ""}
+}
+
+func (s *TelegramBotService) cmdMgoAuditDevices(ctx context.Context, mode string, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: fmt.Sprintf("用法:/%s 关键词", mode)}
+ }
+ keyword := strings.TrimSpace(strings.Join(args, " "))
+ var rows []struct {
+ Username string
+ DeviceID string
+ DeviceName string
+ Client string
+ LastIP string
+ LastSeenAt time.Time
+ }
+ q := s.repo.DB.WithContext(ctx).Table("user_devices").
+ Select("users.username, user_devices.device_id, user_devices.device_name, user_devices.client, user_devices.last_ip, user_devices.last_seen_at").
+ Joins("JOIN users ON users.id = user_devices.user_id").
+ Order("user_devices.last_seen_at desc").
+ Limit(20)
+ switch mode {
+ case "auditip":
+ q = q.Where("user_devices.last_ip LIKE ?", "%"+keyword+"%")
+ case "auditdevice":
+ q = q.Where("user_devices.device_name LIKE ? OR user_devices.device_id LIKE ?", "%"+keyword+"%", "%"+keyword+"%")
+ case "auditclient":
+ q = q.Where("user_devices.client LIKE ?", "%"+keyword+"%")
+ case "udeviceid":
+ q = q.Where("user_devices.device_id LIKE ?", "%"+keyword+"%")
+ }
+ if err := q.Scan(&rows).Error; err != nil {
+ return telegramCommandReply{Text: "查询失败:" + err.Error()}
+ }
+ if len(rows) == 0 {
+ return telegramCommandReply{Text: "没有匹配记录。"}
+ }
+ var out []string
+ for i, r := range rows {
+ out = append(out, fmt.Sprintf("%d. %s / %s / %s / %s / %s", i+1, r.Username, blankDash(r.LastIP), blankDash(r.DeviceName), blankDash(r.Client), r.LastSeenAt.Format("2006-01-02 15:04")))
+ }
+ return telegramCommandReply{Text: "审计结果\n\n" + strings.Join(out, "\n") + ""}
+}
diff --git a/internal/service/telegram_polling.go b/internal/service/telegram_polling.go
new file mode 100644
index 0000000..399f1b3
--- /dev/null
+++ b/internal/service/telegram_polling.go
@@ -0,0 +1,208 @@
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// TelegramPollingStartResult describes what happened when local long polling
+// was requested. The admin UI uses it to avoid a silent "started" toast when
+// no Telegram channel can actually poll.
+type TelegramPollingStartResult struct {
+ Message string `json:"message"`
+ Started int `json:"started"`
+ AlreadyRunning int `json:"already_running"`
+ Skipped int `json:"skipped"`
+ Errors []string `json:"errors,omitempty"`
+}
+
+// StartPolling 为所有已启用的 Telegram 通知渠道启动长轮询。
+func (s *TelegramBotService) StartPolling(ctx context.Context) TelegramPollingStartResult {
+ result := TelegramPollingStartResult{Message: "telegram polling started"}
+ channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
+ if err != nil {
+ s.log.Error("failed to list telegram channels for polling", zap.Error(err))
+ result.Message = "failed to list telegram channels"
+ result.Errors = append(result.Errors, err.Error())
+ return result
+ }
+ if len(channels) == 0 {
+ result.Message = "no telegram channels configured"
+ result.Errors = append(result.Errors, "没有配置 Telegram 通知渠道")
+ return result
+ }
+
+ for _, ch := range channels {
+ if !ch.Enabled {
+ result.Skipped++
+ result.Errors = append(result.Errors, ch.Name+": 通知渠道未启用")
+ continue
+ }
+ configStr := ch.Config
+ if s.crypto != nil && configStr != "" {
+ configStr = s.crypto.Decrypt(configStr)
+ }
+ var rawCfg map[string]any
+ if err := json.Unmarshal([]byte(configStr), &rawCfg); err != nil {
+ result.Skipped++
+ result.Errors = append(result.Errors, ch.Name+": Telegram 配置解析失败: "+err.Error())
+ continue
+ }
+ cfg := telegramStringConfigFromAny(rawCfg)
+ botToken := cfg["bot_token"]
+ if botToken == "" {
+ result.Skipped++
+ result.Errors = append(result.Errors, ch.Name+": Telegram Bot Token 为空")
+ continue
+ }
+ s.pollingMu.Lock()
+ if _, running := s.pollingCancel[botToken]; running {
+ s.pollingMu.Unlock()
+ result.AlreadyRunning++
+ continue
+ }
+ s.pollingMu.Unlock()
+
+ if err := registerTelegramBotCommands(ctx, cfg); err != nil && s.log != nil {
+ s.log.Warn("telegram setMyCommands failed", zap.Error(sanitizeTelegramError(err)))
+ }
+ if err := deleteTelegramWebhook(ctx, cfg); err != nil {
+ result.Skipped++
+ result.Errors = append(result.Errors, ch.Name+": "+sanitizeTelegramError(err).Error())
+ continue
+ }
+
+ s.pollingMu.Lock()
+ if _, running := s.pollingCancel[botToken]; running {
+ s.pollingMu.Unlock()
+ result.AlreadyRunning++
+ continue
+ }
+ pollCtx, cancel := context.WithCancel(context.Background())
+ s.pollingCancel[botToken] = cancel
+ s.pollingMu.Unlock()
+
+ channel := ch
+ go s.pollLoop(pollCtx, cfg, &channel)
+ result.Started++
+ s.log.Info("started telegram polling", zap.String("channel", ch.Name))
+ }
+ if result.Started == 0 && result.AlreadyRunning == 0 {
+ result.Message = "no enabled telegram channels started"
+ }
+ return result
+}
+
+// StopPolling 停止所有 Telegram 长轮询。
+func (s *TelegramBotService) StopPolling() int {
+ s.pollingMu.Lock()
+ defer s.pollingMu.Unlock()
+ stopped := 0
+ for token, cancel := range s.pollingCancel {
+ cancel()
+ delete(s.pollingCancel, token)
+ stopped++
+ }
+ s.log.Info("telegram polling stopped")
+ return stopped
+}
+
+// pollLoop 对单个 Bot Token 执行长轮询。
+func (s *TelegramBotService) pollLoop(ctx context.Context, cfg map[string]string, channel *model.NotifyChannel) {
+ var offset int64 = 0
+ pollURL, err := telegramMethodURL(cfg, cfg["bot_token"], "getUpdates")
+ if err != nil {
+ s.log.Warn("telegram polling config invalid", zap.Error(err))
+ return
+ }
+ clients := telegramHTTPClients(45*time.Second, cfg)
+
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ default:
+ }
+
+ reqBody, _ := json.Marshal(map[string]interface{}{
+ "offset": offset,
+ "timeout": 30,
+ })
+ respBody, err := telegramPollingRequest(ctx, clients, pollURL, string(reqBody))
+ if err != nil {
+ s.log.Debug("telegram polling failed", zap.Error(err))
+ time.Sleep(5 * time.Second)
+ continue
+ }
+
+ var result struct {
+ OK bool `json:"ok"`
+ Result []TelegramUpdate `json:"result"`
+ }
+ if err := json.Unmarshal(respBody, &result); err != nil || !result.OK {
+ time.Sleep(3 * time.Second)
+ continue
+ }
+
+ for _, upd := range result.Result {
+ if upd.UpdateID >= int(offset) {
+ offset = int64(upd.UpdateID) + 1
+ }
+ if !telegramUpdateActionable(upd) {
+ continue
+ }
+ go func(u TelegramUpdate) {
+ handlerCtx, cancel := context.WithTimeout(ctx, 2*time.Minute)
+ defer cancel()
+ _ = s.handleTelegramUpdate(handlerCtx, u, channel)
+ }(upd)
+ }
+ }
+}
+
+// telegramUpdateActionable 判断一条 update 是否需要分发处理。
+// 长轮询默认会返回 message 与 callback_query 两类更新;命令消息需有文本,
+// 而内联按钮回调(callback_query)必须被分发,否则成人目录显隐开关会失效。
+func telegramUpdateActionable(upd TelegramUpdate) bool {
+ if upd.CallbackQuery != nil {
+ return true
+ }
+ return upd.Message != nil && upd.Message.Text != ""
+}
+
+func telegramPollingRequest(ctx context.Context, clients []*http.Client, pollURL, body string) ([]byte, error) {
+ var lastErr error
+ for _, client := range clients {
+ req, err := http.NewRequestWithContext(ctx, http.MethodPost, pollURL, strings.NewReader(body))
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("Content-Type", "application/json")
+ resp, err := client.Do(req)
+ if err != nil {
+ lastErr = sanitizeTelegramError(err)
+ continue
+ }
+ respBody, _ := io.ReadAll(resp.Body)
+ _ = resp.Body.Close()
+ if resp.StatusCode >= 400 {
+ lastErr = fmt.Errorf("telegram api error %d: %s", resp.StatusCode, sanitizeTelegramText(string(respBody)))
+ continue
+ }
+ return respBody, nil
+ }
+ if lastErr != nil {
+ return nil, lastErr
+ }
+ return nil, errors.New("telegram polling failed")
+}
diff --git a/internal/service/telegram_redeem.go b/internal/service/telegram_redeem.go
new file mode 100644
index 0000000..dc9862f
--- /dev/null
+++ b/internal/service/telegram_redeem.go
@@ -0,0 +1,185 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+ "gorm.io/gorm"
+)
+
+var (
+ errRegistrationCodeAlreadyUsed = errors.New("registration code already used")
+ errRegistrationCodeExpired = errors.New("registration code expired")
+)
+
+func (s *TelegramBotService) cmdRedeem(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "请发送:/redeem 兑换码\n未绑定账号时自动尝试注册码;已绑定账号时自动尝试续期码。"}
+ }
+ code := strings.Join(args, " ")
+ if s.boundUser(ctx, msg.From.ID) == nil {
+ return s.redeemRegisterFlow(ctx, channel, msg, code)
+ }
+ return s.redeemRenewFlow(ctx, msg, code)
+}
+
+func (s *TelegramBotService) cmdRedeemRegister(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "请发送:/redeem_register 注册兑换码"}
+ }
+ return s.redeemRegisterFlow(ctx, channel, msg, strings.Join(args, " "))
+}
+
+func (s *TelegramBotService) cmdRedeemRenew(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "请发送:/redeem_renew 续期兑换码"}
+ }
+ return s.redeemRenewFlow(ctx, msg, strings.Join(args, " "))
+}
+
+func (s *TelegramBotService) redeemRegisterFlow(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, raw string) telegramCommandReply {
+ if channel == nil {
+ channel = s.findChannelForMessage(ctx, msg)
+ }
+ if dec := s.telegramUserBindDecision(ctx, channel, msg.From.ID); dec != bindAllowed {
+ return telegramCommandReply{Text: telegramBindRejectText(dec, "兑换注册账号")}
+ }
+ rc, errMsg := s.lookupRedeemableCode(ctx, raw, model.RegistrationCodeRegister)
+ if rc == nil {
+ return telegramCommandReply{Text: errMsg}
+ }
+ if s.auth == nil {
+ return telegramCommandReply{Text: "注册服务暂不可用。"}
+ }
+ if binding := s.telegramBinding(ctx, msg.From.ID); binding != nil {
+ if u, _ := s.repo.User.FindByID(ctx, binding.UserID); u != nil {
+ return telegramCommandReply{Text: fmt.Sprintf("当前 Telegram 已绑定账号 %s,无需再用注册码。", u.Username)}
+ }
+ }
+ user, password, claimedCode, err := s.createUserFromRegistrationCode(ctx, rc.Code)
+ if err != nil {
+ if errors.Is(err, errRegistrationCodeAlreadyUsed) {
+ return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
+ }
+ if errors.Is(err, errRegistrationCodeExpired) {
+ return telegramCommandReply{Text: "兑换码已过期。"}
+ }
+ if errors.Is(err, ErrUserLimitReached) {
+ return telegramCommandReply{Text: "注册失败:用户数量已达授权上限。"}
+ }
+ return telegramCommandReply{Text: "注册失败:" + err.Error()}
+ }
+ if claimedCode == nil {
+ return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
+ }
+ _ = s.upsertTelegramBinding(ctx, msg, user.ID)
+ return telegramCommandReply{
+ Text: fmt.Sprintf("兑换成功并已创建账号:\n用户名:%s\n密码:%s\n到期:%s\n\n请尽快用「改用户名/改密码」修改为你自己的凭据。",
+ user.Username, password, formatExpiry(s.userExpiry(ctx, user.ID))),
+ Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}},
+ }
+}
+
+func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context, rawCode string) (*model.User, string, *model.RegistrationCode, error) {
+ code := normalizeRedemptionCode(rawCode)
+ if code == "" {
+ return nil, "", nil, errRegistrationCodeAlreadyUsed
+ }
+ password := randomCode(10)
+ var created model.User
+ var claimed model.RegistrationCode
+ err := s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
+ if err := tx.Where("code = ? AND kind = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", code, model.RegistrationCodeRegister).
+ First(&claimed).Error; err != nil {
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return errRegistrationCodeAlreadyUsed
+ }
+ return err
+ }
+ if claimed.IsExpired() {
+ return errRegistrationCodeExpired
+ }
+ var count int64
+ if err := tx.Model(&model.User{}).Count(&count).Error; err != nil {
+ return err
+ }
+ if count >= LicensedMaxUsers(ctx, s.repo) {
+ return ErrUserLimitReached
+ }
+ hash, err := hashPassword(password)
+ if err != nil {
+ return err
+ }
+ codePrefix := strings.ToLower(claimed.Code)
+ if len(codePrefix) > 8 {
+ codePrefix = codePrefix[:8]
+ }
+ created = model.User{
+ Username: "u" + codePrefix,
+ PasswordHash: hash,
+ Role: "user",
+ Tier: "free",
+ HideAdult: true,
+ ExpiredAt: renewExpiry(nil, claimed.DurationDays),
+ }
+ if err := tx.Create(&created).Error; err != nil {
+ return err
+ }
+ if err := tx.Create(DefaultPermissions(created.ID)).Error; err != nil {
+ return err
+ }
+ now := time.Now()
+ res := tx.Model(&model.RegistrationCode{}).
+ Where("id = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", claimed.ID).
+ Updates(map[string]any{
+ "used_by_user_id": created.ID,
+ "used_count": gorm.Expr("used_count + 1"),
+ "used_at": gorm.Expr("CASE WHEN used_count + 1 >= CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END THEN ? ELSE used_at END", now),
+ })
+ if res.Error != nil {
+ return res.Error
+ }
+ if res.RowsAffected == 0 {
+ return errRegistrationCodeAlreadyUsed
+ }
+ claimed.UsedByUserID = created.ID
+ claimed.UsedCount++
+ if claimed.UsedCount >= claimed.EffectiveMaxUses() {
+ claimed.UsedAt = &now
+ }
+ return nil
+ })
+ if err != nil {
+ return nil, "", nil, err
+ }
+ return &created, password, &claimed, nil
+}
+
+func (s *TelegramBotService) redeemRenewFlow(ctx context.Context, msg *TelegramMessage, raw string) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号再续期。"}
+ }
+ rc, errMsg := s.lookupRedeemableCode(ctx, raw, model.RegistrationCodeRenew)
+ if rc == nil {
+ return telegramCommandReply{Text: errMsg}
+ }
+ if err := s.repo.RegCode.MarkUsed(ctx, rc.ID, user.ID); err != nil {
+ return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
+ }
+ if err := s.applyRenewal(ctx, user.ID, rc.DurationDays); err != nil {
+ return telegramCommandReply{Text: "续期失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("续期成功 ✅ 当前到期:%s", formatExpiry(s.userExpiry(ctx, user.ID)))}
+}
+
+func (s *TelegramBotService) userExpiry(ctx context.Context, userID string) *time.Time {
+ if u, _ := s.repo.User.FindByID(ctx, userID); u != nil {
+ return u.ExpiredAt
+ }
+ return nil
+}
diff --git a/internal/service/telegram_reply.go b/internal/service/telegram_reply.go
new file mode 100644
index 0000000..f710cd2
--- /dev/null
+++ b/internal/service/telegram_reply.go
@@ -0,0 +1,127 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+ "time"
+
+ "go.uber.org/zap"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+const defaultTelegramMessageDeleteDelay = 120 * time.Second
+
+type telegramSendMessageResponse struct {
+ OK bool `json:"ok"`
+ Result struct {
+ MessageID int `json:"message_id"`
+ } `json:"result"`
+}
+
+// reply 通过 Telegram Bot API 发送回复消息。
+func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error {
+ cfg := s.telegramChannelConfig(channel)
+ if strings.TrimSpace(cfg["bot_token"]) == "" {
+ return fmt.Errorf("bot_token not configured")
+ }
+
+ payload := map[string]interface{}{
+ "chat_id": strconv.Itoa(chatID),
+ "text": reply.Text,
+ "parse_mode": "HTML",
+ }
+ if len(reply.Buttons) > 0 {
+ keyboard := make([][]map[string]string, 0, len(reply.Buttons))
+ for _, row := range reply.Buttons {
+ buttons := make([]map[string]string, 0, len(row))
+ for _, button := range row {
+ buttons = append(buttons, map[string]string{
+ "text": button.Text,
+ "callback_data": button.Data,
+ })
+ }
+ keyboard = append(keyboard, buttons)
+ }
+ payload["reply_markup"] = map[string]interface{}{"inline_keyboard": keyboard}
+ }
+ var sent telegramSendMessageResponse
+ if err := telegramPostJSONDecode(ctx, cfg, "sendMessage", payload, 15*time.Second, &sent); err != nil {
+ return err
+ }
+ if sent.Result.MessageID > 0 {
+ s.scheduleTelegramMessageDelete(cfg, chatID, sent.Result.MessageID)
+ }
+ return nil
+}
+
+func (s *TelegramBotService) replyForMessage(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, reply telegramCommandReply) error {
+ if msg == nil {
+ return nil
+ }
+ if strings.TrimSpace(reply.Text) == "" {
+ return nil
+ }
+ return s.reply(ctx, channel, msg.Chat.ID, reply)
+}
+
+func (s *TelegramBotService) deleteTelegramSourceMessage(channel *model.NotifyChannel, chatID, messageID int) {
+ if messageID <= 0 {
+ return
+ }
+ s.scheduleTelegramMessageDelete(s.telegramChannelConfig(channel), chatID, messageID)
+}
+
+func (s *TelegramBotService) scheduleTelegramMessageDelete(cfg map[string]string, chatID, messageID int) {
+ if chatID == 0 || messageID <= 0 || strings.TrimSpace(cfg["bot_token"]) == "" {
+ return
+ }
+ delay := telegramMessageDeleteDelay(cfg)
+ if delay < 0 {
+ return
+ }
+ cfgCopy := make(map[string]string, len(cfg))
+ for k, v := range cfg {
+ cfgCopy[k] = v
+ }
+ go func() {
+ if delay > 0 {
+ timer := time.NewTimer(delay)
+ defer timer.Stop()
+ <-timer.C
+ }
+ deleteCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
+ defer cancel()
+ err := telegramPostJSON(deleteCtx, cfgCopy, "deleteMessage", map[string]interface{}{
+ "chat_id": strconv.Itoa(chatID),
+ "message_id": messageID,
+ }, 10*time.Second)
+ if err != nil && s.log != nil {
+ s.log.Debug("telegram deleteMessage failed",
+ zap.Int("chat_id", chatID),
+ zap.Int("message_id", messageID),
+ zap.Error(sanitizeTelegramError(err)),
+ )
+ }
+ }()
+}
+
+func telegramMessageDeleteDelay(cfg map[string]string) time.Duration {
+ for _, key := range []string{"auto_delete_seconds", "message_delete_seconds", "delete_after_seconds"} {
+ raw := strings.TrimSpace(cfg[key])
+ if raw == "" {
+ continue
+ }
+ seconds, err := strconv.Atoi(raw)
+ if err != nil {
+ continue
+ }
+ if seconds < 0 {
+ return -1
+ }
+ return time.Duration(seconds) * time.Second
+ }
+ return defaultTelegramMessageDeleteDelay
+}
diff --git a/internal/service/telegram_stats.go b/internal/service/telegram_stats.go
new file mode 100644
index 0000000..faf62b9
--- /dev/null
+++ b/internal/service/telegram_stats.go
@@ -0,0 +1,224 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strings"
+
+ "gorm.io/gorm"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+// cmdStatus 处理 /status 命令。
+func (s *TelegramBotService) cmdStatus(ctx context.Context) (telegramCommandReply, error) {
+ libraryIDs, err := s.activeTelegramStatsLibraryIDs(ctx)
+ if err != nil {
+ return telegramCommandReply{}, err
+ }
+ var mediaCount int64
+ s.mediaStatsQuery(libraryIDs).Count(&mediaCount)
+
+ var totalSize int64
+ if err := s.mediaStatsQuery(libraryIDs).Select("COALESCE(SUM(size_bytes), 0)").Row().Scan(&totalSize); err != nil {
+ return telegramCommandReply{}, err
+ }
+ totalSizeGB := float64(totalSize) / 1024 / 1024 / 1024
+
+ return telegramCommandReply{Text: fmt.Sprintf(
+ "系统运行状态\n\n"+
+ "🎬 媒体总数: %d\n"+
+ "💾 存储占用: %.1f GB",
+ mediaCount, totalSizeGB,
+ )}, nil
+}
+
+// cmdSearch 处理 /search 命令。
+func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (telegramCommandReply, error) {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "请提供搜索关键词\n例: /search 哥斯拉"}, nil
+ }
+
+ keyword := strings.Join(args, " ")
+ var results []model.Media
+ err := s.repo.DB.Where("title LIKE ?", "%"+keyword+"%").
+ Order("year DESC").Limit(8).
+ Find(&results).Error
+ if err != nil {
+ return telegramCommandReply{}, err
+ }
+
+ if len(results) == 0 {
+ return telegramCommandReply{Text: fmt.Sprintf("未找到与 %s 相关的媒体", keyword)}, nil
+ }
+
+ var sb strings.Builder
+ sb.WriteString(fmt.Sprintf("搜索: %s\n\n", keyword))
+ for i, m := range results {
+ year := ""
+ if m.Year > 0 {
+ year = fmt.Sprintf(" (%d)", m.Year)
+ }
+ ep := ""
+ if m.SeasonNum > 0 && m.EpisodeNum > 0 {
+ ep = fmt.Sprintf(" S%02dE%02d", m.SeasonNum, m.EpisodeNum)
+ }
+ sb.WriteString(fmt.Sprintf("%d. %s%s%s — %s\n", i+1, m.Title, year, ep, formatSize(m.SizeBytes)))
+ }
+
+ return telegramCommandReply{Text: sb.String()}, nil
+}
+
+// cmdDownloads 处理 /downloads 命令。
+func (s *TelegramBotService) cmdDownloads(ctx context.Context) (telegramCommandReply, error) {
+ type Row struct {
+ Title string
+ Status string
+ }
+ var rows []Row
+ if err := s.repo.DB.Raw(
+ "SELECT COALESCE(NULLIF(title,''),'下载任务') as title, COALESCE(status,'unknown') as status FROM download_tasks ORDER BY created_at DESC LIMIT 8",
+ ).Scan(&rows).Error; err != nil {
+ return telegramCommandReply{}, err
+ }
+
+ if len(rows) == 0 {
+ return telegramCommandReply{Text: "当前没有下载任务。"}, nil
+ }
+
+ var sb strings.Builder
+ sb.WriteString(fmt.Sprintf("下载任务 (%d)\n\n", len(rows)))
+ for _, r := range rows {
+ icon := "⏳"
+ switch r.Status {
+ case "completed":
+ icon = "✅"
+ case "downloading":
+ icon = "📥"
+ case "error":
+ icon = "❌"
+ }
+ name := strings.TrimSpace(r.Title)
+ if name == "" {
+ name = "下载任务"
+ }
+ if len(name) > 60 {
+ name = name[:57] + "..."
+ }
+ sb.WriteString(fmt.Sprintf("%s %s\n", icon, name))
+ }
+
+ return telegramCommandReply{Text: sb.String()}, nil
+}
+
+// cmdStats 处理 /stats 命令。
+func (s *TelegramBotService) cmdStats(ctx context.Context) (telegramCommandReply, error) {
+ libs, err := s.activeTelegramStatsLibraries(ctx)
+ if err != nil {
+ return telegramCommandReply{}, err
+ }
+ libraryIDs := make([]string, 0, len(libs))
+ for _, lib := range libs {
+ libraryIDs = append(libraryIDs, lib.ID)
+ }
+ var totalMedia int64
+ s.mediaStatsQuery(libraryIDs).Count(&totalMedia)
+
+ var totalSize int64
+ if err := s.mediaStatsQuery(libraryIDs).Select("COALESCE(SUM(size_bytes), 0)").Row().Scan(&totalSize); err != nil {
+ return telegramCommandReply{}, err
+ }
+
+ type LibStat struct {
+ Name string
+ Type string
+ Count int64
+ }
+ stats := make([]LibStat, 0, len(libs))
+ for _, lib := range libs {
+ var count int64
+ if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("library_id = ?", lib.ID).Count(&count).Error; err != nil {
+ return telegramCommandReply{}, err
+ }
+ stats = append(stats, LibStat{Name: lib.Name, Type: lib.Type, Count: count})
+ }
+
+ var sb strings.Builder
+ sb.WriteString("媒体库统计\n\n")
+ sb.WriteString(fmt.Sprintf("📚 总数: %d\n", totalMedia))
+ sb.WriteString(fmt.Sprintf("💾 大小: %s\n", formatSize(totalSize)))
+
+ if len(stats) > 0 {
+ sb.WriteString("\n各库分布:\n")
+ for _, l := range stats {
+ icon := "🎬"
+ switch l.Type {
+ case "tv":
+ icon = "📺"
+ case "anime":
+ icon = "🍥"
+ case "music":
+ icon = "🎵"
+ }
+ sb.WriteString(fmt.Sprintf("%s %s: %d\n", icon, l.Name, l.Count))
+ }
+ }
+
+ return telegramCommandReply{Text: sb.String()}, nil
+}
+
+func (s *TelegramBotService) activeTelegramStatsLibraries(ctx context.Context) ([]model.Library, error) {
+ if s == nil || s.repo == nil || s.repo.Library == nil {
+ return nil, nil
+ }
+ libs, err := s.repo.Library.List(ctx)
+ if err != nil {
+ return nil, err
+ }
+ libs = FilterDisplayCloudLibraries(ctx, s.repo, libs)
+ out := libs[:0]
+ for _, lib := range libs {
+ if lib.Enabled {
+ out = append(out, lib)
+ }
+ }
+ return out, nil
+}
+
+func (s *TelegramBotService) activeTelegramStatsLibraryIDs(ctx context.Context) ([]string, error) {
+ libs, err := s.activeTelegramStatsLibraries(ctx)
+ if err != nil {
+ return nil, err
+ }
+ ids := make([]string, 0, len(libs))
+ for _, lib := range libs {
+ ids = append(ids, lib.ID)
+ }
+ return ids, nil
+}
+
+func (s *TelegramBotService) mediaStatsQuery(libraryIDs []string) *gorm.DB {
+ q := s.repo.DB.Model(&model.Media{})
+ if len(libraryIDs) == 0 {
+ return q.Where("1 = 0")
+ }
+ return q.Where("library_id IN ?", libraryIDs)
+}
+
+// formatSize 格式化字节数为可读字符串。
+func formatSize(bytes int64) string {
+ if bytes <= 0 {
+ return "0 B"
+ }
+ units := []string{"B", "KB", "MB", "GB", "TB"}
+ v := float64(bytes)
+ i := 0
+ for v >= 1024 && i < len(units)-1 {
+ v /= 1024
+ i++
+ }
+ if i == 0 {
+ return fmt.Sprintf("%.0f %s", v, units[i])
+ }
+ return fmt.Sprintf("%.1f %s", v, units[i])
+}
diff --git a/internal/service/telegram_unbind.go b/internal/service/telegram_unbind.go
new file mode 100644
index 0000000..41e067c
--- /dev/null
+++ b/internal/service/telegram_unbind.go
@@ -0,0 +1,220 @@
+package service
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *TelegramBotService) cmdUnbind(ctx context.Context, args []string) telegramCommandReply {
+ targets := parseTelegramUnbindTargets(args)
+ if len(targets) == 0 {
+ return telegramCommandReply{Text: "用法:/unbind 用户名1 用户名2\n也支持逗号分隔,或使用 tg:TelegramID 按 Telegram ID 解绑。此命令只解绑 Bot,不删除媒体账号。"}
+ }
+ var removed int64
+ var done []string
+ var skipped []string
+ var missing []string
+ for _, target := range targets {
+ if tgIDRaw, ok := strings.CutPrefix(strings.ToLower(target), "tg:"); ok {
+ tgID, err := strconv.ParseInt(tgIDRaw, 10, 64)
+ if err != nil || tgID == 0 {
+ missing = append(missing, target)
+ continue
+ }
+ n, err := s.deleteTelegramBindings(ctx, "telegram_user_id = ?", tgID)
+ if err != nil {
+ return telegramCommandReply{Text: "解绑失败:" + err.Error()}
+ }
+ if n == 0 {
+ missing = append(missing, target)
+ continue
+ }
+ removed += n
+ done = append(done, target)
+ continue
+ }
+
+ user, _ := s.repo.User.FindByUsername(ctx, target)
+ if user == nil {
+ user, _ = s.repo.User.FindByID(ctx, target)
+ }
+ if user == nil {
+ missing = append(missing, target)
+ continue
+ }
+ if user.Role == "admin" {
+ skipped = append(skipped, user.Username+"(管理员)")
+ continue
+ }
+ n, err := s.deleteTelegramBindings(ctx, "user_id = ?", user.ID)
+ if err != nil {
+ return telegramCommandReply{Text: "解绑失败:" + err.Error()}
+ }
+ if n == 0 {
+ missing = append(missing, user.Username+"(未绑定)")
+ continue
+ }
+ removed += n
+ done = append(done, user.Username)
+ }
+ return formatUnbindResult("批量解绑完成", removed, done, skipped, missing)
+}
+
+func (s *TelegramBotService) cmdUnbindDuplicates(ctx context.Context) telegramCommandReply {
+ if s == nil || s.repo == nil || s.repo.DB == nil {
+ return telegramCommandReply{Text: "仓库不可用。"}
+ }
+ var bindings []model.TelegramBinding
+ if err := s.repo.DB.WithContext(ctx).Order("updated_at desc, created_at desc").Find(&bindings).Error; err != nil {
+ return telegramCommandReply{Text: "读取绑定失败:" + err.Error()}
+ }
+ seenTelegram := make(map[int64]string)
+ seenUser := make(map[string]string)
+ var removeIDs []string
+ var removedLabels []string
+ for _, binding := range bindings {
+ remove := false
+ if binding.UserID == "" || binding.TelegramUserID == 0 {
+ remove = true
+ } else if user, _ := s.repo.User.FindByID(ctx, binding.UserID); user == nil {
+ remove = true
+ } else if _, ok := seenTelegram[binding.TelegramUserID]; ok {
+ remove = true
+ } else if _, ok := seenUser[binding.UserID]; ok {
+ remove = true
+ }
+ if remove {
+ removeIDs = append(removeIDs, binding.ID)
+ removedLabels = append(removedLabels, fmt.Sprintf("tg:%d", binding.TelegramUserID))
+ continue
+ }
+ seenTelegram[binding.TelegramUserID] = binding.ID
+ seenUser[binding.UserID] = binding.ID
+ }
+ if len(removeIDs) == 0 {
+ return telegramCommandReply{Text: "未发现重复或无效绑定。"}
+ }
+ n, err := s.deleteTelegramBindings(ctx, "id IN ?", removeIDs)
+ if err != nil {
+ return telegramCommandReply{Text: "清理失败:" + err.Error()}
+ }
+ return formatUnbindResult("重复/无效绑定清理完成", n, removedLabels, nil, nil)
+}
+
+func (s *TelegramBotService) cmdUnbindInactive(ctx context.Context, args []string) telegramCommandReply {
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "用法:/unbind_inactive 天数\n例如 /unbind_inactive 30 会解绑 30 天未登录的普通用户 Bot 绑定,不删除账号。"}
+ }
+ days, err := strconv.Atoi(strings.TrimSpace(args[0]))
+ if err != nil || days < 1 {
+ return telegramCommandReply{Text: "天数必须是大于 0 的整数。"}
+ }
+ users, err := s.repo.User.List(ctx)
+ if err != nil {
+ return telegramCommandReply{Text: "读取用户失败:" + err.Error()}
+ }
+ cutoff := time.Now().Add(-time.Duration(days) * 24 * time.Hour)
+ var userIDs []string
+ var done []string
+ for _, user := range users {
+ if user.Role == "admin" {
+ continue
+ }
+ lastActive := user.CreatedAt
+ if user.LastLoginAt != nil {
+ lastActive = *user.LastLoginAt
+ }
+ if lastActive.IsZero() || lastActive.After(cutoff) {
+ continue
+ }
+ var count int64
+ _ = s.repo.DB.WithContext(ctx).Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error
+ if count == 0 {
+ continue
+ }
+ userIDs = append(userIDs, user.ID)
+ done = append(done, user.Username)
+ }
+ if len(userIDs) == 0 {
+ return telegramCommandReply{Text: fmt.Sprintf("未发现 %d 天未登录且已绑定 Bot 的普通用户。", days)}
+ }
+ n, err := s.deleteTelegramBindings(ctx, "user_id IN ?", userIDs)
+ if err != nil {
+ return telegramCommandReply{Text: "解绑失败:" + err.Error()}
+ }
+ return formatUnbindResult(fmt.Sprintf("已解绑 %d 天未登录用户", days), n, done, nil, nil)
+}
+
+func parseTelegramUnbindTargets(args []string) []string {
+ seen := make(map[string]struct{})
+ var targets []string
+ for _, arg := range args {
+ for _, part := range strings.FieldsFunc(arg, func(r rune) bool {
+ return r == ',' || r == ',' || r == ';' || r == ';' || r == '\n' || r == '\t'
+ }) {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ continue
+ }
+ key := strings.ToLower(part)
+ if _, ok := seen[key]; ok {
+ continue
+ }
+ seen[key] = struct{}{}
+ targets = append(targets, part)
+ }
+ }
+ return targets
+}
+
+func (s *TelegramBotService) deleteTelegramBindings(ctx context.Context, query string, args ...interface{}) (int64, error) {
+ if s == nil || s.repo == nil || s.repo.DB == nil {
+ return 0, nil
+ }
+ tx := s.repo.DB.WithContext(ctx).Unscoped().Where(query, args...).Delete(&model.TelegramBinding{})
+ return tx.RowsAffected, tx.Error
+}
+
+func formatUnbindResult(title string, removed int64, done, skipped, missing []string) telegramCommandReply {
+ var sb strings.Builder
+ sb.WriteString("")
+ sb.WriteString(title)
+ sb.WriteString("\n\n")
+ sb.WriteString(fmt.Sprintf("已解绑:%d 条绑定", removed))
+ if len(done) > 0 {
+ sb.WriteString("\n目标:")
+ sb.WriteString(formatShortList(done, 12))
+ }
+ if len(skipped) > 0 {
+ sb.WriteString("\n跳过:")
+ sb.WriteString(formatShortList(skipped, 8))
+ }
+ if len(missing) > 0 {
+ sb.WriteString("\n未找到/未绑定:")
+ sb.WriteString(formatShortList(missing, 8))
+ }
+ return telegramCommandReply{Text: sb.String()}
+}
+
+func formatShortList(items []string, limit int) string {
+ if len(items) == 0 {
+ return ""
+ }
+ if limit < 1 {
+ limit = 1
+ }
+ out := items
+ if len(out) > limit {
+ out = out[:limit]
+ }
+ text := "" + strings.Join(out, "、") + ""
+ if len(items) > limit {
+ text += fmt.Sprintf(" 等 %d 项", len(items))
+ }
+ return text
+}
diff --git a/internal/service/telegram_user_self.go b/internal/service/telegram_user_self.go
new file mode 100644
index 0000000..61c5c3f
--- /dev/null
+++ b/internal/service/telegram_user_self.go
@@ -0,0 +1,225 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strconv"
+ "strings"
+
+ "github.com/ShukeBta/MediaStationGo/internal/model"
+)
+
+func (s *TelegramBotService) cmdKick(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号:/start 用户名 密码"}
+ }
+ if len(args) == 0 {
+ return telegramCommandReply{Text: "请指定要踢下线的设备:/kick all 或 /kick 设备编号。先用 /devices 查看编号。"}
+ }
+ target := strings.TrimSpace(args[0])
+ if strings.EqualFold(target, "all") || target == "全部" {
+ if s.device != nil {
+ if err := s.device.KickAllDevices(ctx, user.ID); err != nil {
+ return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
+ }
+ } else if err := s.repo.UserDevice.SetKickedByUser(ctx, user.ID, true); err != nil {
+ return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: "已踢下线此账号的全部设备。"}
+ }
+ devices, _ := s.listUserDevices(ctx, user.ID)
+ if len(devices) == 0 {
+ return telegramCommandReply{Text: "当前没有记录到登录设备。"}
+ }
+ var chosen *model.UserDevice
+ if n, err := strconv.Atoi(target); err == nil && n >= 1 && n <= len(devices) {
+ chosen = &devices[n-1]
+ } else {
+ for i := range devices {
+ if devices[i].ID == target || devices[i].DeviceID == target {
+ chosen = &devices[i]
+ break
+ }
+ }
+ }
+ if chosen == nil {
+ return telegramCommandReply{Text: "未找到该设备。请用 /devices 查看设备编号后重试。"}
+ }
+ if err := s.repo.UserDevice.SetKicked(ctx, chosen.ID, true); err != nil {
+ return telegramCommandReply{Text: "踢下线失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("已踢下线:%s。", deviceLabel(chosen.DeviceName, chosen.Client))}
+}
+
+func (s *TelegramBotService) cmdSetName(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "请发送:/setname 当前密码 新用户名"}
+ }
+ return s.selfSetName(ctx, msg, strings.Join(args, " "))
+}
+
+func (s *TelegramBotService) cmdSetPass(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
+ if len(args) < 2 {
+ return telegramCommandReply{Text: "请发送:/setpass 当前密码 新密码"}
+ }
+ return s.selfSetPass(ctx, msg, strings.Join(args, " "))
+}
+
+func (s *TelegramBotService) replyAccount(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号:/start 用户名 密码"}
+ }
+ streak := 0
+ if rec, _ := s.repo.SignIn.Get(ctx, user.ID); rec != nil {
+ streak = rec.StreakDays
+ }
+ devices, _ := s.listUserDevices(ctx, user.ID)
+ text := fmt.Sprintf("我的账号\n\n用户名:%s\n状态:%s\n到期:%s\n连续签到:%d 天\n登录设备:%d 台",
+ user.Username,
+ map[bool]string{true: "正常", false: "已禁用"}[user.IsActive],
+ formatExpiry(user.ExpiredAt), streak, len(devices))
+ return telegramCommandReply{Text: text, Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}}
+}
+
+func (s *TelegramBotService) replySignIn(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号后再签到。"}
+ }
+ res, err := s.signIn(ctx, user.ID)
+ if err != nil {
+ return telegramCommandReply{Text: "签到失败:" + err.Error()}
+ }
+ if res.AlreadySigned {
+ return telegramCommandReply{Text: fmt.Sprintf("今天已经签到过啦~\n连续签到 %d 天,累计 %d 天。", res.Streak, res.Total)}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("签到成功 ✅\n连续签到 %d 天,累计 %d 天。", res.Streak, res.Total)}
+}
+
+func (s *TelegramBotService) replyDevices(ctx context.Context, msg *TelegramMessage) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号。"}
+ }
+ devices, _ := s.listUserDevices(ctx, user.ID)
+ if len(devices) == 0 {
+ return telegramCommandReply{Text: "当前没有记录到登录设备。"}
+ }
+ var sb strings.Builder
+ sb.WriteString("我的登录设备\n点击下方按钮可一键踢下线:\n")
+ var rows [][]telegramInlineButton
+ for i, d := range devices {
+ status := ""
+ if d.Kicked {
+ status = "(已踢下线)"
+ } else if d.Playing {
+ status = "(播放中)"
+ } else if d.Online {
+ status = "(在线)"
+ }
+ sb.WriteString(fmt.Sprintf("\n%d. %s%s\n 最近活跃:%s", i+1, deviceLabel(d.DeviceName, d.Client), status, d.LastSeenAt.Format("01-02 15:04")))
+ if !d.Kicked && !strings.HasPrefix(d.ID, "rt:") {
+ rows = append(rows, []telegramInlineButton{{Text: "🚫 踢下线:" + deviceLabel(d.DeviceName, d.Client), Data: "kick:" + d.ID}})
+ }
+ }
+ rows = append(rows, []telegramInlineButton{{Text: "⬅️ 返回菜单", Data: "menu_main"}})
+ return telegramCommandReply{Text: sb.String(), Buttons: rows}
+}
+
+func (s *TelegramBotService) replyKick(ctx context.Context, msg *TelegramMessage, deviceRowID string) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号。"}
+ }
+ var d model.UserDevice
+ if err := s.repo.DB.WithContext(ctx).Where("id = ? AND user_id = ?", deviceRowID, user.ID).First(&d).Error; err != nil {
+ return telegramCommandReply{Text: "未找到该设备。"}
+ }
+ if err := s.repo.UserDevice.SetKicked(ctx, d.ID, true); err != nil {
+ return telegramCommandReply{Text: "操作失败:" + err.Error()}
+ }
+ return s.replyDevices(ctx, msg)
+}
+
+func (s *TelegramBotService) listUserDevices(ctx context.Context, userID string) ([]model.UserDevice, error) {
+ if s.device != nil {
+ return s.device.ListDevices(ctx, userID)
+ }
+ return s.repo.UserDevice.ListByUser(ctx, userID)
+}
+
+func (s *TelegramBotService) selfSetName(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号。"}
+ }
+ currentPassword, newName := splitCurrentPasswordAndValue(input)
+ if currentPassword == "" || newName == "" {
+ return telegramCommandReply{Text: "请发送:当前密码 新用户名。"}
+ }
+ newName = strings.TrimSpace(newName)
+ if len(newName) < 2 || strings.ContainsAny(newName, " \t\n") {
+ return telegramCommandReply{Text: "用户名至少 2 位且不能含空格,请重试。"}
+ }
+ if reply, ok := s.verifyTelegramSelfPassword(ctx, msg, user, currentPassword); !ok {
+ return reply
+ }
+ if existing, _ := s.repo.User.FindByUsername(ctx, newName); existing != nil && existing.ID != user.ID {
+ return telegramCommandReply{Text: "该用户名已被占用,请换一个。"}
+ }
+ if err := s.repo.User.UpdateFields(ctx, user.ID, map[string]any{"username": newName}); err != nil {
+ return telegramCommandReply{Text: "修改失败:" + err.Error()}
+ }
+ return telegramCommandReply{Text: fmt.Sprintf("用户名已修改为 %s。请用新用户名登录。", newName)}
+}
+
+func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply {
+ user := s.boundUser(ctx, msg.From.ID)
+ if user == nil {
+ return telegramCommandReply{Text: "请先绑定账号。"}
+ }
+ currentPassword, newPass := splitCurrentPasswordAndValue(input)
+ if currentPassword == "" || newPass == "" {
+ return telegramCommandReply{Text: "请发送:当前密码 新密码。"}
+ }
+ newPass = strings.TrimSpace(newPass)
+ if s.auth == nil {
+ return telegramCommandReply{Text: "服务暂不可用。"}
+ }
+ if err := s.auth.ChangePassword(ctx, user.ID, currentPassword, newPass); err != nil {
+ if errors.Is(err, ErrInvalidCredentials) {
+ _ = s.unbindTelegramUser(ctx, msg.From.ID)
+ return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"}
+ }
+ return telegramCommandReply{Text: "修改失败:" + err.Error()}
+ }
+ if s.device != nil {
+ _ = s.device.KickAllDevices(ctx, user.ID)
+ }
+ return telegramCommandReply{Text: "密码已修改,请用新密码重新登录第三方客户端。"}
+}
+
+func splitCurrentPasswordAndValue(input string) (string, string) {
+ fields := strings.Fields(strings.TrimSpace(input))
+ if len(fields) < 2 {
+ return "", ""
+ }
+ return fields[0], strings.TrimSpace(strings.Join(fields[1:], " "))
+}
+
+func (s *TelegramBotService) verifyTelegramSelfPassword(ctx context.Context, msg *TelegramMessage, user *model.User, currentPassword string) (telegramCommandReply, bool) {
+ if s.auth == nil {
+ return telegramCommandReply{Text: "服务暂不可用。"}, false
+ }
+ if err := s.auth.VerifyPassword(ctx, user.ID, currentPassword); err != nil {
+ if errors.Is(err, ErrInvalidCredentials) {
+ _ = s.unbindTelegramUser(ctx, msg.From.ID)
+ return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"}, false
+ }
+ return telegramCommandReply{Text: "验证失败:" + err.Error()}, false
+ }
+ return telegramCommandReply{}, true
+}
diff --git a/internal/service/test_db_test.go b/internal/service/test_db_test.go
new file mode 100644
index 0000000..f7e082f
--- /dev/null
+++ b/internal/service/test_db_test.go
@@ -0,0 +1,26 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/glebarez/sqlite"
+ "gorm.io/gorm"
+ "gorm.io/gorm/logger"
+)
+
+func newServiceTestDB(t *testing.T, models ...any) *gorm.DB {
+ t.Helper()
+ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if sqlDB, err := db.DB(); err == nil {
+ t.Cleanup(func() { _ = sqlDB.Close() })
+ }
+ if len(models) > 0 {
+ if err := db.AutoMigrate(models...); err != nil {
+ t.Fatal(err)
+ }
+ }
+ return db
+}
diff --git a/internal/service/tmdb.go b/internal/service/tmdb.go
index 00d339d..40b9035 100644
--- a/internal/service/tmdb.go
+++ b/internal/service/tmdb.go
@@ -18,7 +18,6 @@ package service
import (
"context"
"encoding/json"
- "errors"
"fmt"
"net/http"
"net/url"
@@ -142,184 +141,6 @@ type Match struct {
NSFW bool `json:"nsfw,omitempty"`
}
-type tmdbMovieSearchResult struct {
- ID int `json:"id"`
- Title string `json:"title"`
- OriginalTitle string `json:"original_title"`
- OriginalLanguage string `json:"original_language"`
- Overview string `json:"overview"`
- PosterPath string `json:"poster_path"`
- BackdropPath string `json:"backdrop_path"`
- ReleaseDate string `json:"release_date"`
- VoteAverage float32 `json:"vote_average"`
- GenreIDs []int `json:"genre_ids"`
-}
-
-type tmdbTVSearchResult struct {
- ID int `json:"id"`
- Name string `json:"name"`
- OriginalName string `json:"original_name"`
- OriginalLanguage string `json:"original_language"`
- OriginCountry []string `json:"origin_country"`
- Overview string `json:"overview"`
- PosterPath string `json:"poster_path"`
- BackdropPath string `json:"backdrop_path"`
- FirstAirDate string `json:"first_air_date"`
- VoteAverage float32 `json:"vote_average"`
- GenreIDs []int `json:"genre_ids"`
-}
-
-// SearchMovie issues `/search/movie` and returns the best match, or nil
-// when no result is found. The `year` argument is optional (0 = any).
-func (t *TMDbProvider) SearchMovie(ctx context.Context, query string, year int) (*Match, error) {
- matches, err := t.SearchMovieCandidates(ctx, query, year)
- if err != nil || len(matches) == 0 {
- return nil, err
- }
- return matches[0], nil
-}
-
-// SearchMovieCandidates returns the first TMDb result page as manual-scrape
-// candidates. Automatic scrape still uses SearchMovie's first-result behavior,
-// while manual correction can show alternatives when the top result is wrong.
-func (t *TMDbProvider) SearchMovieCandidates(ctx context.Context, query string, year int) ([]*Match, error) {
- if query == "" {
- return nil, errors.New("empty query")
- }
-
- // Resolve API key from config or database
- apiKey := t.resolveAPIKey(ctx)
- if apiKey == "" {
- return nil, nil
- }
- base := t.resolveBaseURL(ctx)
-
- q := url.Values{}
- q.Set("api_key", apiKey)
- q.Set("query", query)
- q.Set("language", "zh-CN")
- q.Set("include_adult", "false")
- if year > 0 {
- q.Set("year", fmt.Sprintf("%d", year))
- }
- u := base + "/search/movie?" + q.Encode()
-
- type page struct {
- Results []tmdbMovieSearchResult `json:"results"`
- }
-
- var p page
- if err := t.getJSON(ctx, u, &p); err != nil {
- return nil, err
- }
- if len(p.Results) == 0 {
- return nil, nil
- }
- out := make([]*Match, 0, len(p.Results))
- for _, r := range p.Results {
- out = append(out, t.movieSearchResultToMatch(r))
- }
- return out, nil
-}
-
-func (t *TMDbProvider) movieSearchResultToMatch(r tmdbMovieSearchResult) *Match {
- m := &Match{
- TMDbID: r.ID,
- Title: r.Title,
- OriginalName: r.OriginalTitle,
- Overview: r.Overview,
- Rating: r.VoteAverage,
- Languages: nonEmptyStrings(r.OriginalLanguage),
- Genres: genreIDStrings(r.GenreIDs),
- }
- if r.PosterPath != "" {
- m.PosterURL = t.imgCDN + "/w500" + r.PosterPath
- }
- if r.BackdropPath != "" {
- m.BackdropURL = t.imgCDN + "/w1280" + r.BackdropPath
- }
- if len(r.ReleaseDate) >= 4 {
- _, _ = fmt.Sscanf(r.ReleaseDate[:4], "%d", &m.Year)
- }
- return m
-}
-
-// SearchTV issues `/search/tv` and returns the best match. Used by anime /
-// tv libraries before falling back to SearchMovie.
-func (t *TMDbProvider) SearchTV(ctx context.Context, query string, year int) (*Match, error) {
- matches, err := t.SearchTVCandidates(ctx, query, year)
- if err != nil || len(matches) == 0 {
- return nil, err
- }
- return matches[0], nil
-}
-
-// SearchTVCandidates returns the first TMDb TV result page for manual scrape.
-func (t *TMDbProvider) SearchTVCandidates(ctx context.Context, query string, year int) ([]*Match, error) {
- if query == "" {
- return nil, errors.New("empty query")
- }
-
- apiKey := t.resolveAPIKey(ctx)
- if apiKey == "" {
- return nil, nil
- }
- base := t.resolveBaseURL(ctx)
-
- q := url.Values{}
- q.Set("api_key", apiKey)
- q.Set("query", query)
- q.Set("language", "zh-CN")
- q.Set("include_adult", "false")
- if year > 0 {
- q.Set("first_air_date_year", fmt.Sprintf("%d", year))
- }
- u := base + "/search/tv?" + q.Encode()
-
- type page struct {
- Results []tmdbTVSearchResult `json:"results"`
- }
-
- var p page
- if err := t.getJSON(ctx, u, &p); err != nil {
- return nil, err
- }
- if len(p.Results) == 0 {
- return nil, nil
- }
- out := make([]*Match, 0, len(p.Results))
- for _, r := range p.Results {
- out = append(out, t.tvSearchResultToMatch(r))
- }
- return out, nil
-}
-
-func (t *TMDbProvider) tvSearchResultToMatch(r tmdbTVSearchResult) *Match {
- m := &Match{
- TMDbID: r.ID,
- Title: r.Name,
- OriginalName: r.OriginalName,
- Overview: r.Overview,
- Rating: r.VoteAverage,
- Languages: nonEmptyStrings(r.OriginalLanguage),
- Countries: deduplicate(r.OriginCountry),
- Genres: genreIDStrings(r.GenreIDs),
- }
- if m.Title == "" {
- m.Title = r.OriginalName
- }
- if r.PosterPath != "" {
- m.PosterURL = t.imgCDN + "/w500" + r.PosterPath
- }
- if r.BackdropPath != "" {
- m.BackdropURL = t.imgCDN + "/w1280" + r.BackdropPath
- }
- if len(r.FirstAirDate) >= 4 {
- _, _ = fmt.Sscanf(r.FirstAirDate[:4], "%d", &m.Year)
- }
- return m
-}
-
func (t *TMDbProvider) getJSON(ctx context.Context, url string, out any) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
diff --git a/internal/service/tmdb_search.go b/internal/service/tmdb_search.go
new file mode 100644
index 0000000..a59ca03
--- /dev/null
+++ b/internal/service/tmdb_search.go
@@ -0,0 +1,185 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "net/url"
+)
+
+type tmdbMovieSearchResult struct {
+ ID int `json:"id"`
+ Title string `json:"title"`
+ OriginalTitle string `json:"original_title"`
+ OriginalLanguage string `json:"original_language"`
+ Overview string `json:"overview"`
+ PosterPath string `json:"poster_path"`
+ BackdropPath string `json:"backdrop_path"`
+ ReleaseDate string `json:"release_date"`
+ VoteAverage float32 `json:"vote_average"`
+ GenreIDs []int `json:"genre_ids"`
+}
+
+type tmdbTVSearchResult struct {
+ ID int `json:"id"`
+ Name string `json:"name"`
+ OriginalName string `json:"original_name"`
+ OriginalLanguage string `json:"original_language"`
+ OriginCountry []string `json:"origin_country"`
+ Overview string `json:"overview"`
+ PosterPath string `json:"poster_path"`
+ BackdropPath string `json:"backdrop_path"`
+ FirstAirDate string `json:"first_air_date"`
+ VoteAverage float32 `json:"vote_average"`
+ GenreIDs []int `json:"genre_ids"`
+}
+
+// SearchMovie issues `/search/movie` and returns the best match, or nil
+// when no result is found. The `year` argument is optional (0 = any).
+func (t *TMDbProvider) SearchMovie(ctx context.Context, query string, year int) (*Match, error) {
+ matches, err := t.SearchMovieCandidates(ctx, query, year)
+ if err != nil || len(matches) == 0 {
+ return nil, err
+ }
+ return matches[0], nil
+}
+
+// SearchMovieCandidates returns the first TMDb result page as manual-scrape
+// candidates. Automatic scrape still uses SearchMovie's first-result behavior,
+// while manual correction can show alternatives when the top result is wrong.
+func (t *TMDbProvider) SearchMovieCandidates(ctx context.Context, query string, year int) ([]*Match, error) {
+ if query == "" {
+ return nil, errors.New("empty query")
+ }
+
+ apiKey := t.resolveAPIKey(ctx)
+ if apiKey == "" {
+ return nil, nil
+ }
+ base := t.resolveBaseURL(ctx)
+
+ q := url.Values{}
+ q.Set("api_key", apiKey)
+ q.Set("query", query)
+ q.Set("language", "zh-CN")
+ q.Set("include_adult", "false")
+ if year > 0 {
+ q.Set("year", fmt.Sprintf("%d", year))
+ }
+ u := base + "/search/movie?" + q.Encode()
+
+ type page struct {
+ Results []tmdbMovieSearchResult `json:"results"`
+ }
+
+ var p page
+ if err := t.getJSON(ctx, u, &p); err != nil {
+ return nil, err
+ }
+ if len(p.Results) == 0 {
+ return nil, nil
+ }
+ out := make([]*Match, 0, len(p.Results))
+ for _, r := range p.Results {
+ out = append(out, t.movieSearchResultToMatch(r))
+ }
+ return out, nil
+}
+
+func (t *TMDbProvider) movieSearchResultToMatch(r tmdbMovieSearchResult) *Match {
+ m := &Match{
+ TMDbID: r.ID,
+ Title: r.Title,
+ OriginalName: r.OriginalTitle,
+ Overview: r.Overview,
+ Rating: r.VoteAverage,
+ Languages: nonEmptyStrings(r.OriginalLanguage),
+ Genres: genreIDStrings(r.GenreIDs),
+ }
+ if r.PosterPath != "" {
+ m.PosterURL = t.imgCDN + "/w500" + r.PosterPath
+ }
+ if r.BackdropPath != "" {
+ m.BackdropURL = t.imgCDN + "/w1280" + r.BackdropPath
+ }
+ if len(r.ReleaseDate) >= 4 {
+ _, _ = fmt.Sscanf(r.ReleaseDate[:4], "%d", &m.Year)
+ }
+ return m
+}
+
+// SearchTV issues `/search/tv` and returns the best match. Used by anime /
+// tv libraries before falling back to SearchMovie.
+func (t *TMDbProvider) SearchTV(ctx context.Context, query string, year int) (*Match, error) {
+ matches, err := t.SearchTVCandidates(ctx, query, year)
+ if err != nil || len(matches) == 0 {
+ return nil, err
+ }
+ return matches[0], nil
+}
+
+// SearchTVCandidates returns the first TMDb TV result page for manual scrape.
+func (t *TMDbProvider) SearchTVCandidates(ctx context.Context, query string, year int) ([]*Match, error) {
+ if query == "" {
+ return nil, errors.New("empty query")
+ }
+
+ apiKey := t.resolveAPIKey(ctx)
+ if apiKey == "" {
+ return nil, nil
+ }
+ base := t.resolveBaseURL(ctx)
+
+ q := url.Values{}
+ q.Set("api_key", apiKey)
+ q.Set("query", query)
+ q.Set("language", "zh-CN")
+ q.Set("include_adult", "false")
+ if year > 0 {
+ q.Set("first_air_date_year", fmt.Sprintf("%d", year))
+ }
+ u := base + "/search/tv?" + q.Encode()
+
+ type page struct {
+ Results []tmdbTVSearchResult `json:"results"`
+ }
+
+ var p page
+ if err := t.getJSON(ctx, u, &p); err != nil {
+ return nil, err
+ }
+ if len(p.Results) == 0 {
+ return nil, nil
+ }
+ out := make([]*Match, 0, len(p.Results))
+ for _, r := range p.Results {
+ out = append(out, t.tvSearchResultToMatch(r))
+ }
+ return out, nil
+}
+
+func (t *TMDbProvider) tvSearchResultToMatch(r tmdbTVSearchResult) *Match {
+ m := &Match{
+ TMDbID: r.ID,
+ Title: r.Name,
+ OriginalName: r.OriginalName,
+ Overview: r.Overview,
+ Rating: r.VoteAverage,
+ Languages: nonEmptyStrings(r.OriginalLanguage),
+ Countries: deduplicate(r.OriginCountry),
+ Genres: genreIDStrings(r.GenreIDs),
+ }
+ if m.Title == "" {
+ m.Title = r.OriginalName
+ }
+ if r.PosterPath != "" {
+ m.PosterURL = t.imgCDN + "/w500" + r.PosterPath
+ }
+ if r.BackdropPath != "" {
+ m.BackdropURL = t.imgCDN + "/w1280" + r.BackdropPath
+ }
+ if len(r.FirstAirDate) >= 4 {
+ _, _ = fmt.Sscanf(r.FirstAirDate[:4], "%d", &m.Year)
+ }
+ return m
+}
diff --git a/internal/service/token_svc_pending_test.go b/internal/service/token_svc_pending_test.go
index 6fb2e77..d79653b 100644
--- a/internal/service/token_svc_pending_test.go
+++ b/internal/service/token_svc_pending_test.go
@@ -4,9 +4,7 @@ import (
"testing"
"time"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -15,13 +13,7 @@ import (
func newTokenTestRepo(t *testing.T) *repository.Container {
t.Helper()
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatal(err)
- }
- if err := db.AutoMigrate(&model.User{}, &model.RefreshToken{}, &model.Setting{}); err != nil {
- t.Fatal(err)
- }
+ db := newServiceTestDB(t, &model.User{}, &model.RefreshToken{}, &model.Setting{})
return repository.New(db)
}
diff --git a/internal/service/transfer.go b/internal/service/transfer.go
index a7d3c6c..07e76a2 100644
--- a/internal/service/transfer.go
+++ b/internal/service/transfer.go
@@ -101,3 +101,36 @@ func copyFile(src, dst string) error {
}
return f.Close()
}
+
+// moveFile tries os.Rename first (instant on same fs), then falls back
+// to copy + remove for cross-device moves.
+//
+// If dst already exists, moveFile returns an error instead of overwriting it.
+// OrganizeMedia checks this before calling transferFile; this remains the
+// second line of defense against different releases collapsing to one name.
+func moveFile(src, dst string) error {
+ if _, err := os.Stat(dst); err == nil {
+ return fmt.Errorf("destination already exists: %s", dst)
+ }
+ if err := os.Rename(src, dst); err == nil {
+ return nil
+ }
+ in, err := os.Open(src) // #nosec G304 -- src is selected from configured media/download roots by the organizer.
+ if err != nil {
+ return err
+ }
+ defer in.Close()
+ f, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) // #nosec G304,G302 -- dst is organizer-generated; media files must remain readable by local players.
+ if err != nil {
+ return err
+ }
+ if _, werr := io.Copy(f, in); werr != nil {
+ _ = f.Close()
+ _ = os.Remove(dst)
+ return werr
+ }
+ if cerr := f.Close(); cerr != nil {
+ return cerr
+ }
+ return os.Remove(src)
+}
diff --git a/internal/service/watcher_test.go b/internal/service/watcher_test.go
index c36d8b7..46000e1 100644
--- a/internal/service/watcher_test.go
+++ b/internal/service/watcher_test.go
@@ -6,9 +6,7 @@ import (
"testing"
"github.com/fsnotify/fsnotify"
- "github.com/glebarez/sqlite"
"go.uber.org/zap"
- "gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -25,13 +23,7 @@ func TestWatcherRefreshMapsHostLibraryPathToContainerPath(t *testing.T) {
t.Setenv("MEDIASTATION_MEDIA_DIR", hostMedia)
t.Setenv("MEDIASTATION_MEDIA_CONTAINER_DIR", containerMedia)
- db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
- if err != nil {
- t.Fatalf("open db: %v", err)
- }
- if err := db.AutoMigrate(&model.Library{}); err != nil {
- t.Fatalf("migrate: %v", err)
- }
+ db := newServiceTestDB(t, &model.Library{})
repos := repository.New(db)
lib := model.Library{
Base: model.Base{ID: "lib-tv"},
diff --git a/web/eslint.config.js b/web/eslint.config.js
index 267f86c..fde0fd4 100644
--- a/web/eslint.config.js
+++ b/web/eslint.config.js
@@ -7,6 +7,16 @@ import tseslint from 'typescript-eslint'
export default tseslint.config(
{ ignores: ['dist', 'node_modules'] },
js.configs.recommended,
+ {
+ files: ['public/**/*.js'],
+ languageOptions: {
+ ecmaVersion: 2022,
+ globals: {
+ ...globals.browser,
+ ...globals.serviceworker,
+ },
+ },
+ },
...tseslint.configs.recommended,
{
files: ['**/*.{ts,tsx}'],
diff --git a/web/public/artwork-cache-sw.js b/web/public/artwork-cache-sw.js
index 1faa511..c1dbafc 100644
--- a/web/public/artwork-cache-sw.js
+++ b/web/public/artwork-cache-sw.js
@@ -43,6 +43,31 @@ async function cacheArtwork(request) {
const contentType = response.headers.get('Content-Type') || ''
if (response.ok && contentType.toLowerCase().startsWith('image/')) {
await cache.put(cacheKey, response.clone())
+ await deleteOldArtworkVariants(cache, cacheKey)
}
return response
}
+
+async function deleteOldArtworkVariants(cache, currentRequest) {
+ const currentURL = new URL(currentRequest.url)
+ const currentIdentity = artworkIdentity(currentURL)
+ if (!currentIdentity) return
+
+ const keys = await cache.keys()
+ await Promise.all(keys.map(async (key) => {
+ if (key.url === currentRequest.url) return
+ const keyURL = new URL(key.url)
+ if (artworkIdentity(keyURL) !== currentIdentity) return
+ await cache.delete(key)
+ }))
+}
+
+function artworkIdentity(url) {
+ if (url.pathname === '/api/img') {
+ return `${url.origin}${url.pathname}?url=${url.searchParams.get('url') || ''}`
+ }
+ if (url.pathname.startsWith('/api/cloud/play/')) {
+ return `${url.origin}${url.pathname}`
+ }
+ return ''
+}
diff --git a/web/src/App.tsx b/web/src/App.tsx
index fad9dbd..707e3c1 100644
--- a/web/src/App.tsx
+++ b/web/src/App.tsx
@@ -1,108 +1,11 @@
-import { Component, Suspense, lazy, type ErrorInfo, type ReactNode } from 'react'
+import { Component, Suspense, type ErrorInfo, type ReactNode } from 'react'
import { Navigate, Route, Routes } from 'react-router-dom'
+import { appRoutes, type AppRoute } from './appRoutes'
import { Layout } from './components/Layout'
import { RequireAdmin, RequireAuth } from './components/RequireAuth'
import { LoginPage } from './pages/LoginPage'
-// Lazy-loaded routes — the login screen and the layout shell ship in the
-// initial bundle; everything else is fetched on first navigation.
-const HomePage = lazy(() => import('./pages/HomePage').then((m) => ({ default: m.HomePage })))
-const LibraryPage = lazy(() =>
- import('./pages/LibraryPage').then((m) => ({ default: m.LibraryPage })),
-)
-const LibrariesPage = lazy(() =>
- import('./pages/LibrariesPage').then((m) => ({ default: m.LibrariesPage })),
-)
-const SearchPage = lazy(() =>
- import('./pages/SearchPage').then((m) => ({ default: m.SearchPage })),
-)
-const FavouritesPage = lazy(() =>
- import('./pages/FavouritesPage').then((m) => ({ default: m.FavouritesPage })),
-)
-const PlaylistsPage = lazy(() =>
- import('./pages/PlaylistsPage').then((m) => ({ default: m.PlaylistsPage })),
-)
-const PlaylistDetailPage = lazy(() =>
- import('./pages/PlaylistDetailPage').then((m) => ({ default: m.PlaylistDetailPage })),
-)
-const MediaDetailPage = lazy(() =>
- import('./pages/MediaDetailPage').then((m) => ({ default: m.MediaDetailPage })),
-)
-const PlayerPage = lazy(() =>
- import('./pages/PlayerPage').then((m) => ({ default: m.PlayerPage })),
-)
-const AdminPage = lazy(() => import('./pages/AdminPage').then((m) => ({ default: m.AdminPage })))
-const DownloadsPage = lazy(() =>
- import('./pages/DownloadsPage').then((m) => ({ default: m.DownloadsPage })),
-)
-const SubscriptionsPage = lazy(() =>
- import('./pages/SubscriptionsPage').then((m) => ({ default: m.SubscriptionsPage })),
-)
-const ProfilePage = lazy(() =>
- import('./pages/ProfilePage').then((m) => ({ default: m.ProfilePage })),
-)
-const StatsPage = lazy(() => import('./pages/StatsPage').then((m) => ({ default: m.StatsPage })))
-const DiscoverPage = lazy(() =>
- import('./pages/DiscoverPage').then((m) => ({ default: m.DiscoverPage })),
-)
-const TasksPage = lazy(() => import('./pages/TasksPage').then((m) => ({ default: m.TasksPage })))
-const RecycleBinPage = lazy(() =>
- import('./pages/RecycleBinPage').then((m) => ({ default: m.RecycleBinPage })),
-)
-const DlnaPage = lazy(() => import('./pages/DlnaPage').then((m) => ({ default: m.DlnaPage })))
-const FileManagerPage = lazy(() =>
- import('./pages/FileManagerPage').then((m) => ({ default: m.FileManagerPage })),
-)
-const StoragePage = lazy(() =>
- import('./pages/StoragePage').then((m) => ({ default: m.StoragePage })),
-)
-const DuplicatesPage = lazy(() =>
- import('./pages/DuplicatesPage').then((m) => ({ default: m.DuplicatesPage })),
-)
-const SchedulerPage = lazy(() =>
- import('./pages/SchedulerPage').then((m) => ({ default: m.SchedulerPage })),
-)
-const WatchHistoryPage = lazy(() =>
- import('./pages/WatchHistoryPage').then((m) => ({ default: m.WatchHistoryPage })),
-)
-const PosterWallPage = lazy(() =>
- import('./pages/PosterWallPage').then((m) => ({ default: m.PosterWallPage })),
-)
-const SitesPage = lazy(() =>
- import('./pages/SitesPage').then((m) => ({ default: m.SitesPage })),
-)
-const SiteSearchPage = lazy(() =>
- import('./pages/SiteSearchPage').then((m) => ({ default: m.SiteSearchPage })),
-)
-const AIAssistantPage = lazy(() =>
- import('./pages/AIAssistantPage').then((m) => ({ default: m.AIAssistantPage })),
-)
-const StrmPage = lazy(() =>
- import('./pages/StrmPage').then((m) => ({ default: m.StrmPage })),
-)
-const ProfileManagementPage = lazy(() =>
- import('./pages/ProfileManagementPage').then((m) => ({ default: m.ProfileManagementPage })),
-)
-const NotifyChannelsPage = lazy(() =>
- import('./pages/NotifyChannelsPage').then((m) => ({ default: m.NotifyChannelsPage })),
-)
-const SettingsPage = lazy(() =>
- import('./pages/SettingsPage').then((m) => ({ default: m.SettingsPage })),
-)
-const AssistantChatPage = lazy(() =>
- import('./pages/AssistantChatPage').then((m) => ({ default: m.AssistantChatPage })),
-)
-const DownloadClientsPage = lazy(() =>
- import('./pages/DownloadClientsPage').then((m) => ({ default: m.DownloadClientsPage })),
-)
-const StorageConfigPage = lazy(() =>
- import('./pages/StorageConfigPage').then((m) => ({ default: m.StorageConfigPage })),
-)
-const LicensePage = lazy(() =>
- import('./pages/LicensePage').then((m) => ({ default: m.LicensePage })),
-)
-
const Loading = () => 加载中…
class AppErrorBoundary extends Component<{ children: ReactNode }, { hasError: boolean }> {
@@ -140,177 +43,35 @@ class AppErrorBoundary extends Component<{ children: ReactNode }, { hasError: bo
}
}
+function routeElement(route: AppRoute) {
+ if (!route.adminOnly) return route.element
+ return {route.element}
+}
+
export default function App() {
return (
}>
- } />
-
-
-
- }
- >
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
- } />
+ } />
-
-
+
+
+
}
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
-
-
- }
- />
-
- } />
+ >
+ {appRoutes.map((route) => (
+
+ ))}
+
+ } />
diff --git a/web/src/api/ai.ts b/web/src/api/ai.ts
index 978499d..38e9e3d 100644
--- a/web/src/api/ai.ts
+++ b/web/src/api/ai.ts
@@ -24,6 +24,7 @@ export interface ExternalMediaResult {
bangumi_id?: number
douban_id?: string
subscribe_keyword: string
+ subscribe_aliases?: string[]
total_episodes?: number
downloaded_episodes?: number
local_media_count?: number
diff --git a/web/src/api/apiConfig.ts b/web/src/api/apiConfig.ts
deleted file mode 100644
index 11268a2..0000000
--- a/web/src/api/apiConfig.ts
+++ /dev/null
@@ -1,58 +0,0 @@
-// API 配置 API 模块
-import { api } from './client'
-import type { ApiConfig, ApiProvider } from '../types'
-
-// 获取所有 API 配置
-export async function listApiConfigs(): Promise {
- const resp = await api.get('/api-config')
- return resp.data as unknown as ApiConfig[]
-}
-
-// 获取提供者列表
-export async function getProviders(): Promise {
- const resp = await api.get('/api-config/providers/list')
- return resp.data as unknown as ApiProvider[]
-}
-
-// 获取指定提供者的配置
-export async function getApiConfig(provider: string): Promise {
- const resp = await api.get(`/api-config/${provider}`)
- return resp.data as unknown as ApiConfig
-}
-
-// 获取生效的配置
-export async function getEffectiveConfig(provider: string): Promise {
- const resp = await api.get(`/api-config/${provider}/effective`)
- return resp.data as unknown as ApiConfig
-}
-
-// 更新 API 配置
-export interface UpdateApiConfigRequest {
- api_key?: string
- base_url?: string
- extra?: string
- enabled?: boolean
-}
-
-export async function updateApiConfig(
- provider: string,
- data: UpdateApiConfigRequest
-): Promise {
- const resp = await api.post(`/api-config/${provider}`, data)
- return resp.data as unknown as ApiConfig
-}
-
-// 删除 API 配置
-export async function deleteApiConfig(provider: string): Promise {
- await api.delete(`/api-config/${provider}`)
-}
-
-// 测试 API 连接
-export interface TestApiConfigResponse {
- result: 'success' | 'error' | 'invalid' | 'unknown'
-}
-
-export async function testApiConfig(provider: string): Promise {
- const resp = await api.post(`/api-config/${provider}/test`)
- return resp.data as unknown as TestApiConfigResponse
-}
diff --git a/web/src/api/client.ts b/web/src/api/client.ts
index 15f323c..1d50707 100644
--- a/web/src/api/client.ts
+++ b/web/src/api/client.ts
@@ -151,11 +151,13 @@ export function hlsURL(mediaId: string): string {
// imageURL converts a remote poster URL into a same-origin proxy URL so it
// can never be blocked by CORS / GFW. Empty strings pass through unchanged.
-export function imageURL(remote?: string): string {
+export function imageURL(remote?: string, version?: string): string {
if (!remote) return ''
- if (remote.startsWith('/api/img')) return remote
- if (remote.startsWith('/api/')) return withQuery(remote, tokenQuery())
- return `/api/img?url=${encodeURIComponent(remote)}&${tokenQuery()}`
+ const versionQuery = version ? `v=${encodeURIComponent(version)}` : ''
+ if (remote.startsWith('/api/img')) return withQuery(withoutAuthQuery(remote), versionQuery)
+ if (remote.startsWith('/api/cloud/play/')) return withQuery(withoutAuthQuery(remote), versionQuery)
+ if (remote.startsWith('/api/')) return withQuery(withQuery(remote, tokenQuery()), versionQuery)
+ return withQuery(`/api/img?url=${encodeURIComponent(remote)}`, versionQuery)
}
function withQuery(url: string, query: string): string {
@@ -163,6 +165,20 @@ function withQuery(url: string, query: string): string {
return `${url}${url.includes('?') ? '&' : '?'}${query}`
}
+function withoutAuthQuery(url: string): string {
+ const hashIndex = url.indexOf('#')
+ const beforeHash = hashIndex >= 0 ? url.slice(0, hashIndex) : url
+ const hash = hashIndex >= 0 ? url.slice(hashIndex) : ''
+ const queryIndex = beforeHash.indexOf('?')
+ if (queryIndex < 0) return url
+
+ const path = beforeHash.slice(0, queryIndex)
+ const params = new URLSearchParams(beforeHash.slice(queryIndex + 1))
+ ;['token', 'api_key', 'apiKey', 'ApiKey'].forEach((key) => params.delete(key))
+ const query = params.toString()
+ return `${path}${query ? `?${query}` : ''}${hash}`
+}
+
// getToken returns the current auth token
export function getToken(): string | null {
return useAuthStore.getState().token
diff --git a/web/src/api/discover.ts b/web/src/api/discover.ts
index c9a51bd..c945bff 100644
--- a/web/src/api/discover.ts
+++ b/web/src/api/discover.ts
@@ -16,6 +16,7 @@ export interface DiscoverItem extends Partial {
year?: number
rating?: number
subscribe_keyword?: string
+ subscribe_aliases?: string[]
}
export interface DiscoverSection {
diff --git a/web/src/api/discover_extra.ts b/web/src/api/discover_extra.ts
deleted file mode 100644
index 476d741..0000000
--- a/web/src/api/discover_extra.ts
+++ /dev/null
@@ -1,17 +0,0 @@
-import { api } from './client'
-import type { DiscoverItem, DiscoverSection } from '../types'
-
-// discoverExtraAPI wraps the Vue-style multi-section feed used by the
-// React DiscoverPage rails. Use the existing /discover/trending and
-// /discover/popular helpers for the simple cases.
-export const discoverExtraAPI = {
- sections: () =>
- api.get<{ sections: DiscoverSection[] }>('/discover/sections').then((r) => r.data.sections),
-
- feed: (sectionKeys: string[]) =>
- api
- .get>('/discover/feed', {
- params: { sections: sectionKeys.join(',') },
- })
- .then((r) => r.data),
-}
diff --git a/web/src/api/library.ts b/web/src/api/library.ts
index 6b2cfef..f359e49 100644
--- a/web/src/api/library.ts
+++ b/web/src/api/library.ts
@@ -1,5 +1,6 @@
import { api, BATCH_REQUEST_TIMEOUT, LONG_REQUEST_TIMEOUT } from './client'
import type { Library, Media, ScanResult } from '../types'
+import type { SeriesCard } from '../utils/groupSeries'
export interface MediaPage {
items: Media[]
@@ -15,6 +16,13 @@ export interface MediaSearchPage {
page_size?: number
}
+export interface SeriesPage {
+ items: SeriesCard[]
+ total: number
+ page: number
+ page_size: number
+}
+
export interface ManualScrapeCandidate {
source: string
media_type?: string
@@ -35,6 +43,15 @@ export interface ManualScrapeCandidate {
nsfw?: boolean
}
+export interface ScrapeOptions {
+ episode_artwork?: boolean
+ episode_images?: boolean
+ refresh_matched?: boolean
+ include_matched?: boolean
+}
+
+export type ManualScrapeApplyOptions = ScrapeOptions
+
export interface MediaMetadataUpdate {
title?: string
original_name?: string
@@ -71,26 +88,51 @@ export const libraryAPI = {
scan: (id: string) =>
api.post(`/libraries/${id}/scan`, null, { timeout: BATCH_REQUEST_TIMEOUT }).then((r) => r.data),
- scrape: (id: string) =>
- api.post(`/libraries/${id}/scrape`, null, { timeout: BATCH_REQUEST_TIMEOUT }).then((r) => r.data),
+ scrape: (id: string, options?: ScrapeOptions) =>
+ api.post(`/libraries/${id}/scrape`, options ?? null, { timeout: BATCH_REQUEST_TIMEOUT }).then((r) => r.data),
- listMedia: (id: string, page = 1, pageSize = 50) =>
+ listMedia: (id: string, page = 1, pageSize = 50, options?: { groupVersions?: boolean }) =>
api
.get(`/libraries/${id}/media`, {
+ params: {
+ page,
+ page_size: pageSize,
+ group_versions: options?.groupVersions === false ? 0 : undefined,
+ },
+ timeout: LONG_REQUEST_TIMEOUT,
+ })
+ .then((r) => r.data),
+
+ listSeries: (id: string, page = 1, pageSize = 500) =>
+ api
+ .get(`/libraries/${id}/series`, {
params: { page, page_size: pageSize },
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
+
+ listSeriesEpisodes: (id: string, key: string) =>
+ api
+ .get<{ items: Media[]; total: number }>(`/libraries/${id}/series/episodes`, {
+ params: { key },
+ timeout: LONG_REQUEST_TIMEOUT,
+ })
+ .then((r) => r.data),
}
export const mediaAPI = {
search: (q: string, limit = 50) =>
api.get('/media', { params: { q, limit } }).then((r) => r.data),
- searchPage: (q: string, page = 1, pageSize = 50) =>
+ searchPage: (q: string, page = 1, pageSize = 50, options?: { groupVersions?: boolean }) =>
api
.get('/media', {
- params: { q, page, page_size: pageSize },
+ params: {
+ q,
+ page,
+ page_size: pageSize,
+ group_versions: options?.groupVersions === false ? 0 : undefined,
+ },
timeout: LONG_REQUEST_TIMEOUT,
})
.then((r) => r.data),
@@ -105,11 +147,25 @@ export const mediaAPI = {
.get<{ items: ManualScrapeCandidate[] }>(`/media/${id}/scrape/search`, { params })
.then((r) => r.data.items),
- applyManualScrape: (id: string, match: ManualScrapeCandidate) =>
- api.post(`/media/${id}/scrape/apply`, match, { timeout: LONG_REQUEST_TIMEOUT }).then((r) => r.data),
-
- applyManualScrapeBatch: (mediaIDs: string[], match: ManualScrapeCandidate) =>
+ applyManualScrape: (id: string, match: ManualScrapeCandidate, options?: ManualScrapeApplyOptions) =>
api
- .post<{ applied: number; errors?: string[] }>('/media/scrape/apply', { media_ids: mediaIDs, match }, { timeout: BATCH_REQUEST_TIMEOUT })
+ .post(
+ `/media/${id}/scrape/apply`,
+ episodeImageOption(options) === undefined ? match : { ...match, episode_images: episodeImageOption(options) },
+ { timeout: LONG_REQUEST_TIMEOUT },
+ )
+ .then((r) => r.data),
+
+ applyManualScrapeBatch: (mediaIDs: string[], match: ManualScrapeCandidate, options?: ManualScrapeApplyOptions) =>
+ api
+ .post<{ applied: number; errors?: string[] }>(
+ '/media/scrape/apply',
+ { media_ids: mediaIDs, match, episode_images: episodeImageOption(options) },
+ { timeout: BATCH_REQUEST_TIMEOUT },
+ )
.then((r) => r.data),
}
+
+function episodeImageOption(options?: ScrapeOptions): boolean | undefined {
+ return options?.episode_images ?? options?.episode_artwork
+}
diff --git a/web/src/api/media_extra.ts b/web/src/api/media_extra.ts
deleted file mode 100644
index fea1026..0000000
--- a/web/src/api/media_extra.ts
+++ /dev/null
@@ -1,19 +0,0 @@
-import { api } from './client'
-import type { Media } from '../types'
-
-// Auxiliary media surfaces used by the home page rails and the admin
-// dashboard "library composition" card.
-export const mediaExtraAPI = {
- recent: (limit = 12) =>
- api.get('/media/recent', { params: { limit } }).then((r) => r.data),
-
- stats: () =>
- api
- .get<{
- by_type: { movies: number; tv: number; anime: number; music: number; unscraped: number }
- total: number
- total_size: number
- total_seconds: number
- }>('/media/stats')
- .then((r) => r.data),
-}
diff --git a/web/src/api/permissions.ts b/web/src/api/permissions.ts
deleted file mode 100644
index e53cbb1..0000000
--- a/web/src/api/permissions.ts
+++ /dev/null
@@ -1,53 +0,0 @@
-import { api } from './client'
-
-export interface UserPermission {
- user_id: string
- can_play_media: boolean
- can_favorite: boolean
- can_view_history: boolean
- can_view_dashboard: boolean
- can_view_discover: boolean
- can_manage_downloads: boolean
- can_manage_subscriptions: boolean
- can_manage_sites: boolean
- can_manage_files: boolean
- can_manage_strm: boolean
- can_cast: boolean
- can_use_ai_assistant: boolean
- can_access_settings: boolean
- updated_at: string
-}
-
-type ApiEnvelope = {
- code?: number
- message?: string
- data?: T
-}
-
-function unwrap(raw: T | ApiEnvelope): T {
- if (raw && typeof raw === 'object' && 'data' in raw && (raw as ApiEnvelope).data !== undefined) {
- return (raw as ApiEnvelope).data as T
- }
- return raw as T
-}
-
-export const permissionsAPI = {
- // Caller's effective permissions; admins always get the all-true set.
- mine: () => api.get>('/auth/permissions').then((r) => unwrap(r.data)),
-
- // Admin endpoints
- get: (userID: string) =>
- api
- .get>(`/admin/users/${userID}/permissions`)
- .then((r) => unwrap(r.data)),
-
- save: (userID: string, p: UserPermission) =>
- api
- .put>(`/admin/users/${userID}/permissions`, p)
- .then((r) => unwrap(r.data)),
-
- reset: (userID: string) =>
- api
- .post>(`/admin/users/${userID}/permissions/reset`)
- .then((r) => unwrap(r.data)),
-}
diff --git a/web/src/api/playback.ts b/web/src/api/playback.ts
index 9665d49..90b364d 100644
--- a/web/src/api/playback.ts
+++ b/web/src/api/playback.ts
@@ -27,6 +27,11 @@ export interface ExternalPlayer {
url: string
}
+function publicOriginHeader() {
+ if (typeof window === 'undefined' || !window.location?.origin) return undefined
+ return { 'X-MediaStation-Public-Origin': window.location.origin }
+}
+
export const playbackAPI = {
recordProgress: (mediaId: string, positionMs: number, durationMs: number) =>
api
@@ -70,9 +75,15 @@ export const playbackAPI = {
externalPlayers: (mediaId: string) =>
api
- .get<{ players: ExternalPlayer[]; url?: string }>(`/playback/${mediaId}/external-players`)
+ .get<{ players: ExternalPlayer[]; url?: string }>(`/playback/${mediaId}/external-players`, {
+ headers: publicOriginHeader(),
+ })
.then((r) => r.data),
externalURL: (mediaId: string) =>
- api.get<{ url: string }>(`/playback/${mediaId}/external-url`).then((r) => r.data),
+ api
+ .get<{ url: string }>(`/playback/${mediaId}/external-url`, {
+ headers: publicOriginHeader(),
+ })
+ .then((r) => r.data),
}
diff --git a/web/src/api/series.ts b/web/src/api/series.ts
deleted file mode 100644
index c421e08..0000000
--- a/web/src/api/series.ts
+++ /dev/null
@@ -1,14 +0,0 @@
-import { api } from './client'
-import type { Media } from '../types'
-
-export interface SeasonGroup {
- season: number
- episodes: Media[]
-}
-
-export const seriesAPI = {
- seasons: (libraryID: string) =>
- api
- .get<{ seasons: SeasonGroup[] }>(`/libraries/${libraryID}/seasons`)
- .then((r) => r.data.seasons),
-}
diff --git a/web/src/api/stats_extra.ts b/web/src/api/stats_extra.ts
deleted file mode 100644
index 7cec793..0000000
--- a/web/src/api/stats_extra.ts
+++ /dev/null
@@ -1,40 +0,0 @@
-import { api } from './client'
-import type { Hardware, Library, Media } from '../types'
-
-// statsExtraAPI exposes the admin dashboard surfaces beyond /stats.
-export const statsExtraAPI = {
- overview: () =>
- api
- .get<{
- libraries: number
- media_count: number
- users_count: number
- total_size: number
- total_seconds: number
- generated_at: string
- }>('/stats/overview')
- .then((r) => r.data),
-
- trend: (days = 14) =>
- api
- .get<{ trend: { day: string; count: number }[]; days: number }>('/stats/trend', {
- params: { days },
- })
- .then((r) => r.data),
-
- topContent: (limit = 10) =>
- api
- .get<{
- items: { media: Media; play_count: number; last_played: string }[]
- }>('/stats/top-content', { params: { limit } })
- .then((r) => r.data),
-
- libraries: () =>
- api
- .get<{
- libraries: { library: Library; item_count: number; total_size: number }[]
- }>('/stats/libraries')
- .then((r) => r.data),
-
- monitor: () => api.get('/stats/monitor').then((r) => r.data),
-}
diff --git a/web/src/api/storage_config.ts b/web/src/api/storage_config.ts
index d46da42..0c0057d 100644
--- a/web/src/api/storage_config.ts
+++ b/web/src/api/storage_config.ts
@@ -1,6 +1,6 @@
import { api, BATCH_REQUEST_TIMEOUT, LONG_REQUEST_TIMEOUT } from './client'
-export type StorageType = 'alist' | 'openlist' | 's3' | 'webdav' | 'cloud115' | 'quark' | 'clouddrive2'
+export type StorageType = 'alist' | 'openlist' | 'webdav' | 'cloud115' | 'clouddrive2'
export interface CloudEntry {
id: string
@@ -141,6 +141,20 @@ export const cloudAPI = {
})
.then((r) => r.data),
+ mkdir: (type: StorageType, dir: string, name: string) =>
+ api
+ .post<{ entry: CloudEntry }>(`/admin/cloud/${type}/mkdir`, { dir, name }, {
+ timeout: LONG_REQUEST_TIMEOUT,
+ })
+ .then((r) => r.data),
+
+ rename: (type: StorageType, ref: string, name: string) =>
+ api
+ .put<{ entry: CloudEntry }>(`/admin/cloud/${type}/rename`, { ref, name }, {
+ timeout: LONG_REQUEST_TIMEOUT,
+ })
+ .then((r) => r.data),
+
import: (type: StorageType, ref: string, name: string, size: number) =>
api
.post(`/admin/cloud/${type}/import`, { ref, name, size })
diff --git a/web/src/api/strm.ts b/web/src/api/strm.ts
index 5a53b17..8b9ed70 100644
--- a/web/src/api/strm.ts
+++ b/web/src/api/strm.ts
@@ -15,6 +15,7 @@ export type GenerateSTRMResult = {
generated: number
updated: number
skipped: number
+ cleaned: number
errors?: string[]
items?: Array<{
media_id: string
diff --git a/web/src/api/subscriptions.ts b/web/src/api/subscriptions.ts
index 105a2a8..b284aa7 100644
--- a/web/src/api/subscriptions.ts
+++ b/web/src/api/subscriptions.ts
@@ -18,6 +18,28 @@ export function buildSiteSearchFeedURL(keyword: string, source?: string, aliases
return `site-search://search?${params.toString()}`
}
+export function buildSubscriptionAliases(item: {
+ title?: string
+ original_name?: string
+ subscribe_keyword?: string
+ subscribe_aliases?: string[]
+ year?: number
+}) {
+ const withYear = (value?: string) => {
+ const title = (value || '').trim()
+ if (!title) return ''
+ return item.year && item.year > 0 ? `${title} ${item.year}` : title
+ }
+ return [
+ ...(item.subscribe_aliases || []),
+ item.title || '',
+ item.original_name || '',
+ withYear(item.title),
+ withYear(item.original_name),
+ item.subscribe_keyword || '',
+ ]
+}
+
export const subscriptionsAPI = {
list: () =>
api.get<{ items: Subscription[] }>('/subscriptions').then((r) => r.data.items),
@@ -38,6 +60,8 @@ export const subscriptionsAPI = {
poster_url?: string
backdrop_url?: string
overview?: string
+ original_name?: string
+ year?: number
resolution?: string
quality?: string
effects?: string
diff --git a/web/src/api/tools.ts b/web/src/api/tools.ts
index 383062c..085f7a4 100644
--- a/web/src/api/tools.ts
+++ b/web/src/api/tools.ts
@@ -1,4 +1,5 @@
import { api } from './client'
+import type { ScrapeOptions } from './library'
// toolsAPI groups admin-only endpoints that don't fit the other domain
// modules: organizing media files into the canonical naming layout, and
@@ -55,6 +56,7 @@ export const toolsAPI = {
organized: number
skipped: number
replaced?: number
+ reclassified?: number
source_path?: string
dest_path?: string
errors?: string[]
@@ -101,15 +103,15 @@ export const toolsAPI = {
// repairAndRescrapeAll 触发「全库修复+重刮」:先从媒体路径中的
// {tmdb-N}/{bangumi-N} 占位符回填缺失/错误的外部 ID,再批量重刮整库。
// 后端异步执行,立即返回;进度通过 WS "scrape" topic 推送。
- repairAndRescrapeAll: () =>
+ repairAndRescrapeAll: (options?: ScrapeOptions) =>
api
- .post<{ status: string }>('/admin/media/repair-rescrape', {})
+ .post<{ status: string }>('/admin/media/repair-rescrape', options ?? {})
.then((r) => r.data),
// repairAndRescrapeLibrary 触发「单库修复+重刮」:只对指定媒体库回填占位符
// 外部 ID 并重刮,不影响其它库。后端异步执行,进度通过 WS "scrape" topic 推送。
- repairAndRescrapeLibrary: (libraryID: string) =>
+ repairAndRescrapeLibrary: (libraryID: string, options?: ScrapeOptions) =>
api
- .post<{ status: string }>(`/admin/libraries/${libraryID}/repair-rescrape`, {})
+ .post<{ status: string }>(`/admin/libraries/${libraryID}/repair-rescrape`, options ?? {})
.then((r) => r.data),
}
diff --git a/web/src/appRoutes.tsx b/web/src/appRoutes.tsx
new file mode 100644
index 0000000..5ce167e
--- /dev/null
+++ b/web/src/appRoutes.tsx
@@ -0,0 +1,106 @@
+/* eslint-disable react-refresh/only-export-components */
+import { lazy, type ReactElement } from 'react'
+import { Navigate } from 'react-router-dom'
+
+const HomePage = lazy(() => import('./pages/HomePage').then((m) => ({ default: m.HomePage })))
+const LibraryPage = lazy(() => import('./pages/LibraryPage').then((m) => ({ default: m.LibraryPage })))
+const LibrariesPage = lazy(() => import('./pages/LibrariesPage').then((m) => ({ default: m.LibrariesPage })))
+const SearchPage = lazy(() => import('./pages/SearchPage').then((m) => ({ default: m.SearchPage })))
+const FavouritesPage = lazy(() => import('./pages/FavouritesPage').then((m) => ({ default: m.FavouritesPage })))
+const PlaylistsPage = lazy(() => import('./pages/PlaylistsPage').then((m) => ({ default: m.PlaylistsPage })))
+const PlaylistDetailPage = lazy(() =>
+ import('./pages/PlaylistDetailPage').then((m) => ({ default: m.PlaylistDetailPage })),
+)
+const MediaDetailPage = lazy(() => import('./pages/MediaDetailPage').then((m) => ({ default: m.MediaDetailPage })))
+const PlayerPage = lazy(() => import('./pages/PlayerPage').then((m) => ({ default: m.PlayerPage })))
+const AdminPage = lazy(() => import('./pages/AdminPage').then((m) => ({ default: m.AdminPage })))
+const DownloadsPage = lazy(() => import('./pages/DownloadsPage').then((m) => ({ default: m.DownloadsPage })))
+const SubscriptionsPage = lazy(() =>
+ import('./pages/SubscriptionsPage').then((m) => ({ default: m.SubscriptionsPage })),
+)
+const ProfilePage = lazy(() => import('./pages/ProfilePage').then((m) => ({ default: m.ProfilePage })))
+const StatsPage = lazy(() => import('./pages/StatsPage').then((m) => ({ default: m.StatsPage })))
+const DiscoverPage = lazy(() => import('./pages/DiscoverPage').then((m) => ({ default: m.DiscoverPage })))
+const TasksPage = lazy(() => import('./pages/TasksPage').then((m) => ({ default: m.TasksPage })))
+const RecycleBinPage = lazy(() => import('./pages/RecycleBinPage').then((m) => ({ default: m.RecycleBinPage })))
+const DlnaPage = lazy(() => import('./pages/DlnaPage').then((m) => ({ default: m.DlnaPage })))
+const FileManagerPage = lazy(() =>
+ import('./pages/FileManagerPage').then((m) => ({ default: m.FileManagerPage })),
+)
+const StoragePage = lazy(() => import('./pages/StoragePage').then((m) => ({ default: m.StoragePage })))
+const DuplicatesPage = lazy(() => import('./pages/DuplicatesPage').then((m) => ({ default: m.DuplicatesPage })))
+const SchedulerPage = lazy(() => import('./pages/SchedulerPage').then((m) => ({ default: m.SchedulerPage })))
+const WatchHistoryPage = lazy(() =>
+ import('./pages/WatchHistoryPage').then((m) => ({ default: m.WatchHistoryPage })),
+)
+const PosterWallPage = lazy(() => import('./pages/PosterWallPage').then((m) => ({ default: m.PosterWallPage })))
+const SitesPage = lazy(() => import('./pages/SitesPage').then((m) => ({ default: m.SitesPage })))
+const SiteSearchPage = lazy(() => import('./pages/SiteSearchPage').then((m) => ({ default: m.SiteSearchPage })))
+const AIAssistantPage = lazy(() =>
+ import('./pages/AIAssistantPage').then((m) => ({ default: m.AIAssistantPage })),
+)
+const StrmPage = lazy(() => import('./pages/StrmPage').then((m) => ({ default: m.StrmPage })))
+const ProfileManagementPage = lazy(() =>
+ import('./pages/ProfileManagementPage').then((m) => ({ default: m.ProfileManagementPage })),
+)
+const NotifyChannelsPage = lazy(() =>
+ import('./pages/NotifyChannelsPage').then((m) => ({ default: m.NotifyChannelsPage })),
+)
+const SettingsPage = lazy(() => import('./pages/SettingsPage').then((m) => ({ default: m.SettingsPage })))
+const AssistantChatPage = lazy(() =>
+ import('./pages/AssistantChatPage').then((m) => ({ default: m.AssistantChatPage })),
+)
+const DownloadClientsPage = lazy(() =>
+ import('./pages/DownloadClientsPage').then((m) => ({ default: m.DownloadClientsPage })),
+)
+const StorageConfigPage = lazy(() =>
+ import('./pages/StorageConfigPage').then((m) => ({ default: m.StorageConfigPage })),
+)
+const LicensePage = lazy(() => import('./pages/LicensePage').then((m) => ({ default: m.LicensePage })))
+
+export type AppRoute = {
+ path?: string
+ index?: boolean
+ element: ReactElement
+ adminOnly?: boolean
+}
+
+export const appRoutes: AppRoute[] = [
+ { index: true, element: },
+ { path: 'libraries', element: },
+ { path: 'library/:id', element: },
+ { path: 'discover', element: },
+ { path: 'search', element: },
+ { path: 'favourites', element: },
+ { path: 'playlists', element: },
+ { path: 'playlist/:id', element: },
+ { path: 'media/:id', element: },
+ { path: 'play/:id', element: },
+ { path: 'downloads', element: },
+ { path: 'subscriptions', element: },
+ { path: 'profile', element: },
+ { path: 'dlna', element: },
+ { path: 'history', element: },
+ { path: 'poster-wall', element: },
+ { path: 'site-search', element: },
+ { path: 'ai', element: },
+ { path: 'play-profiles', element: },
+ { path: 'api-configs', element: },
+ { path: 'tools', element: },
+ { path: 'sites', element: , adminOnly: true },
+ { path: 'files', element: , adminOnly: true },
+ { path: 'storage', element: , adminOnly: true },
+ { path: 'duplicates', element: , adminOnly: true },
+ { path: 'scheduler', element: , adminOnly: true },
+ { path: 'tasks', element: , adminOnly: true },
+ { path: 'recycle', element: , adminOnly: true },
+ { path: 'strm', element: , adminOnly: true },
+ { path: 'notify-channels', element: , adminOnly: true },
+ { path: 'settings', element: , adminOnly: true },
+ { path: 'assistant', element: , adminOnly: true },
+ { path: 'download-clients', element: , adminOnly: true },
+ { path: 'license', element: , adminOnly: true },
+ { path: 'storage-config', element: , adminOnly: true },
+ { path: 'stats', element: , adminOnly: true },
+ { path: 'admin', element: , adminOnly: true },
+]
diff --git a/web/src/components/APIConfigsPanel.tsx b/web/src/components/APIConfigsPanel.tsx
index d7f4b34..ca45b4b 100644
--- a/web/src/components/APIConfigsPanel.tsx
+++ b/web/src/components/APIConfigsPanel.tsx
@@ -3,7 +3,7 @@ import toast from 'react-hot-toast'
import { Eye, KeyRound, Pencil, Save, Trash2, X } from 'lucide-react'
import { apiConfigsAPI, type APIConfig } from '../api/api_configs'
-import { confirmAction } from './ConfirmDialog'
+import { confirmAction } from './confirmAction'
// Compact inline-editable provider table for use inside AdminPage's "外部API" tab.
export function APIConfigsPanel() {
diff --git a/web/src/components/ConfirmDialog.tsx b/web/src/components/ConfirmDialog.tsx
index 3e18a5f..07c554a 100644
--- a/web/src/components/ConfirmDialog.tsx
+++ b/web/src/components/ConfirmDialog.tsx
@@ -1,7 +1,6 @@
-import { createRoot } from 'react-dom/client'
import { AlertTriangle } from 'lucide-react'
-type ConfirmOptions = {
+export type ConfirmOptions = {
title?: string
message: string
confirmText?: string
@@ -9,21 +8,7 @@ type ConfirmOptions = {
danger?: boolean
}
-export function confirmAction(options: ConfirmOptions): Promise {
- return new Promise((resolve) => {
- const host = document.createElement('div')
- document.body.appendChild(host)
- const root = createRoot(host)
- const close = (value: boolean) => {
- root.unmount()
- host.remove()
- resolve(value)
- }
- root.render()
- })
-}
-
-function ConfirmDialog({
+export function ConfirmDialog({
options,
onClose,
}: {
diff --git a/web/src/components/EpisodeArtworkToggle.tsx b/web/src/components/EpisodeArtworkToggle.tsx
new file mode 100644
index 0000000..b2be863
--- /dev/null
+++ b/web/src/components/EpisodeArtworkToggle.tsx
@@ -0,0 +1,57 @@
+import { ImageOff, ImagePlus } from 'lucide-react'
+
+type EpisodeArtworkToggleProps = {
+ checked: boolean
+ onChange: (checked: boolean) => void
+ title?: string
+ className?: string
+}
+
+export function EpisodeArtworkToggle({
+ checked,
+ onChange,
+ title = '关闭后仍会写入每集文字元数据,只跳过每集图片',
+ className = '',
+}: EpisodeArtworkToggleProps) {
+ const Icon = checked ? ImagePlus : ImageOff
+ const controlStyle = checked
+ ? { backgroundColor: 'var(--app-brand-soft)', borderColor: 'var(--app-brand-border)', color: 'var(--app-brand-text)' }
+ : undefined
+ const emphasisStyle = checked
+ ? { backgroundColor: 'var(--app-brand-emphasis)', color: 'var(--app-brand-text)' }
+ : undefined
+
+ return (
+
+ )
+}
diff --git a/web/src/components/GlobalEvents.tsx b/web/src/components/GlobalEvents.tsx
index 36a7a7b..074a34f 100644
--- a/web/src/components/GlobalEvents.tsx
+++ b/web/src/components/GlobalEvents.tsx
@@ -37,7 +37,18 @@ export function GlobalEvents() {
}
}
if (topic === 'scrape' && p.finished) {
- toast.success(`刮削完成:成功匹配 ${p.matched ?? 0} 项`)
+ const processed = Number(p.processed ?? 0)
+ const matched = Number(p.matched ?? 0)
+ const failed = Number(p.failed ?? 0)
+ if (failed > 0) {
+ toast.error(`刮削完成但有错误:处理 ${processed} 项 · 成功匹配 ${matched} 项 · 失败 ${failed} 项`)
+ return
+ }
+ if (processed > 0) {
+ toast.success(`刮削完成:处理 ${processed} 项 · 成功匹配 ${matched} 项`)
+ } else {
+ toast.success(`刮削完成:没有待刮削媒体 · 成功匹配 ${matched} 项`)
+ }
}
if (topic === 'subscription') {
const queued = (p.queued as number | undefined) ?? 0
diff --git a/web/src/components/Layout.tsx b/web/src/components/Layout.tsx
index b00925f..3ff9555 100644
--- a/web/src/components/Layout.tsx
+++ b/web/src/components/Layout.tsx
@@ -1,343 +1,75 @@
-import { useEffect, useMemo, useRef, useState } from 'react'
-import { Link, NavLink, Outlet, useLocation, useNavigate } from 'react-router-dom'
+import { Link, Outlet, useLocation, useNavigate } from 'react-router-dom'
import { AnimatePresence, motion } from 'framer-motion'
-import toast from 'react-hot-toast'
-import {
- Activity, Clock, CloudDownload, Compass,
- Cast, Globe, HardDrive, Heart, Home, Image, KeySquare,
- ListMusic, LogOut, MessageSquareText, Rss, Search, Trash2,
- Settings, Sliders, Sparkles, UserCog,
- Library as LibraryIcon, User as UserIcon, ChevronDown, Menu, X
-} from 'lucide-react'
+import { Menu, MessageSquareText, Search, Sparkles } from 'lucide-react'
import clsx from 'clsx'
import { AppFooter } from './AppFooter'
import { useAuthStore } from '../stores/auth'
-import { usePermissionStore } from '../stores/permissions'
import { usePlayProfileStore } from '../stores/playProfile'
-import { imageURL } from '../api/client'
-import { mediaAPI } from '../api/library'
-import { playProfilesAPI } from '../api/play_profiles'
-import { requestPIN } from './PinDialog'
-import type { Media, PlayProfile } from '../types'
-import { groupSeries, seriesCardLink } from '../utils/groupSeries'
+import { LayoutSearchBox } from './LayoutSearchBox'
+import { LayoutSidebarContent } from './LayoutSidebarContent'
+import { LayoutThemeToggle } from './LayoutThemeToggle'
+import { LayoutUserMenu } from './LayoutUserMenu'
+import { useLayoutSearch } from './useLayoutSearch'
+import { useLayoutPermissions } from './useLayoutPermissions'
+import { useLayoutProfiles } from './useLayoutProfiles'
+import { useLayoutSidebar } from './useLayoutSidebar'
+import { useThemeMode } from './useThemeMode'
export function Layout() {
const navigate = useNavigate()
const location = useLocation()
const user = useAuthStore((s) => s.user)
const logout = useAuthStore((s) => s.logout)
- const permissions = usePermissionStore((s) => s.permissions)
- const isSuper = usePermissionStore((s) => s.isSuper)
- const isPermissionLoading = usePermissionStore((s) => s.isLoading)
- const fetchPermissions = usePermissionStore((s) => s.fetchPermissions)
const activeProfileId = usePlayProfileStore((s) => s.activeProfileId)
const setActiveProfile = usePlayProfileStore((s) => s.setActiveProfile)
- const [isSidebarOpen, setIsSidebarOpen] = useState(true)
- const [isMobileDrawerOpen, setIsMobileDrawerOpen] = useState(false)
- const [isProfileOpen, setIsProfileOpen] = useState(false)
- const [openGroups, setOpenGroups] = useState>({ media: true })
- const [profiles, setProfiles] = useState([])
- const [searchFocused, setSearchFocused] = useState(false)
- const [searchQuery, setSearchQuery] = useState('')
- const [searchItems, setSearchItems] = useState([])
- const [searchLoading, setSearchLoading] = useState(false)
- const [searchTotal, setSearchTotal] = useState(0)
- const [searchError, setSearchError] = useState('')
- const searchSeq = useRef(0)
- const searchCards = useMemo(() => groupSeries(searchItems).slice(0, 8), [searchItems])
+ const theme = useThemeMode()
+ const search = useLayoutSearch({
+ pathname: location.pathname,
+ locationSearch: location.search,
+ navigate,
+ })
+ const { can, isAdmin } = useLayoutPermissions(user)
+ const {
+ isMobileDrawerOpen,
+ isRouteIn,
+ isSidebarOpen,
+ openGroups,
+ setIsMobileDrawerOpen,
+ setIsSidebarOpen,
+ toggleGroup,
+ } = useLayoutSidebar(location.pathname)
+ const {
+ activeProfile,
+ isProfileOpen,
+ profiles,
+ setIsProfileOpen,
+ switchProfile,
+ useDefaultProfile,
+ } = useLayoutProfiles({ activeProfileId, setActiveProfile, user })
- // Auto-collapse sidebar on smaller tablet screens, and auto-hide drawer on path change
- useEffect(() => {
- const handleResize = () => {
- if (window.innerWidth < 1024) {
- setIsSidebarOpen(false)
- } else {
- setIsSidebarOpen(true)
- }
- }
- handleResize()
- window.addEventListener('resize', handleResize)
- return () => window.removeEventListener('resize', handleResize)
- }, [])
-
- useEffect(() => {
- setIsMobileDrawerOpen(false)
- }, [location.pathname])
-
- useEffect(() => {
- if (user && !isPermissionLoading && Object.keys(permissions ?? {}).length === 0) {
- fetchPermissions().catch(() => undefined)
- }
- }, [fetchPermissions, isPermissionLoading, permissions, user])
-
- useEffect(() => {
- if (location.pathname === '/search') {
- const query = new URLSearchParams(location.search).get('q') ?? ''
- setSearchQuery(query)
- }
- }, [location.pathname, location.search])
-
- useEffect(() => {
- const query = searchQuery.trim()
- const seq = ++searchSeq.current
- if (!searchFocused || !query) {
- setSearchItems([])
- setSearchTotal(0)
- setSearchError('')
- setSearchLoading(false)
- return
- }
-
- setSearchLoading(true)
- setSearchError('')
- const timer = window.setTimeout(() => {
- mediaAPI
- .search(query, 24)
- .then((data) => {
- if (seq !== searchSeq.current) return
- setSearchItems(data.items ?? [])
- setSearchTotal(data.total ?? (data.items ?? []).length)
- })
- .catch(() => {
- if (seq !== searchSeq.current) return
- setSearchItems([])
- setSearchTotal(0)
- setSearchError('搜索失败,请稍后再试')
- })
- .finally(() => {
- if (seq === searchSeq.current) setSearchLoading(false)
- })
- }, 220)
-
- return () => window.clearTimeout(timer)
- }, [searchFocused, searchQuery])
-
- useEffect(() => {
- if (!user) {
- setProfiles([])
- setActiveProfile(null)
- return
- }
- playProfilesAPI
- .list()
- .then((rows) => {
- setProfiles(rows)
- const active = rows.find((p) => p.id === activeProfileId)
- if (!active) {
- const defaultProfile = rows.find((p) => p.is_default && !p.require_pin)
- setActiveProfile(defaultProfile?.id ?? null)
- }
- })
- .catch(() => undefined)
- }, [activeProfileId, setActiveProfile, user])
-
- const isAdmin = user?.role === 'admin'
- const can = (key: string) => isAdmin || isSuper || (permissions ?? {})[key] === true
- const activeProfile = profiles.find((p) => p.id === activeProfileId) ?? null
- const sidebarExpanded = isSidebarOpen || isMobileDrawerOpen
- const isRouteIn = (paths: string[]) =>
- paths.some((path) => (path === '/' ? location.pathname === '/' : location.pathname.startsWith(path)))
- const toggleGroup = (key: string) =>
- setOpenGroups((current) => ({ ...current, [key]: !current[key] }))
-
- useEffect(() => {
- const groupPaths: Record = {
- media: ['/', '/libraries', '/library', '/poster-wall', '/discover', '/search', '/dlna', '/ai'],
- personal: ['/favourites', '/playlists', '/playlist', '/history', '/profile', '/play-profiles'],
- downloads: ['/downloads', '/download-clients', '/subscriptions', '/site-search'],
- tools: ['/storage', '/storage-config', '/files', '/strm', '/duplicates', '/tasks', '/scheduler', '/recycle', '/stats'],
- system: ['/admin', '/sites', '/notify-channels', '/license', '/settings', '/assistant'],
- }
- const active = Object.entries(groupPaths).find(([, paths]) => isRouteIn(paths))?.[0]
- if (active) {
- setOpenGroups((current) => (current[active] ? current : { ...current, [active]: true }))
- }
- // eslint-disable-next-line react-hooks/exhaustive-deps
- }, [location.pathname])
-
- const handleSearchSubmit = (e: React.FormEvent) => {
- e.preventDefault()
- if (searchQuery.trim()) {
- navigate(`/search?q=${encodeURIComponent(searchQuery.trim())}`)
- setSearchFocused(false)
- }
- }
-
- const handleProfileSwitch = async (profile: PlayProfile) => {
- if (activeProfileId === profile.id) {
- setIsProfileOpen(false)
- return
- }
- try {
- let pinToken: string | null = null
- if (profile.require_pin) {
- const pin = await requestPIN({ profileName: profile.name })
- if (!pin) return
- const verified = await playProfilesAPI.verifyPin(profile.id, pin)
- pinToken = verified.token
- }
- setActiveProfile(profile.id, pinToken)
- setIsProfileOpen(false)
- toast.success(`已切换到「${profile.name}」`)
- } catch (err: unknown) {
- const msg =
- (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? 'PIN 验证失败'
- toast.error(msg)
- }
+ const handleLogout = () => {
+ logout()
+ navigate('/login')
}
const sidebarContent = (
-
- {/* Brand Logo & Brand Title */}
-
-
-

- {(isSidebarOpen || isMobileDrawerOpen) && (
-
- MediaStationGo
-
- )}
-
-
- {/* Toggle Collapse Button for Large Screen */}
-
-
- {/* Mobile Drawer Close Button */}
-
-
-
- {/* Navigation List */}
-
-
- {/* Sidebar Logout Action */}
-
-
-
-
+ setIsSidebarOpen((current) => !current)}
+ onCloseMobileDrawer={() => setIsMobileDrawerOpen(false)}
+ onLogout={handleLogout}
+ />
)
return (
-
+
{/* 1. Desktop Persistent Sidebar */}