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(`<?xml version="1.0" encoding="utf-8"?> +<d:multistatus xmlns:d="DAV:"> + <d:response><d:href>/dav/Anime/JianLai/</d:href><d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat></d:response> + <d:response><d:href>/dav/Anime/JianLai/tvshow.nfo</d:href><d:propstat><d:prop><d:displayname>tvshow.nfo</d:displayname><d:getcontentlength>64</d:getcontentlength><d:resourcetype/></d:prop></d:propstat></d:response> + <d:response><d:href>/dav/Anime/JianLai/poster.jpg</d:href><d:propstat><d:prop><d:displayname>poster.jpg</d:displayname><d:getcontentlength>1024</d:getcontentlength><d:resourcetype/></d:prop></d:propstat></d:response> + <d:response><d:href>/dav/Anime/JianLai/Season1/</d:href><d:propstat><d:prop><d:displayname>Season1</d:displayname><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat></d:response> +</d:multistatus>`)) + case "/dav/Anime/JianLai/Season1": + _, _ = w.Write([]byte(`<?xml version="1.0" encoding="utf-8"?> +<d:multistatus xmlns:d="DAV:"> + <d:response><d:href>/dav/Anime/JianLai/Season1/</d:href><d:propstat><d:prop><d:resourcetype><d:collection/></d:resourcetype></d:prop></d:propstat></d:response> + <d:response><d:href>/dav/Anime/JianLai/Season1/JianLai.S01E01.mkv</d:href><d:propstat><d:prop><d:displayname>JianLai.S01E01.mkv</d:displayname><d:getcontentlength>2048</d:getcontentlength><d:resourcetype/></d:prop></d:propstat></d:response> + <d:response><d:href>/dav/Anime/JianLai/Season1/JianLai.S01E01.nfo</d:href><d:propstat><d:prop><d:displayname>JianLai.S01E01.nfo</d:displayname><d:getcontentlength>128</d:getcontentlength><d:resourcetype/></d:prop></d:propstat></d:response> +</d:multistatus>`)) + 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(`<tvshow><title>剑来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 := `<movie><title>` + 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 */} -
- - MediaStationGo - {(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 */}
) } - -interface SidebarGroupProps { - id: string; - icon: React.ReactNode; - label: string; - children: React.ReactNode; - collapsed?: boolean; - open?: boolean; - active?: boolean; - onToggle: (id: string) => void; -} - -function SidebarGroup({ id, icon, label, children, collapsed, open, active, onToggle }: SidebarGroupProps) { - return ( -
- - - {!collapsed && open && ( - -
- {children} -
-
- )} -
-
- ) -} - -interface SidebarLinkProps { - to: string; - icon: React.ReactNode; - label: string; - end?: boolean; - collapsed?: boolean; - child?: boolean; -} - -function SidebarLink({ to, icon, label, end, collapsed, child }: SidebarLinkProps) { - return ( - - clsx( - "flex items-center gap-3.5 rounded-xl px-4 py-3 text-sm font-semibold transition-all duration-300 relative group", - child && "py-2.5 text-[13px]", - isActive - ? "bg-[#111827] text-white shadow-sm" - : "text-gray-500 hover:bg-gray-50 hover:text-gray-900" - ) - } - > - {({ isActive }) => ( - <> - - {icon} - - {!collapsed && ( - - {label} - - )} - {collapsed && ( -
- {label} -
- )} - - )} -
- ) -} diff --git a/web/src/components/LayoutSearchBox.tsx b/web/src/components/LayoutSearchBox.tsx new file mode 100644 index 0000000..703029d --- /dev/null +++ b/web/src/components/LayoutSearchBox.tsx @@ -0,0 +1,143 @@ +import { FormEvent } from 'react' +import { Link } from 'react-router-dom' +import { AnimatePresence, motion } from 'framer-motion' +import { Library as LibraryIcon, Search } from 'lucide-react' +import clsx from 'clsx' + +import { imageURL } from '../api/client' +import { seriesCardLink, type SeriesCard } from '../utils/groupSeries' + +type LayoutSearchBoxProps = { + query: string + focused: boolean + loading: boolean + error: string + cards: SeriesCard[] + total: number + onQueryChange: (value: string) => void + onFocusedChange: (focused: boolean) => void + onSubmit: (event: FormEvent) => void +} + +export function LayoutSearchBox({ + query, + focused, + loading, + error, + cards, + total, + onQueryChange, + onFocusedChange, + onSubmit, +}: LayoutSearchBoxProps) { + const trimmedQuery = query.trim() + + return ( +
+ + + + onQueryChange(event.target.value)} + onMouseDown={() => onFocusedChange(true)} + onClick={() => onFocusedChange(true)} + onFocus={() => onFocusedChange(true)} + onBlur={() => window.setTimeout(() => onFocusedChange(false), 120)} + placeholder="搜索电影、电视剧、演员、种子站点..." + className="w-full rounded-full border border-[var(--app-border)] bg-[var(--app-control-bg)] py-2.5 pl-11 pr-12 text-sm text-[var(--app-text)] placeholder:text-[var(--app-muted)] outline-none transition-all duration-300 focus:border-brand-500 focus:bg-[var(--app-panel)] focus:ring-4 focus:ring-brand-100/40" + /> +
+ + Enter + +
+ + {focused && trimmedQuery && ( + event.preventDefault()} + className="absolute left-0 right-0 top-full z-50 mt-3 overflow-hidden rounded-2xl border border-[var(--app-border)] bg-[var(--app-panel)] shadow-2xl" + > +
+ {loading && ( +
+ + 搜索中... +
+ )} + {!loading && error && ( +
{error}
+ )} + {!loading && !error && cards.length === 0 && ( +
没有找到匹配的本地媒体
+ )} + {!loading && !error && cards.length > 0 && ( +
+ {cards.map((card) => ( + onFocusedChange(false)} + /> + ))} +
+ )} +
+ onFocusedChange(false)} + className="flex items-center justify-between border-t border-[var(--app-border)] px-4 py-3 text-sm font-semibold text-brand-500 hover:bg-[var(--app-hover)]" + > + 查看全部搜索结果 + + {total > 0 ? `${total} 个条目` : 'Enter'} + + +
+ )} +
+
+ ) +} + +function SearchResultItem({ card, onClick }: { card: SeriesCard; onClick: () => void }) { + return ( + +
+ {card.rep.poster_url ? ( + {card.rep.title} + ) : ( +
+ +
+ )} +
+
+
+ {card.rep.title || card.rep.original_name || '未命名媒体'} +
+
+ {card.rep.year ? {card.rep.year} : null} + {card.count > 1 ? `${card.count} 集/条目` : '单条媒体'} + {card.rep.width ? {card.rep.width}x{card.rep.height} : null} +
+
+ + ) +} diff --git a/web/src/components/LayoutSidebarContent.tsx b/web/src/components/LayoutSidebarContent.tsx new file mode 100644 index 0000000..33df06a --- /dev/null +++ b/web/src/components/LayoutSidebarContent.tsx @@ -0,0 +1,130 @@ +import { Link } from 'react-router-dom' +import { motion } from 'framer-motion' +import { LogOut, Menu, X } from 'lucide-react' +import clsx from 'clsx' +import { LAYOUT_NAV_GROUPS, NAV_GROUP_PATHS, type LayoutNavItem } from './layoutNavigation' +import { SidebarGroup, SidebarLink } from './LayoutSidebarNav' + +type LayoutSidebarContentProps = { + isSidebarOpen: boolean + isMobileDrawerOpen: boolean + openGroups: Record + isAdmin: boolean + username?: string + can: (key: string) => boolean + isRouteIn: (paths: string[]) => boolean + onToggleGroup: (id: string) => void + onToggleSidebar: () => void + onCloseMobileDrawer: () => void + onLogout: () => void +} + +export function LayoutSidebarContent({ + isSidebarOpen, + isMobileDrawerOpen, + openGroups, + isAdmin, + username, + can, + isRouteIn, + onToggleGroup, + onToggleSidebar, + onCloseMobileDrawer, + onLogout, +}: LayoutSidebarContentProps) { + const sidebarExpanded = isSidebarOpen || isMobileDrawerOpen + const isItemVisible = (item: LayoutNavItem) => + (!item.adminOnly || isAdmin) && (!item.permission || can(item.permission)) + const visibleGroups = LAYOUT_NAV_GROUPS + .filter((group) => !group.adminOnly || isAdmin) + .map((group) => ({ + group, + items: group.items.filter(isItemVisible), + })) + .filter(({ items }) => items.length > 0) + + return ( +
+
+ + MediaStationGo + {sidebarExpanded && ( + + MediaStationGo + + )} + + + + + +
+ + + +
+ +
+
+ ) +} diff --git a/web/src/components/LayoutSidebarNav.tsx b/web/src/components/LayoutSidebarNav.tsx new file mode 100644 index 0000000..765b39b --- /dev/null +++ b/web/src/components/LayoutSidebarNav.tsx @@ -0,0 +1,122 @@ +import type { ReactNode } from 'react' +import { NavLink } from 'react-router-dom' +import { AnimatePresence, motion } from 'framer-motion' +import { ChevronDown } from 'lucide-react' +import clsx from 'clsx' + +type SidebarGroupProps = { + id: string + icon: ReactNode + label: string + children: ReactNode + collapsed?: boolean + open?: boolean + active?: boolean + onToggle: (id: string) => void +} + +export function SidebarGroup({ id, icon, label, children, collapsed, open, active, onToggle }: SidebarGroupProps) { + return ( +
+ + + {!collapsed && open && ( + +
+ {children} +
+
+ )} +
+
+ ) +} + +type SidebarLinkProps = { + to: string + icon: ReactNode + label: string + end?: boolean + collapsed?: boolean + child?: boolean +} + +export function SidebarLink({ to, icon, label, end, collapsed, child }: SidebarLinkProps) { + return ( + + clsx( + 'relative flex items-center gap-3.5 rounded-xl px-4 py-3 text-sm font-semibold transition-all duration-300 group', + child && 'py-2.5 text-[13px]', + isActive + ? 'bg-[var(--app-active-bg)] text-[var(--app-active-text)] shadow-sm' + : 'text-[var(--app-muted)] hover:bg-[var(--app-hover)] hover:text-[var(--app-text)]', + ) + } + > + {({ isActive }) => ( + <> + + {icon} + + {!collapsed && ( + + {label} + + )} + {collapsed && ( +
+ {label} +
+ )} + + )} +
+ ) +} diff --git a/web/src/components/LayoutThemeToggle.tsx b/web/src/components/LayoutThemeToggle.tsx new file mode 100644 index 0000000..16d27ca --- /dev/null +++ b/web/src/components/LayoutThemeToggle.tsx @@ -0,0 +1,48 @@ +import { Monitor, Moon, Sun } from 'lucide-react' +import clsx from 'clsx' + +import type { ThemeMode } from './useThemeMode' + +type LayoutThemeToggleProps = { + mode: ThemeMode + onChange: (mode: ThemeMode) => void +} + +const options: Array<{ + mode: ThemeMode + label: string + icon: typeof Sun +}> = [ + { mode: 'light', label: '白天模式', icon: Sun }, + { mode: 'dark', label: '夜晚模式', icon: Moon }, + { mode: 'system', label: '跟随系统', icon: Monitor }, +] + +export function LayoutThemeToggle({ mode, onChange }: LayoutThemeToggleProps) { + return ( +
+ {options.map((option) => { + const Icon = option.icon + const active = mode === option.mode + return ( + + ) + })} +
+ ) +} diff --git a/web/src/components/LayoutUserMenu.tsx b/web/src/components/LayoutUserMenu.tsx new file mode 100644 index 0000000..d6f0a24 --- /dev/null +++ b/web/src/components/LayoutUserMenu.tsx @@ -0,0 +1,143 @@ +import { Link } from 'react-router-dom' +import { AnimatePresence, motion } from 'framer-motion' +import { ChevronDown, LogOut, Settings, User as UserIcon, UserCog } from 'lucide-react' +import clsx from 'clsx' + +import type { PlayProfile } from '../types' + +type LayoutUser = { + username?: string + role?: string +} + +type LayoutUserMenuProps = { + user: LayoutUser | null | undefined + isOpen: boolean + profiles: PlayProfile[] + activeProfileId: string | null + activeProfile: PlayProfile | null + onToggle: () => void + onClose: () => void + onUseDefaultProfile: () => void + onSwitchProfile: (profile: PlayProfile) => void + onLogout: () => void +} + +export function LayoutUserMenu({ + user, + isOpen, + profiles, + activeProfileId, + activeProfile, + onToggle, + onClose, + onUseDefaultProfile, + onSwitchProfile, + onLogout, +}: LayoutUserMenuProps) { + return ( +
+ + + + {isOpen && ( + <> +
+ + } label="个人基本信息" onClick={onClose} /> + {user?.role === 'admin' && ( + } label="管理主控制台" onClick={onClose} /> + )} +
+
+

+ 当前观影 Profile +

+
+ + {profiles.map((profile) => ( + + ))} +
+
+ } label="管理观影 Profile" onClick={onClose} /> +
+ + + + )} + +
+ ) +} + +function UserMenuLink({ + to, + icon, + label, + onClick, +}: { + to: string + icon: React.ReactNode + label: string + onClick: () => void +}) { + return ( + + {icon} + {label} + + ) +} + +function profileButtonClass(active: boolean): string { + return clsx( + 'flex w-full items-center justify-between rounded-xl px-2.5 py-2 text-left text-xs transition-colors', + active + ? 'bg-[var(--app-active-bg)] text-[var(--app-active-text)]' + : 'text-[var(--app-subtle)] hover:bg-[var(--app-hover)]', + ) +} diff --git a/web/src/components/ManagementShortcuts.tsx b/web/src/components/ManagementShortcuts.tsx index 7fc8e38..701dbd1 100644 --- a/web/src/components/ManagementShortcuts.tsx +++ b/web/src/components/ManagementShortcuts.tsx @@ -6,6 +6,7 @@ type ShortcutItem = { title: string description: string badge?: string + group?: string } type ManagementShortcutsProps = { @@ -15,43 +16,80 @@ type ManagementShortcutsProps = { } export function ManagementShortcuts({ title, description, items }: ManagementShortcutsProps) { + const groups = groupShortcutItems(items) + const hasNamedGroups = groups.some((group) => group.label) + return ( -
-
+
+
-

{title}

- {description &&

{description}

} +

{title}

+ {description &&

{description}

}
-
- {items.map((item) => ( - -
-
-
-

{item.title}

- {item.badge && ( - - {item.badge} - - )} -
-

- {item.description} -

-
- + {hasNamedGroups ? ( +
+ {groups.map((group) => ( +
+ {group.label &&

{group.label}

} +
- - ))} -
+ ))} +
+ ) : ( + + )}
) } + +function groupShortcutItems(items: ShortcutItem[]) { + const groups: { label: string; items: ShortcutItem[] }[] = [] + items.forEach((item) => { + const label = item.group ?? '' + const group = groups.find((candidate) => candidate.label === label) + if (group) { + group.items.push(item) + return + } + groups.push({ label, items: [item] }) + }) + return groups +} + +function ShortcutGrid({ items }: { items: ShortcutItem[] }) { + return ( +
+ {items.map((item) => ( + + ))} +
+ ) +} + +function ShortcutCard({ item }: { item: ShortcutItem }) { + return ( + +
+
+
+

{item.title}

+ {item.badge && ( + + {item.badge} + + )} +
+

{item.description}

+
+ +
+ + ) +} diff --git a/web/src/components/ManualScrapeDialog.tsx b/web/src/components/ManualScrapeDialog.tsx index 7076c02..677e35c 100644 --- a/web/src/components/ManualScrapeDialog.tsx +++ b/web/src/components/ManualScrapeDialog.tsx @@ -1,10 +1,11 @@ -import { useEffect, useMemo, useState } from 'react' +import { useEffect, useMemo, useState, type Dispatch, type SetStateAction } from 'react' import { Check, LoaderCircle, Search, Sparkles, X } from 'lucide-react' import toast from 'react-hot-toast' import { imageURL } from '../api/client' import { mediaAPI, type ManualScrapeCandidate } from '../api/library' import type { Media } from '../types' +import { EpisodeArtworkToggle } from './EpisodeArtworkToggle' interface ManualScrapeDialogProps { open: boolean @@ -13,12 +14,12 @@ interface ManualScrapeDialogProps { defaultQuery?: string mediaType?: string scopeLabel?: string + episodeArtwork?: boolean onClose: () => void onApplied?: () => void } const providers = [ - { value: 'all', label: '全部源' }, { value: 'tmdb', label: 'TMDb' }, { value: 'douban', label: '豆瓣' }, { value: 'bangumi', label: 'Bangumi' }, @@ -33,11 +34,13 @@ export function ManualScrapeDialog({ defaultQuery, mediaType, scopeLabel, + episodeArtwork, onClose, onApplied, }: ManualScrapeDialogProps) { const [query, setQuery] = useState('') - const [provider, setProvider] = useState('all') + const [selectedProviders, setSelectedProviders] = useState([]) + const [includeEpisodeArtwork, setIncludeEpisodeArtwork] = useState(false) const [searching, setSearching] = useState(false) const [applyingKey, setApplyingKey] = useState('') const [items, setItems] = useState([]) @@ -50,10 +53,11 @@ export function ManualScrapeDialog({ useEffect(() => { if (!open) return setQuery(defaultQuery || media?.title || '') - setProvider('all') + setSelectedProviders([]) + setIncludeEpisodeArtwork(episodeArtwork ?? false) setItems([]) setApplyingKey('') - }, [defaultQuery, media?.title, open]) + }, [defaultQuery, episodeArtwork, media?.title, open]) if (!open || !media) return null @@ -67,7 +71,7 @@ export function ManualScrapeDialog({ try { const results = await mediaAPI.manualScrapeSearch(media.id, { query: text, - provider, + provider: selectedProviders.length > 0 ? selectedProviders.join(',') : 'all', media_type: mediaType, }) setItems(results) @@ -84,11 +88,12 @@ export function ManualScrapeDialog({ const key = candidateKey(item) setApplyingKey(key) try { + const options = { episode_images: includeEpisodeArtwork } if (targetIds.length > 1) { - const result = await mediaAPI.applyManualScrapeBatch(targetIds, item) + const result = await mediaAPI.applyManualScrapeBatch(targetIds, item, options) toast.success(`已应用到 ${result.applied} 个媒体`) } else { - await mediaAPI.applyManualScrape(media.id, item) + await mediaAPI.applyManualScrape(media.id, item, options) toast.success('已应用手动匹配') } onApplied?.() @@ -117,15 +122,30 @@ export function ManualScrapeDialog({
- +
+ + {providers.map((item) => { + const active = selectedProviders.includes(item.value) + return ( + + ) + })} +
: } 搜索 + {isEpisodeArtworkTarget(media, mediaType, targetIds.length) && ( + + )}
@@ -190,6 +217,29 @@ function candidateKey(item: ManualScrapeCandidate): string { return `${item.source}:${item.tmdb_id || item.bangumi_id || item.douban_id || item.thetvdb_id || item.title}:${item.media_type || ''}` } +function toggleProvider(value: string, setSelectedProviders: Dispatch>) { + setSelectedProviders((current) => { + if (current.includes(value)) { + return current.filter((item) => item !== value) + } + return [...current, value] + }) +} + +function providerButtonClass(active: boolean): string { + return ( + 'inline-flex h-11 items-center gap-1.5 rounded-xl border px-3 text-xs font-bold transition ' + + (active + ? 'border-brand-300 bg-brand-50 text-brand-700' + : 'border-sand-200 bg-white text-sand-600 hover:border-brand-200 hover:text-brand-600') + ) +} + +function isEpisodeArtworkTarget(media: Media, mediaType?: string, targetCount = 1): boolean { + const type = (mediaType || '').toLowerCase() + return type === 'tv' || type === 'anime' || type === 'variety' || media.season_num > 0 || media.episode_num > 0 || targetCount > 1 +} + function candidateIDText(item: ManualScrapeCandidate): string { const parts = [ item.tmdb_id ? `TMDb ${item.tmdb_id}` : '', diff --git a/web/src/components/MediaCard.tsx b/web/src/components/MediaCard.tsx index ad667ac..2cca6f2 100644 --- a/web/src/components/MediaCard.tsx +++ b/web/src/components/MediaCard.tsx @@ -19,25 +19,26 @@ export const MediaCard = ({ const ref = useRef(null) const href = linkTo ?? `/media/${media.id}` const [posterFit, setPosterFit] = useState<'cover' | 'contain'>('cover') + const posterSrc = imageURL(media.poster_url, media.updated_at) useEffect(() => { setPosterFit('cover') - }, [media.poster_url]) + }, [media.poster_url, media.updated_at]) const card = ( {/* Poster Wrapper */} -
+
{media.poster_url ? ( <> {posterFit === 'contain' && ( )} {media.title} ) : ( -
+
No Poster
@@ -70,7 +71,7 @@ export const MediaCard = ({ {/* Episode count badge */} {count !== undefined && count > 1 && ( - + {count} 集 @@ -78,7 +79,7 @@ export const MediaCard = ({ {/* Rating Badge */} {(rating || (media as any).rating) && ( - + {(rating || (media as any).rating).toFixed(1)} @@ -96,7 +97,7 @@ export const MediaCard = ({ 立即观影 -

+

{media.overview || "暂无简介内容"}

@@ -104,7 +105,7 @@ export const MediaCard = ({ {/* Progress Bar overlay */} {progress !== undefined && progress > 0 && progress < 1 && ( -
+
{/* Media Metadata Info */} -
-

+

+

{media.title}

-
+
{media.year > 0 ? media.year : "未知年份"} {media.video_codec && ( - + {media.video_codec} )} diff --git a/web/src/components/PasswordDialog.tsx b/web/src/components/PasswordDialog.tsx index 63d5f6b..b6d9c93 100644 --- a/web/src/components/PasswordDialog.tsx +++ b/web/src/components/PasswordDialog.tsx @@ -1,28 +1,13 @@ import { FormEvent, useState } from 'react' -import { createRoot } from 'react-dom/client' import { KeyRound } from 'lucide-react' -type PasswordOptions = { +export type PasswordOptions = { title?: string message?: string confirmText?: string } -export function requestPassword(options: PasswordOptions): Promise { - return new Promise((resolve) => { - const host = document.createElement('div') - document.body.appendChild(host) - const root = createRoot(host) - const close = (value: string | null) => { - root.unmount() - host.remove() - resolve(value) - } - root.render() - }) -} - -function PasswordDialog({ +export function PasswordDialog({ options, onClose, }: { diff --git a/web/src/components/PermissionGuard.tsx b/web/src/components/PermissionGuard.tsx index ac7e495..5301c02 100644 --- a/web/src/components/PermissionGuard.tsx +++ b/web/src/components/PermissionGuard.tsx @@ -50,18 +50,3 @@ export function PermissionGuard({ return <>{fallback} } - -// 权限检查工具函数 -export function checkPermission( - permission: string, - isSuper: boolean, - tier: string, - role: string, - permissions: Record -): boolean { - // 超级用户有所有权限 - if (isSuper || tier === 'plus' || role === 'admin') { - return true - } - return permissions[permission] === true -} diff --git a/web/src/components/PinDialog.tsx b/web/src/components/PinDialog.tsx index 736d950..8370410 100644 --- a/web/src/components/PinDialog.tsx +++ b/web/src/components/PinDialog.tsx @@ -1,28 +1,13 @@ import { FormEvent, useState } from 'react' -import { createRoot } from 'react-dom/client' import { LockKeyhole } from 'lucide-react' -type PinOptions = { +export type PinOptions = { title?: string message?: string profileName: string } -export function requestPIN(options: PinOptions): Promise { - return new Promise((resolve) => { - const host = document.createElement('div') - document.body.appendChild(host) - const root = createRoot(host) - const close = (value: string | null) => { - root.unmount() - host.remove() - resolve(value) - } - root.render() - }) -} - -function PinDialog({ +export function PinDialog({ options, onClose, }: { diff --git a/web/src/components/confirmAction.tsx b/web/src/components/confirmAction.tsx new file mode 100644 index 0000000..52cd3ae --- /dev/null +++ b/web/src/components/confirmAction.tsx @@ -0,0 +1,17 @@ +import { createRoot } from 'react-dom/client' + +import { ConfirmDialog, type ConfirmOptions } from './ConfirmDialog' + +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() + }) +} diff --git a/web/src/components/layoutNavigation.ts b/web/src/components/layoutNavigation.ts new file mode 100644 index 0000000..1a39389 --- /dev/null +++ b/web/src/components/layoutNavigation.ts @@ -0,0 +1,126 @@ +import type { LucideIcon } from 'lucide-react' +import { + Activity, + Cast, + Clock, + CloudDownload, + Compass, + Globe, + HardDrive, + Heart, + Home, + Image, + KeySquare, + Library, + ListMusic, + MessageSquareText, + Rss, + Search, + Settings, + Sliders, + Sparkles, + Trash2, + User, + UserCog, +} from 'lucide-react' + +export type LayoutNavGroupID = 'media' | 'personal' | 'downloads' | 'tools' | 'system' + +export type LayoutNavItem = { + to: string + label: string + icon: LucideIcon + end?: boolean + permission?: string + adminOnly?: boolean +} + +export type LayoutNavGroup = { + id: LayoutNavGroupID + label: string + icon: LucideIcon + activePaths: string[] + adminOnly?: boolean + items: LayoutNavItem[] +} + +export const LAYOUT_NAV_GROUPS: LayoutNavGroup[] = [ + { + id: 'media', + label: '媒体浏览', + icon: Home, + activePaths: ['/', '/libraries', '/library', '/poster-wall', '/discover', '/search', '/dlna', '/ai'], + items: [ + { to: '/', label: '系统首页', icon: Home, end: true }, + { to: '/libraries', label: '媒体库', icon: Library }, + { to: '/poster-wall', label: '海报墙', icon: Image }, + { to: '/discover', label: '精彩发现', icon: Compass, permission: 'can_view_discover' }, + { to: '/search', label: '智能搜索', icon: Search, permission: 'can_use_ai' }, + { to: '/dlna', label: 'DLNA 投屏', icon: Cast, permission: 'can_cast' }, + { to: '/ai', label: 'AI 助理', icon: Sparkles, permission: 'can_use_ai_assistant' }, + ], + }, + { + id: 'personal', + label: '个人观影', + icon: User, + activePaths: ['/favourites', '/playlists', '/playlist', '/history', '/profile', '/play-profiles'], + items: [ + { to: '/favourites', label: '我的收藏', icon: Heart }, + { to: '/playlists', label: '播放列表', icon: ListMusic }, + { to: '/history', label: '观看历史', icon: Clock }, + { to: '/profile', label: '账号信息', icon: User }, + { to: '/play-profiles', label: '观影 Profile', icon: UserCog }, + ], + }, + { + id: 'downloads', + label: '下载与订阅', + icon: CloudDownload, + activePaths: ['/downloads', '/download-clients', '/subscriptions', '/site-search'], + items: [ + { to: '/downloads', label: '下载中心', icon: CloudDownload, permission: 'can_manage_downloads' }, + { to: '/subscriptions', label: '订阅管理', icon: Rss, permission: 'can_manage_subscriptions' }, + { to: '/site-search', label: '站点检索', icon: Search, permission: 'can_manage_sites' }, + { to: '/download-clients', label: '下载器管理', icon: Sliders, adminOnly: true }, + ], + }, + { + id: 'tools', + label: '文件与自动化', + icon: HardDrive, + activePaths: ['/storage', '/storage-config', '/files', '/strm', '/duplicates', '/tasks', '/scheduler', '/recycle', '/stats'], + adminOnly: true, + items: [ + { to: '/storage', label: '存储与文件', icon: HardDrive }, + { to: '/storage-config', label: '外部存储', icon: CloudDownload }, + { to: '/files', label: '文件管理', icon: Library }, + { to: '/strm', label: 'STRM 管理', icon: Cast }, + { to: '/duplicates', label: '重复文件', icon: Image }, + { to: '/tasks', label: '任务队列', icon: Activity }, + { to: '/scheduler', label: '计划任务', icon: Clock }, + { to: '/recycle', label: '回收站', icon: Trash2 }, + { to: '/stats', label: '运行状态', icon: Activity }, + ], + }, + { + id: 'system', + label: '系统配置', + icon: Settings, + activePaths: ['/admin', '/sites', '/notify-channels', '/license', '/settings', '/assistant'], + adminOnly: true, + items: [ + { to: '/admin', label: '媒体与用户', icon: Settings }, + { to: '/sites', label: '站点管理', icon: Globe }, + { to: '/notify-channels', label: '通知渠道', icon: MessageSquareText }, + { to: '/assistant', label: 'AI 会话', icon: Sparkles }, + { to: '/license', label: '授权许可', icon: KeySquare }, + { to: '/settings', label: '系统设置', icon: Sliders }, + ], + }, +] + +export const NAV_GROUP_PATHS: Record = LAYOUT_NAV_GROUPS.reduce( + (paths, group) => ({ ...paths, [group.id]: group.activePaths }), + {} as Record, +) diff --git a/web/src/components/requestPIN.tsx b/web/src/components/requestPIN.tsx new file mode 100644 index 0000000..a374656 --- /dev/null +++ b/web/src/components/requestPIN.tsx @@ -0,0 +1,17 @@ +import { createRoot } from 'react-dom/client' + +import { PinDialog, type PinOptions } from './PinDialog' + +export function requestPIN(options: PinOptions): Promise { + return new Promise((resolve) => { + const host = document.createElement('div') + document.body.appendChild(host) + const root = createRoot(host) + const close = (value: string | null) => { + root.unmount() + host.remove() + resolve(value) + } + root.render() + }) +} diff --git a/web/src/components/requestPassword.tsx b/web/src/components/requestPassword.tsx new file mode 100644 index 0000000..b096d7d --- /dev/null +++ b/web/src/components/requestPassword.tsx @@ -0,0 +1,17 @@ +import { createRoot } from 'react-dom/client' + +import { PasswordDialog, type PasswordOptions } from './PasswordDialog' + +export function requestPassword(options: PasswordOptions): Promise { + return new Promise((resolve) => { + const host = document.createElement('div') + document.body.appendChild(host) + const root = createRoot(host) + const close = (value: string | null) => { + root.unmount() + host.remove() + resolve(value) + } + root.render() + }) +} diff --git a/web/src/components/useLayoutPermissions.ts b/web/src/components/useLayoutPermissions.ts new file mode 100644 index 0000000..3ed4c0f --- /dev/null +++ b/web/src/components/useLayoutPermissions.ts @@ -0,0 +1,25 @@ +import { useCallback, useEffect } from 'react' + +import { usePermissionStore } from '../stores/permissions' +import type { User } from '../types' + +export function useLayoutPermissions(user: User | null | undefined) { + const permissions = usePermissionStore((state) => state.permissions) + const isSuper = usePermissionStore((state) => state.isSuper) + const isPermissionLoading = usePermissionStore((state) => state.isLoading) + const fetchPermissions = usePermissionStore((state) => state.fetchPermissions) + + useEffect(() => { + if (user && !isPermissionLoading && Object.keys(permissions ?? {}).length === 0) { + fetchPermissions().catch(() => undefined) + } + }, [fetchPermissions, isPermissionLoading, permissions, user]) + + const isAdmin = user?.role === 'admin' + const can = useCallback( + (key: string) => isAdmin || isSuper || (permissions ?? {})[key] === true, + [isAdmin, isSuper, permissions], + ) + + return { can, isAdmin } +} diff --git a/web/src/components/useLayoutProfiles.ts b/web/src/components/useLayoutProfiles.ts new file mode 100644 index 0000000..d46003d --- /dev/null +++ b/web/src/components/useLayoutProfiles.ts @@ -0,0 +1,86 @@ +import { useCallback, useEffect, useMemo, useState } from 'react' +import toast from 'react-hot-toast' + +import { playProfilesAPI } from '../api/play_profiles' +import type { PlayProfile, User } from '../types' +import { requestPIN } from './requestPIN' + +type UseLayoutProfilesOptions = { + activeProfileId: string | null + setActiveProfile: (id: string | null, pinToken?: string | null) => void + user: User | null | undefined +} + +export function useLayoutProfiles({ + activeProfileId, + setActiveProfile, + user, +}: UseLayoutProfilesOptions) { + const [isProfileOpen, setIsProfileOpen] = useState(false) + const [profiles, setProfiles] = useState([]) + + useEffect(() => { + if (!user) { + setProfiles([]) + setActiveProfile(null) + return + } + playProfilesAPI + .list() + .then((rows) => { + setProfiles(rows) + const active = rows.find((profile) => profile.id === activeProfileId) + if (!active) { + const defaultProfile = rows.find((profile) => profile.is_default && !profile.require_pin) + setActiveProfile(defaultProfile?.id ?? null) + } + }) + .catch(() => undefined) + }, [activeProfileId, setActiveProfile, user]) + + const activeProfile = useMemo( + () => profiles.find((profile) => profile.id === activeProfileId) ?? null, + [activeProfileId, profiles], + ) + + const switchProfile = useCallback( + 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) + } + }, + [activeProfileId, setActiveProfile], + ) + + const useDefaultProfile = useCallback(() => { + setActiveProfile(null) + setIsProfileOpen(false) + }, [setActiveProfile]) + + return { + activeProfile, + isProfileOpen, + profiles, + setIsProfileOpen, + switchProfile, + useDefaultProfile, + } +} diff --git a/web/src/components/useLayoutSearch.ts b/web/src/components/useLayoutSearch.ts new file mode 100644 index 0000000..5cf0485 --- /dev/null +++ b/web/src/components/useLayoutSearch.ts @@ -0,0 +1,85 @@ +import { useEffect, useMemo, useRef, useState, type FormEvent } from 'react' +import type { NavigateFunction } from 'react-router-dom' + +import { mediaAPI } from '../api/library' +import type { Media } from '../types' +import { groupSeries } from '../utils/groupSeries' + +type UseLayoutSearchOptions = { + pathname: string + locationSearch: string + navigate: NavigateFunction +} + +export function useLayoutSearch({ pathname, locationSearch, navigate }: UseLayoutSearchOptions) { + const [focused, setFocused] = useState(false) + const [query, setQuery] = useState('') + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [total, setTotal] = useState(0) + const [error, setError] = useState('') + const searchSeq = useRef(0) + const cards = useMemo(() => groupSeries(items).slice(0, 8), [items]) + + useEffect(() => { + if (pathname === '/search') { + setQuery(new URLSearchParams(locationSearch).get('q') ?? '') + } + }, [pathname, locationSearch]) + + useEffect(() => { + const trimmedQuery = query.trim() + const seq = ++searchSeq.current + if (!focused || !trimmedQuery) { + setItems([]) + setTotal(0) + setError('') + setLoading(false) + return + } + + setLoading(true) + setError('') + const timer = window.setTimeout(() => { + mediaAPI + .search(trimmedQuery, 24) + .then((data) => { + if (seq !== searchSeq.current) return + setItems(data.items ?? []) + setTotal(data.total ?? (data.items ?? []).length) + }) + .catch(() => { + if (seq !== searchSeq.current) return + setItems([]) + setTotal(0) + setError('搜索失败,请稍后再试') + }) + .finally(() => { + if (seq === searchSeq.current) setLoading(false) + }) + }, 220) + + return () => window.clearTimeout(timer) + }, [focused, query]) + + const submit = (event: FormEvent) => { + event.preventDefault() + const trimmedQuery = query.trim() + if (trimmedQuery) { + navigate(`/search?q=${encodeURIComponent(trimmedQuery)}`) + setFocused(false) + } + } + + return { + cards, + error, + focused, + loading, + query, + total, + setFocused, + setQuery, + submit, + } +} diff --git a/web/src/components/useLayoutSidebar.ts b/web/src/components/useLayoutSidebar.ts new file mode 100644 index 0000000..b7e6587 --- /dev/null +++ b/web/src/components/useLayoutSidebar.ts @@ -0,0 +1,50 @@ +import { useCallback, useEffect, useState } from 'react' + +import { NAV_GROUP_PATHS } from './layoutNavigation' + +export function useLayoutSidebar(pathname: string) { + const [isSidebarOpen, setIsSidebarOpen] = useState(true) + const [isMobileDrawerOpen, setIsMobileDrawerOpen] = useState(false) + const [openGroups, setOpenGroups] = useState>({ media: true }) + + useEffect(() => { + const handleResize = () => { + setIsSidebarOpen(window.innerWidth >= 1024) + } + handleResize() + window.addEventListener('resize', handleResize) + return () => window.removeEventListener('resize', handleResize) + }, []) + + useEffect(() => { + setIsMobileDrawerOpen(false) + }, [pathname]) + + const isRouteIn = useCallback( + (paths: string[]) => + paths.some((path) => (path === '/' ? pathname === '/' : pathname.startsWith(path))), + [pathname], + ) + + const toggleGroup = useCallback( + (key: string) => setOpenGroups((current) => ({ ...current, [key]: !current[key] })), + [], + ) + + useEffect(() => { + const active = Object.entries(NAV_GROUP_PATHS).find(([, paths]) => isRouteIn(paths))?.[0] + if (active) { + setOpenGroups((current) => (current[active] ? current : { ...current, [active]: true })) + } + }, [isRouteIn]) + + return { + isMobileDrawerOpen, + isRouteIn, + isSidebarOpen, + openGroups, + setIsMobileDrawerOpen, + setIsSidebarOpen, + toggleGroup, + } +} diff --git a/web/src/components/useThemeMode.ts b/web/src/components/useThemeMode.ts new file mode 100644 index 0000000..736ed8b --- /dev/null +++ b/web/src/components/useThemeMode.ts @@ -0,0 +1,64 @@ +import { useCallback, useEffect, useState } from 'react' + +export type ThemeMode = 'light' | 'dark' | 'system' +export type ResolvedTheme = 'light' | 'dark' + +const THEME_STORAGE_KEY = 'mediastationgo.theme' +const DARK_QUERY = '(prefers-color-scheme: dark)' + +function readStoredTheme(): ThemeMode { + if (typeof window === 'undefined') return 'system' + const value = window.localStorage.getItem(THEME_STORAGE_KEY) + return value === 'light' || value === 'dark' || value === 'system' ? value : 'system' +} + +function systemTheme(): ResolvedTheme { + if (typeof window === 'undefined') return 'light' + return window.matchMedia(DARK_QUERY).matches ? 'dark' : 'light' +} + +function resolveTheme(mode: ThemeMode): ResolvedTheme { + return mode === 'system' ? systemTheme() : mode +} + +function applyTheme(mode: ThemeMode) { + if (typeof document === 'undefined') return + const resolved = resolveTheme(mode) + const root = document.documentElement + root.dataset.themeMode = mode + root.dataset.theme = resolved + root.style.colorScheme = resolved +} + +export function initializeThemeMode() { + applyTheme(readStoredTheme()) +} + +export function useThemeMode() { + const [mode, setModeState] = useState(() => readStoredTheme()) + const [resolvedTheme, setResolvedTheme] = useState(() => resolveTheme(readStoredTheme())) + + useEffect(() => { + applyTheme(mode) + setResolvedTheme(resolveTheme(mode)) + window.localStorage.setItem(THEME_STORAGE_KEY, mode) + }, [mode]) + + useEffect(() => { + const media = window.matchMedia(DARK_QUERY) + const update = () => { + if (mode === 'system') { + applyTheme(mode) + setResolvedTheme(resolveTheme(mode)) + } + } + media.addEventListener('change', update) + return () => media.removeEventListener('change', update) + }, [mode]) + + const setMode = useCallback((nextMode: ThemeMode) => { + setModeState(nextMode) + }, []) + + return { mode, resolvedTheme, setMode } +} diff --git a/web/src/index.css b/web/src/index.css index a37c674..7d812cd 100644 --- a/web/src/index.css +++ b/web/src/index.css @@ -4,9 +4,71 @@ html, body, #root { height: 100%; } +:root { + --app-bg: #f9fafb; + --app-panel: #ffffff; + --app-panel-soft: #f9fafb; + --app-header-bg: rgba(255, 255, 255, 0.82); + --app-control-bg: rgba(249, 250, 251, 0.72); + --app-hover: #f3f4f6; + --app-active-bg: #111827; + --app-active-text: #ffffff; + --app-active-icon: #c9954a; + --app-command-bg: #111827; + --app-command-text: #ffffff; + --app-tooltip-bg: #111827; + --app-tooltip-text: #ffffff; + --app-brand-soft: rgba(255, 248, 231, 0.9); + --app-brand-text: #9a6a1e; + --app-brand-border: rgba(234, 214, 182, 0.9); + --app-brand-emphasis: rgba(201, 149, 74, 0.18); + --app-danger-soft: rgba(254, 242, 242, 0.9); + --app-text: #111827; + --app-subtle: #374151; + --app-muted: #6b7280; + --app-border: rgba(229, 231, 235, 0.82); + --app-shadow: rgba(17, 24, 39, 0.08); + --app-hero-bg: radial-gradient(circle at 80% 20%, rgba(212, 175, 55, 0.22), transparent 34%), linear-gradient(135deg, #fff7ed, #f8fafc 52%, #eef2ff); + --app-hero-overlay: linear-gradient(90deg, #ffffff 0%, rgba(255,255,255,0.96) 37%, rgba(255,255,255,0.62) 68%, rgba(255,255,255,0.2) 100%); + --app-hero-fade: linear-gradient(to top, #ffffff, transparent); + --app-poster-shell: #ffffff; + --app-poster-empty: linear-gradient(135deg, #f9fafb, #fff7ed); +} + +:root[data-theme='dark'] { + --app-bg: #070a0f; + --app-panel: #0f141d; + --app-panel-soft: #111827; + --app-header-bg: rgba(15, 20, 29, 0.86); + --app-control-bg: rgba(17, 24, 39, 0.72); + --app-hover: #1f2937; + --app-active-bg: #1f2937; + --app-active-text: #f9fafb; + --app-active-icon: #e3b56d; + --app-command-bg: #f9fafb; + --app-command-text: #0b1018; + --app-tooltip-bg: #f9fafb; + --app-tooltip-text: #0b1018; + --app-brand-soft: rgba(201, 149, 74, 0.14); + --app-brand-text: #e3b56d; + --app-brand-border: rgba(201, 149, 74, 0.34); + --app-brand-emphasis: rgba(201, 149, 74, 0.18); + --app-danger-soft: rgba(127, 29, 29, 0.22); + --app-text: #f9fafb; + --app-subtle: #d1d5db; + --app-muted: #9ca3af; + --app-border: rgba(55, 65, 81, 0.82); + --app-shadow: rgba(0, 0, 0, 0.45); + --app-hero-bg: radial-gradient(circle at 78% 18%, rgba(201, 149, 74, 0.18), transparent 34%), linear-gradient(135deg, #0b1018, #101826 54%, #171f2d); + --app-hero-overlay: linear-gradient(90deg, rgba(7,10,15,0.98) 0%, rgba(15,20,29,0.92) 42%, rgba(15,20,29,0.62) 70%, rgba(15,20,29,0.24) 100%); + --app-hero-fade: linear-gradient(to top, #0f141d, transparent); + --app-poster-shell: #141b26; + --app-poster-empty: linear-gradient(135deg, #111827, #1f2937); +} + body { - background-color: #f9fafb; /* Premium Swiss Editorial Light/Slate Background */ - color: #111827; /* Deep Graphite Charcoal Body Text */ + background-color: var(--app-bg); + color: var(--app-text); margin: 0; -webkit-font-smoothing: antialiased; -moz-osx-font-smoothing: grayscale; @@ -15,7 +77,7 @@ body { /* ─── Selection ─── */ ::selection { background: rgba(201, 149, 74, 0.2); - color: #111827; + color: var(--app-text); } /* ─── Focus Ring ─── */ @@ -31,6 +93,77 @@ body { ::-webkit-scrollbar-thumb { background: #e5e7eb; border-radius: 999px; } ::-webkit-scrollbar-thumb:hover { background: #c9954a; } +:root[data-theme='dark'] ::-webkit-scrollbar-thumb { background: #374151; } + +:root[data-theme='dark'] .bg-white, +:root[data-theme='dark'] .bg-gray-50, +:root[data-theme='dark'] .bg-gray-50\/50, +:root[data-theme='dark'] .bg-gray-50\/70, +:root[data-theme='dark'] .bg-white\/80, +:root[data-theme='dark'] .bg-white\/82, +:root[data-theme='dark'] .bg-white\/85 { + background-color: var(--app-panel) !important; +} + +:root[data-theme='dark'] .bg-gray-100, +:root[data-theme='dark'] .bg-gray-100\/50 { + background-color: var(--app-hover) !important; +} + +:root[data-theme='dark'] .bg-\[\#f9fafb\] { + background-color: var(--app-bg) !important; +} + +:root[data-theme='dark'] .text-gray-900, +:root[data-theme='dark'] .text-gray-950, +:root[data-theme='dark'] .text-\[\#111827\], +:root[data-theme='dark'] .text-ink-600, +:root[data-theme='dark'] .text-ink-500, +:root[data-theme='dark'] .text-ink-400, +:root[data-theme='dark'] .text-ink-300, +:root[data-theme='dark'] .text-ink-200, +:root[data-theme='dark'] .text-ink-100 { + color: var(--app-text) !important; +} + +:root[data-theme='dark'] .text-gray-700, +:root[data-theme='dark'] .text-gray-600, +:root[data-theme='dark'] .text-gray-500, +:root[data-theme='dark'] .text-ink-50, +:root[data-theme='dark'] .text-sand-500 { + color: var(--app-muted) !important; +} + +:root[data-theme='dark'] .border-gray-100, +:root[data-theme='dark'] .border-gray-200, +:root[data-theme='dark'] .border-gray-200\/50, +:root[data-theme='dark'] .border-gray-200\/60, +:root[data-theme='dark'] .border-gray-200\/80, +:root[data-theme='dark'] .border-gray-200\/90, +:root[data-theme='dark'] .border-gray-300 { + border-color: var(--app-border) !important; +} + +:root[data-theme='dark'] .shadow-xl, +:root[data-theme='dark'] .shadow-2xl, +:root[data-theme='dark'] .shadow-md, +:root[data-theme='dark'] .shadow-sm { + box-shadow: 0 18px 50px var(--app-shadow) !important; +} + +:root[data-theme='dark'] input, +:root[data-theme='dark'] select, +:root[data-theme='dark'] textarea { + color: var(--app-text) !important; + background-color: var(--app-panel-soft) !important; + border-color: var(--app-border) !important; +} + +:root[data-theme='dark'] input::placeholder, +:root[data-theme='dark'] textarea::placeholder { + color: #6b7280; +} + .scrollbar-hide { -ms-overflow-style: none; scrollbar-width: none; @@ -47,6 +180,10 @@ body { } @layer components { + .theme-hero-bg { background: var(--app-hero-bg); } + .theme-hero-overlay { background: var(--app-hero-overlay); } + .theme-hero-fade { background: var(--app-hero-fade); } + /* ── Unified Swiss-Editorial Cards ── */ .card { @apply rounded-2xl bg-white border border-gray-200/80 shadow-[0_1px_3px_rgba(0,0,0,0.01),0_1px_2px_rgba(0,0,0,0.015)] p-4 sm:p-6; @@ -130,4 +267,28 @@ body { .neon-button { @apply btn-outline; } .input-base { @apply input-field; } .warm-divider { @apply divider; } -} \ No newline at end of file +} + +.episode-artwork-toggle { + background-color: var(--app-panel); + border-color: var(--app-border); + color: var(--app-subtle); +} +.episode-artwork-toggle--on, +.episode-artwork-toggle[data-state='on'] { + background-color: var(--app-brand-soft) !important; + border-color: var(--app-brand-border) !important; + color: var(--app-brand-text) !important; +} +.episode-artwork-toggle__icon, +.episode-artwork-toggle__state { + background-color: var(--app-panel-soft); + color: var(--app-muted); +} +.episode-artwork-toggle--on .episode-artwork-toggle__icon, +.episode-artwork-toggle--on .episode-artwork-toggle__state, +.episode-artwork-toggle[data-state='on'] .episode-artwork-toggle__icon, +.episode-artwork-toggle[data-state='on'] .episode-artwork-toggle__state { + background-color: var(--app-brand-emphasis) !important; + color: var(--app-brand-text) !important; +} diff --git a/web/src/main.tsx b/web/src/main.tsx index 97c27e5..2054a8f 100644 --- a/web/src/main.tsx +++ b/web/src/main.tsx @@ -5,8 +5,11 @@ import { Toaster } from 'react-hot-toast' import App from './App' import { GlobalEvents } from './components/GlobalEvents' +import { initializeThemeMode } from './components/useThemeMode' import './index.css' +initializeThemeMode() + if ('serviceWorker' in navigator && import.meta.env.PROD) { window.addEventListener('load', () => { navigator.serviceWorker.register('/artwork-cache-sw.js').catch(() => undefined) diff --git a/web/src/pages/AIAssistantExternalResults.tsx b/web/src/pages/AIAssistantExternalResults.tsx new file mode 100644 index 0000000..920e9fc --- /dev/null +++ b/web/src/pages/AIAssistantExternalResults.tsx @@ -0,0 +1,99 @@ +import { useState } from 'react' +import { Rss } from 'lucide-react' +import toast from 'react-hot-toast' + +import type { ExternalMediaResult } from '../api/ai' +import { imageURL } from '../api/client' +import { buildSiteSearchFeedURL, buildSubscriptionAliases, subscriptionsAPI } from '../api/subscriptions' + +type AIAssistantExternalResultsProps = { + items: ExternalMediaResult[] +} + +export function AIAssistantExternalResults({ items }: AIAssistantExternalResultsProps) { + const [subscribing, setSubscribing] = useState('') + + if (items.length === 0) return null + + return ( +
+ {items.map((item) => { + const keyword = item.subscribe_keyword || item.title + const key = `${item.source}:${keyword}` + return ( +
+
+
+ {item.poster_url ? ( + {item.title} + ) : null} +
+
+
+ {item.source} + {item.media_type && {item.media_type}} + {item.year ? {item.year} : null} +
+

{item.title}

+

+ {item.overview || `订阅关键词:${keyword}`} +

+ +
+
+
+ ) + })} +
+ ) +} + +async function subscribeExternalItem( + item: ExternalMediaResult, + key: string, + keyword: string, + setSubscribing: (key: string) => void, +) { + setSubscribing(key) + try { + const feed = buildSiteSearchFeedURL(keyword, item.source, buildSubscriptionAliases(item)) + const sub = await subscriptionsAPI.create({ + name: `${item.title} 自动订阅`, + feed_url: feed, + filter: keyword, + media_type: item.media_type, + source: item.source, + poster_url: item.poster_url, + backdrop_url: item.backdrop_url, + overview: item.overview, + original_name: item.original_name, + year: item.year, + total_episodes: item.total_episodes, + enabled: true, + }) + const run = await subscriptionsAPI.runNow(sub.id) + toast.success( + run.queued > 0 + ? `已订阅并加入 ${run.queued} 个下载` + : '已订阅,暂未在 PT 站点找到可下载资源', + ) + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '订阅失败' + toast.error(msg) + } finally { + setSubscribing('') + } +} diff --git a/web/src/pages/AIAssistantHeader.tsx b/web/src/pages/AIAssistantHeader.tsx new file mode 100644 index 0000000..0ee83f8 --- /dev/null +++ b/web/src/pages/AIAssistantHeader.tsx @@ -0,0 +1,40 @@ +import { Sparkles } from 'lucide-react' + +export type AIAssistantStatus = { + enabled: boolean + provider: string + model: string +} + +type AIAssistantHeaderProps = { + status: AIAssistantStatus | null +} + +export function AIAssistantHeader({ status }: AIAssistantHeaderProps) { + return ( +
+
+
+ +
+
+

AI 助手

+

自然语言搜索 · 基于观影历史的智能推荐

+
+
+ {status && ( +
+ + {status.enabled + ? `已连接 · ${status.provider}${status.model ? ' / ' + status.model : ''}` + : '未配置 AI 服务,使用本地规则解析'} +
+ )} +
+ ) +} diff --git a/web/src/pages/AIAssistantPage.tsx b/web/src/pages/AIAssistantPage.tsx index cb1fd99..391a675 100644 --- a/web/src/pages/AIAssistantPage.tsx +++ b/web/src/pages/AIAssistantPage.tsx @@ -1,14 +1,9 @@ -import { FormEvent, useEffect, useMemo, useState } from 'react' import { Link } from 'react-router-dom' -import { Loader2, Rss, Search, Sparkles, Wand2 } from 'lucide-react' -import toast from 'react-hot-toast' -import { aiAPI, type ExternalMediaResult, type SearchIntent } from '../api/ai' -import { imageURL } from '../api/client' -import { buildSiteSearchFeedURL, subscriptionsAPI } from '../api/subscriptions' -import { MediaCard } from '../components/MediaCard' -import type { Media } from '../types' -import { groupSeries, seriesCardLink } from '../utils/groupSeries' +import { AIAssistantHeader } from './AIAssistantHeader' +import { AIAssistantRecommendationsSection } from './AIAssistantRecommendationsSection' +import { AIAssistantSearchSection } from './AIAssistantSearchSection' +import { useAIAssistantPage } from './useAIAssistantPage' // AIAssistantPage exposes the two AI helpers backed by the Go server: // - smart search: parses a natural-language query into a SearchIntent + @@ -20,308 +15,31 @@ import { groupSeries, seriesCardLink } from '../utils/groupSeries' // operation-execute endpoints, so we render the same two capabilities as // a focused two-panel screen. export function AIAssistantPage() { - const [status, setStatus] = useState<{ - enabled: boolean - provider: string - model: string - } | null>(null) - const [query, setQuery] = useState('') - const [searching, setSearching] = useState(false) - const [intent, setIntent] = useState(null) - const [items, setItems] = useState([]) - const [externalItems, setExternalItems] = useState([]) - const [subscribing, setSubscribing] = useState('') - const localCards = useMemo(() => groupSeries(items), [items]) - - const [recs, setRecs] = useState(null) - const [recommending, setRecommending] = useState(false) - - useEffect(() => { - aiAPI - .status() - .then(setStatus) - .catch(() => setStatus({ enabled: false, provider: '', model: '' })) - }, []) - - const onSearch = async (e: FormEvent) => { - e.preventDefault() - if (!query.trim()) return - setSearching(true) - setIntent(null) - setItems([]) - setExternalItems([]) - try { - const r = await aiAPI.smartSearch(query.trim()) - setIntent(r.intent) - setItems(r.items) - setExternalItems(r.external_items ?? []) - if (r.items.length === 0 && (r.external_items ?? []).length === 0) toast('未找到匹配项') - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '搜索失败' - toast.error(msg) - } finally { - setSearching(false) - } - } - - const onRecommend = async () => { - setRecommending(true) - try { - const titles = await aiAPI.recommend() - setRecs(titles) - if (titles.length === 0) toast('暂无可推荐内容,请先观看一些媒体') - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '获取推荐失败' - toast.error(msg) - } finally { - setRecommending(false) - } - } - - const quickHints = [ - '2023 年的科幻电影', - '评分高的动漫', - '最近添加的纪录片', - '中文剧集', - ] + const assistant = useAIAssistantPage() return (
-
-
-
- -
-
-

AI 助手

-

- 自然语言搜索 · 基于观影历史的智能推荐 -

-
-
- {status && ( -
- - {status.enabled - ? `已连接 · ${status.provider}${status.model ? ' / ' + status.model : ''}` - : '未配置 AI 服务,使用本地规则解析'} -
- )} -
+ - {/* Smart search */} -
-

智能搜索

-
- setQuery(e.target.value)} - /> - -
+ -
- {quickHints.map((h) => ( - - ))} -
- - {intent && ( -
-
解析结果
-
- - 查询: {intent.query || '—'} - - {intent.year !== undefined && intent.year > 0 && ( - - 年份: {intent.year} - - )} - {intent.genre && ( - - 类型: {intent.genre} - - )} - {intent.type && ( - - 分类: {intent.type} - - )} - {intent.sort && ( - - 排序: {intent.sort} - - )} - {intent.language && ( - - 语言: {intent.language} - - )} -
-
- )} - - {localCards.length > 0 && ( -
-
- 本地媒体库 · {localCards.length} 个合集 / {items.length} 个条目 -
-
- {localCards.map((card) => ( - - ))} -
-
- )} - - {externalItems.length > 0 && ( -
- {externalItems.map((item) => { - const keyword = item.subscribe_keyword || item.title - const key = `${item.source}:${keyword}` - return ( -
-
-
- {item.poster_url ? ( - {item.title} - ) : null} -
-
-
- {item.source} - {item.media_type && {item.media_type}} - {item.year ? {item.year} : null} -
-

{item.title}

-

- {item.overview || `订阅关键词:${keyword}`} -

- -
-
-
- ) - })} -
- )} -
- - {/* Recommendations */} -
-
-

为你推荐

- -
-

- 推荐基于你的最近观看历史。点击标题在媒体库中查找。 -

- - {recs && recs.length > 0 && ( -
    - {recs.map((t, i) => ( -
  • - - {t} - - -
  • - ))} -
- )} - - {recs && recs.length === 0 && ( -

- 还没有推荐结果 — 先去看几部片子,我再给你挑。 -

- )} -
+ {/* Decorative footer (mirrors the Vue page hint that AI runs locally). */} - {!status?.enabled && ( + {!assistant.status?.enabled && (

提示: 当前未配置外部 AI Provider,系统将使用本地规则引擎解析查询。 管理员可在 API 配置{' '} diff --git a/web/src/pages/AIAssistantRecommendationsSection.tsx b/web/src/pages/AIAssistantRecommendationsSection.tsx new file mode 100644 index 0000000..39536bd --- /dev/null +++ b/web/src/pages/AIAssistantRecommendationsSection.tsx @@ -0,0 +1,47 @@ +import { Link } from 'react-router-dom' +import { Loader2, Search, Wand2 } from 'lucide-react' + +type AIAssistantRecommendationsSectionProps = { + recs: string[] | null + recommending: boolean + onRecommend: () => void +} + +export function AIAssistantRecommendationsSection({ + recs, + recommending, + onRecommend, +}: AIAssistantRecommendationsSectionProps) { + return ( +

+
+

为你推荐

+ +
+

推荐基于你的最近观看历史。点击标题在媒体库中查找。

+ + {recs && recs.length > 0 && ( +
    + {recs.map((title, index) => ( +
  • + + {title} + + +
  • + ))} +
+ )} + + {recs && recs.length === 0 && ( +

还没有推荐结果 — 先去看几部片子,我再给你挑。

+ )} +
+ ) +} diff --git a/web/src/pages/AIAssistantSearchSection.tsx b/web/src/pages/AIAssistantSearchSection.tsx new file mode 100644 index 0000000..badfc9e --- /dev/null +++ b/web/src/pages/AIAssistantSearchSection.tsx @@ -0,0 +1,138 @@ +import type { FormEvent } from 'react' +import { Loader2, Search } from 'lucide-react' + +import type { ExternalMediaResult, SearchIntent } from '../api/ai' +import { MediaCard } from '../components/MediaCard' +import type { Media } from '../types' +import type { SeriesCard } from '../utils/groupSeries' +import { seriesCardLink } from '../utils/groupSeries' +import { AIAssistantExternalResults } from './AIAssistantExternalResults' + +type AIAssistantSearchSectionProps = { + query: string + searching: boolean + intent: SearchIntent | null + items: Media[] + localCards: SeriesCard[] + externalItems: ExternalMediaResult[] + onSearch: (event: FormEvent) => void + setQuery: (query: string) => void +} + +const quickHints = [ + '2023 年的科幻电影', + '评分高的动漫', + '最近添加的纪录片', + '中文剧集', +] + +export function AIAssistantSearchSection({ + query, + searching, + intent, + items, + localCards, + externalItems, + onSearch, + setQuery, +}: AIAssistantSearchSectionProps) { + return ( +
+

智能搜索

+
+ setQuery(e.target.value)} + /> + +
+ +
+ {quickHints.map((hint) => ( + + ))} +
+ + + + +
+ ) +} + +function SearchIntentSummary({ intent }: { intent: SearchIntent | null }) { + if (!intent) return null + return ( +
+
解析结果
+
+ + 查询: {intent.query || '—'} + + {intent.year !== undefined && intent.year > 0 && ( + + 年份: {intent.year} + + )} + {intent.genre && ( + + 类型: {intent.genre} + + )} + {intent.type && ( + + 分类: {intent.type} + + )} + {intent.sort && ( + + 排序: {intent.sort} + + )} + {intent.language && ( + + 语言: {intent.language} + + )} +
+
+ ) +} + +function LocalMediaResults({ + localCards, + itemCount, +}: { + localCards: SeriesCard[] + itemCount: number +}) { + if (localCards.length === 0) return null + return ( +
+
+ 本地媒体库 · {localCards.length} 个合集 / {itemCount} 个条目 +
+
+ {localCards.map((card) => ( + + ))} +
+
+ ) +} diff --git a/web/src/pages/APIConfigsPage.tsx b/web/src/pages/APIConfigsPage.tsx index ef7fbb9..f38dabc 100644 --- a/web/src/pages/APIConfigsPage.tsx +++ b/web/src/pages/APIConfigsPage.tsx @@ -3,7 +3,7 @@ import toast from 'react-hot-toast' import { Eye, KeyRound, Save, Trash2 } from 'lucide-react' import { apiConfigsAPI, type APIConfig } from '../api/api_configs' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' // APIConfigsPage manages third-party API keys (TMDb / Bangumi / TheTVDB / // Fanart / OpenAI / Douban). Plaintext keys are never returned by the diff --git a/web/src/pages/AdminLibraryPanel.tsx b/web/src/pages/AdminLibraryPanel.tsx new file mode 100644 index 0000000..6718ed3 --- /dev/null +++ b/web/src/pages/AdminLibraryPanel.tsx @@ -0,0 +1,115 @@ +import { FormEvent, useEffect, useState } from 'react' +import toast from 'react-hot-toast' +import { Trash2 } from 'lucide-react' + +import { libraryAPI } from '../api/library' +import type { Library } from '../types' +import { confirmAction } from '../components/confirmAction' + +export function AdminLibraryPanel() { + const [libs, setLibs] = useState([]) + const [name, setName] = useState('') + const [path, setPath] = useState('') + const [type, setType] = useState('movie') + + const refresh = () => libraryAPI.list({ includeHidden: true }).then(setLibs) + useEffect(() => { + refresh().catch(() => undefined) + }, []) + + const handleCreate = async (e: FormEvent) => { + e.preventDefault() + try { + await libraryAPI.create(name, path, type) + toast.success('媒体库已创建') + setName('') + setPath('') + await refresh() + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '创建失败' + toast.error(msg) + } + } + + return ( +
+
+ setName(e.target.value)} + /> + setPath(e.target.value)} + /> +

+ Docker 部署时请优先填写容器内路径,例如 /media/电影、/media/电视剧/国产剧;如果误填 NAS + 宿主机路径,系统会尝试按 compose 挂载自动转换。 +

+ + +
+ +
+ + + + + + + + + + + {libs.map((l) => ( + + + + + + + ))} + +
名称路径类型操作
{l.name}{l.path}{l.type} + + +
+
+
+ ) +} diff --git a/web/src/pages/AdminPage.tsx b/web/src/pages/AdminPage.tsx index ac0041c..8d136dd 100644 --- a/web/src/pages/AdminPage.tsx +++ b/web/src/pages/AdminPage.tsx @@ -1,16 +1,10 @@ -import { FormEvent, useEffect, useState } from 'react' +import { useEffect, useState } from 'react' import { useSearchParams } from 'react-router-dom' -import toast from 'react-hot-toast' -import { KeyRound, Loader2, Pencil, Plus, ShieldCheck, Trash2, UserCheck, UserX, X } from 'lucide-react' -import { adminAPI } from '../api/admin' -import { libraryAPI } from '../api/library' -import { licenseAPI, type LicenseStatus } from '../api/license' -import type { Library, User } from '../types' import { APIConfigsPanel } from '../components/APIConfigsPanel' import { ManagementShortcuts } from '../components/ManagementShortcuts' -import { confirmAction } from '../components/ConfirmDialog' -import { requestPassword } from '../components/PasswordDialog' +import { AdminLibraryPanel } from './AdminLibraryPanel' +import { AdminUsersPanel } from './AdminUsersPanel' type AdminTab = 'library' | 'users' | 'api' @@ -44,10 +38,10 @@ export function AdminPage() { title="统一管理入口" description="侧栏保持精简,完整管理能力统一从这里进入。" items={[ - { to: '/sites', title: '站点管理', description: '维护 PT 站点、认证方式和检索配置' }, - { to: '/download-clients', title: '下载器管理', description: '配置 qBittorrent 等下载器连接', badge: '下载' }, - { to: '/files', title: '手动整理', description: '从下载目录选择文件夹并整理入库' }, - { to: '/storage', title: '存储与文件', description: '查看占用、清理重复项和管理文件' }, + { to: '/sites', title: '站点管理', description: '维护 PT 站点、认证方式和检索配置', group: '站点与下载' }, + { to: '/download-clients', title: '下载器管理', description: '配置 qBittorrent 等下载器连接', badge: '下载', group: '站点与下载' }, + { to: '/files', title: '手动整理', description: '从下载目录选择文件夹并整理入库', group: '文件与入库' }, + { to: '/storage', title: '存储与文件', description: '查看占用、清理重复项和管理文件', group: '文件与入库' }, ]} />
@@ -67,388 +61,9 @@ export function AdminPage() { ))}
- {tab === 'library' && } - {tab === 'users' && } + {tab === 'library' && } + {tab === 'users' && } {tab === 'api' && }
) } - -function LibraryPanel() { - const [libs, setLibs] = useState([]) - const [name, setName] = useState('') - const [path, setPath] = useState('') - const [type, setType] = useState('movie') - - const refresh = () => libraryAPI.list({ includeHidden: true }).then(setLibs) - useEffect(() => { - refresh().catch(() => undefined) - }, []) - - const handleCreate = async (e: FormEvent) => { - e.preventDefault() - try { - await libraryAPI.create(name, path, type) - toast.success('媒体库已创建') - setName('') - setPath('') - await refresh() - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '创建失败' - toast.error(msg) - } - } - - return ( -
-
- setName(e.target.value)} - /> - setPath(e.target.value)} - /> -

- Docker 部署时请优先填写容器内路径,例如 /media/电影、/media/电视剧/国产剧;如果误填 NAS - 宿主机路径,系统会尝试按 compose 挂载自动转换。 -

- - -
- -
- - - - - - - - - - - {libs.map((l) => ( - - - - - - - ))} - -
名称路径类型操作
{l.name}{l.path}{l.type} - - -
-
-
- ) -} - -function UsersPanel() { - const [users, setUsers] = useState([]) - const [licenseStatus, setLicenseStatus] = useState(null) - const [username, setUsername] = useState('') - const [password, setPassword] = useState('') - const [editingID, setEditingID] = useState(null) - const [editingUsername, setEditingUsername] = useState('') - const [resettingPasswordID, setResettingPasswordID] = useState(null) - const refresh = async () => { - const [nextUsers, nextLicense] = await Promise.all([ - adminAPI.listUsers(), - licenseAPI.status().catch(() => null), - ]) - setUsers(nextUsers) - setLicenseStatus(nextLicense) - } - useEffect(() => { - refresh().catch(() => undefined) - }, []) - - const unlimitedUsers = - licenseStatus?.active === true && - (licenseStatus.unlimited_users === true || licenseStatus.max_users == null) - const maxUsers = unlimitedUsers ? null : (licenseStatus?.max_users ?? 20) - const userLimitReached = maxUsers != null && users.length >= maxUsers - const userLimitLabel = unlimitedUsers ? '不限制' : String(maxUsers) - - const handleCreate = async (e: FormEvent) => { - e.preventDefault() - try { - await adminAPI.createUser({ username, password }) - toast.success('用户已添加,默认仅允许浏览与播放媒体') - setUsername('') - setPassword('') - await refresh() - } catch (err: unknown) { - const msg = - userCreateErrorMessage(err) ?? - '添加用户失败' - toast.error(msg) - } - } - - const startEdit = (u: User) => { - setEditingID(u.id) - setEditingUsername(u.username) - } - - const saveEdit = async (id: string) => { - try { - await adminAPI.updateUser(id, { username: editingUsername }) - toast.success('用户名已更新') - setEditingID(null) - await refresh() - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '更新失败' - toast.error(msg) - } - } - - const resetPassword = async (u: User) => { - if (resettingPasswordID) return - const nextPassword = await requestPassword({ - title: `重置 ${u.username} 的密码`, - message: '请输入新的临时密码,至少 6 位。保存后该用户可立即使用新密码登录 Web、Bot 与第三方客户端。', - confirmText: '重置密码', - }) - if (!nextPassword) return - if (nextPassword.length < 6) { - toast.error('新密码至少 6 位') - return - } - setResettingPasswordID(u.id) - try { - await adminAPI.resetUserPassword(u.id, nextPassword) - toast.success('密码已重置') - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '重置密码失败' - toast.error(msg) - } finally { - setResettingPasswordID(null) - } - } - - const toggleStatus = async (u: User) => { - const next = !u.is_active - if (!next && u.is_protected) { - toast.error('受保护管理员不可禁用') - return - } - if ( - !next && - !(await confirmAction({ - title: '禁用用户', - message: `禁用「${u.username}」后,Web 与第三方客户端已有登录也会失效。`, - confirmText: '禁用', - })) - ) { - return - } - try { - await adminAPI.setUserStatus(u.id, next) - toast.success(next ? '用户已解禁' : '用户已禁用') - await refresh() - } catch (err: unknown) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '操作失败' - toast.error(msg) - } - } - - return ( -
-
-
-
-

用户管理

-

- 已创建 {users.length}/{userLimitLabel} 个用户;新增用户默认只有媒体库浏览、播放、外部播放器与第三方客户端观看权限。 -

-
- - 默认管理员不可删除 · 最高权限 - -
- setUsername(e.target.value)} - disabled={userLimitReached} - /> - setPassword(e.target.value)} - disabled={userLimitReached} - /> - -
- -
- - - - - - - - - - - - - {users.map((u) => ( - - - - - - - - - ))} - -
用户名角色状态权限说明最近登录操作
- {editingID === u.id ? ( - setEditingUsername(e.target.value)} - /> - ) : ( - - {u.username} - {u.is_default_admin && } - - )} - {u.role === 'admin' ? '管理员' : '观看用户'} - {u.is_active ? '正常' : '已禁用'} - - {u.role === 'admin' ? '全部管理权限' : '仅浏览/播放/外部播放器,无下载与文件操作'} - - {u.last_login_at ? new Date(u.last_login_at).toLocaleString() : '从未登录'} - - {editingID === u.id ? ( - <> - - - - ) : ( - - )} - - - -
-
-
- ) -} - -function userCreateErrorMessage(err: unknown): string | undefined { - const data = (err as { response?: { data?: { error?: string; max_users?: number } } })?.response?.data - if (!data?.error) return undefined - if (data.error === 'user limit reached' && data.max_users != null) { - return `用户数量已达到授权上限:${data.max_users} 人` - } - return data.error -} diff --git a/web/src/pages/AdminUsersForm.tsx b/web/src/pages/AdminUsersForm.tsx new file mode 100644 index 0000000..2033b41 --- /dev/null +++ b/web/src/pages/AdminUsersForm.tsx @@ -0,0 +1,62 @@ +import { FormEvent } from 'react' +import { Plus } from 'lucide-react' + +type AdminUsersFormProps = { + usersCount: number + userLimitLabel: string + username: string + password: string + userLimitReached: boolean + onUsernameChange: (value: string) => void + onPasswordChange: (value: string) => void + onSubmit: (e: FormEvent) => void +} + +export function AdminUsersForm({ + usersCount, + userLimitLabel, + username, + password, + userLimitReached, + onUsernameChange, + onPasswordChange, + onSubmit, +}: AdminUsersFormProps) { + return ( +
+
+
+

用户管理

+

+ 已创建 {usersCount}/{userLimitLabel} 个用户;新增用户默认只有媒体库浏览、播放、外部播放器与第三方客户端观看权限。 +

+
+ + 默认管理员不可删除 · 最高权限 + +
+ onUsernameChange(e.target.value)} + disabled={userLimitReached} + /> + onPasswordChange(e.target.value)} + disabled={userLimitReached} + /> + +
+ ) +} diff --git a/web/src/pages/AdminUsersPanel.tsx b/web/src/pages/AdminUsersPanel.tsx new file mode 100644 index 0000000..6c426f8 --- /dev/null +++ b/web/src/pages/AdminUsersPanel.tsx @@ -0,0 +1,175 @@ +import { FormEvent, useEffect, useState } from 'react' +import toast from 'react-hot-toast' + +import { adminAPI } from '../api/admin' +import { licenseAPI, type LicenseStatus } from '../api/license' +import type { User } from '../types' +import { confirmAction } from '../components/confirmAction' +import { requestPassword } from '../components/requestPassword' +import { AdminUsersForm } from './AdminUsersForm' +import { AdminUsersTable } from './AdminUsersTable' + +export function AdminUsersPanel() { + const [users, setUsers] = useState([]) + const [licenseStatus, setLicenseStatus] = useState(null) + const [username, setUsername] = useState('') + const [password, setPassword] = useState('') + const [editingID, setEditingID] = useState(null) + const [editingUsername, setEditingUsername] = useState('') + const [resettingPasswordID, setResettingPasswordID] = useState(null) + const refresh = async () => { + const [nextUsers, nextLicense] = await Promise.all([ + adminAPI.listUsers(), + licenseAPI.status().catch(() => null), + ]) + setUsers(nextUsers) + setLicenseStatus(nextLicense) + } + useEffect(() => { + refresh().catch(() => undefined) + const timer = window.setInterval(() => refresh().catch(() => undefined), 10000) + return () => window.clearInterval(timer) + }, []) + + const unlimitedUsers = + licenseStatus?.active === true && + (licenseStatus.unlimited_users === true || licenseStatus.max_users == null) + const maxUsers = unlimitedUsers ? null : (licenseStatus?.max_users ?? 20) + const userLimitReached = maxUsers != null && users.length >= maxUsers + const userLimitLabel = unlimitedUsers ? '不限制' : String(maxUsers) + + const handleCreate = async (e: FormEvent) => { + e.preventDefault() + try { + await adminAPI.createUser({ username, password }) + toast.success('用户已添加,默认仅允许浏览与播放媒体') + setUsername('') + setPassword('') + await refresh() + } catch (err: unknown) { + const msg = + userCreateErrorMessage(err) ?? + '添加用户失败' + toast.error(msg) + } + } + + const startEdit = (u: User) => { + setEditingID(u.id) + setEditingUsername(u.username) + } + + const saveEdit = async (id: string) => { + try { + await adminAPI.updateUser(id, { username: editingUsername }) + toast.success('用户名已更新') + setEditingID(null) + await refresh() + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '更新失败' + toast.error(msg) + } + } + + const resetPassword = async (u: User) => { + if (resettingPasswordID) return + const nextPassword = await requestPassword({ + title: `重置 ${u.username} 的密码`, + message: '请输入新的临时密码,至少 6 位。保存后该用户可立即使用新密码登录 Web、Bot 与第三方客户端。', + confirmText: '重置密码', + }) + if (!nextPassword) return + if (nextPassword.length < 6) { + toast.error('新密码至少 6 位') + return + } + setResettingPasswordID(u.id) + try { + await adminAPI.resetUserPassword(u.id, nextPassword) + toast.success('密码已重置') + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '重置密码失败' + toast.error(msg) + } finally { + setResettingPasswordID(null) + } + } + + const toggleStatus = async (u: User) => { + const next = !u.is_active + if (!next && u.is_protected) { + toast.error('受保护管理员不可禁用') + return + } + if ( + !next && + !(await confirmAction({ + title: '禁用用户', + message: `禁用「${u.username}」后,Web 与第三方客户端已有登录也会失效。`, + confirmText: '禁用', + })) + ) { + return + } + try { + await adminAPI.setUserStatus(u.id, next) + toast.success(next ? '用户已解禁' : '用户已禁用') + await refresh() + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '操作失败' + toast.error(msg) + } + } + + const deleteUser = async (u: User) => { + if (u.is_protected) return + if (!(await confirmAction({ title: '删除用户', message: `确定删除「${u.username}」?`, confirmText: '删除' }))) return + await adminAPI.deleteUser(u.id) + toast.success('已删除') + await refresh() + } + + return ( +
+ + + setEditingID(null)} + onStartEdit={startEdit} + onResetPassword={resetPassword} + onToggleStatus={toggleStatus} + onDeleteUser={deleteUser} + /> +
+ ) +} + +function userCreateErrorMessage(err: unknown): string | undefined { + const data = (err as { response?: { data?: { error?: string; max_users?: number } } })?.response?.data + if (!data?.error) return undefined + if (data.error === 'user limit reached' && data.max_users != null) { + return `用户数量已达到授权上限:${data.max_users} 人` + } + return data.error +} diff --git a/web/src/pages/AdminUsersTable.tsx b/web/src/pages/AdminUsersTable.tsx new file mode 100644 index 0000000..9bcebb8 --- /dev/null +++ b/web/src/pages/AdminUsersTable.tsx @@ -0,0 +1,136 @@ +import { KeyRound, Loader2, Pencil, ShieldCheck, Trash2, UserCheck, UserX, X } from 'lucide-react' + +import type { User } from '../types' + +type AdminUsersTableProps = { + users: User[] + editingID: string | null + editingUsername: string + resettingPasswordID: string | null + onEditingUsernameChange: (value: string) => void + onSaveEdit: (id: string) => void + onCancelEdit: () => void + onStartEdit: (user: User) => void + onResetPassword: (user: User) => void + onToggleStatus: (user: User) => void + onDeleteUser: (user: User) => void +} + +export function AdminUsersTable({ + users, + editingID, + editingUsername, + resettingPasswordID, + onEditingUsernameChange, + onSaveEdit, + onCancelEdit, + onStartEdit, + onResetPassword, + onToggleStatus, + onDeleteUser, +}: AdminUsersTableProps) { + return ( +
+ + + + + + + + + + + + + {users.map((u) => ( + + + + + + + + + ))} + +
用户名角色状态权限说明最近登录操作
+ {editingID === u.id ? ( + onEditingUsernameChange(e.target.value)} + /> + ) : ( + + {u.username} + {u.is_default_admin && } + + )} + {u.role === 'admin' ? '管理员' : '观看用户'} + {u.is_active ? '正常' : '已禁用'} + + {u.role === 'admin' ? '全部管理权限' : '仅浏览/播放/外部播放器,无下载与文件操作'} + + + {u.last_login_at ? new Date(u.last_login_at).toLocaleString() : '从未登录'} + {u.realtime_online && 在线} + {(u.realtime_device_count ?? 0) > 0 && {u.realtime_device_count} 台} + + + {editingID === u.id ? ( + <> + + + + ) : ( + + )} + + + +
+
+ ) +} diff --git a/web/src/pages/AssistantChatPage.tsx b/web/src/pages/AssistantChatPage.tsx index 6998722..2cfe669 100644 --- a/web/src/pages/AssistantChatPage.tsx +++ b/web/src/pages/AssistantChatPage.tsx @@ -8,7 +8,7 @@ import { type AssistantSession, type SessionView, } from '../api/assistant' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' // AssistantChatPage is the multi-turn chat surface backed by the Go // AssistantService. It complements the older AIAssistantPage which is diff --git a/web/src/pages/AutoOrganizeSettingsPanel.tsx b/web/src/pages/AutoOrganizeSettingsPanel.tsx new file mode 100644 index 0000000..21d426b --- /dev/null +++ b/web/src/pages/AutoOrganizeSettingsPanel.tsx @@ -0,0 +1,123 @@ +import { RefreshCw } from 'lucide-react' + +import { + type AutoOrganizeConfig, + type AutoOrganizeTab, +} from './autoOrganizeModel' +import { + AutoOrganizeBasicTab, + AutoOrganizeNamingTab, + AutoOrganizeScrapeTab, +} from './AutoOrganizeSettingsTabs' + +type AutoOrganizeSettingsPanelProps = { + config: AutoOrganizeConfig + currentDir: string + activeTab: AutoOrganizeTab + dirty: boolean + loading: boolean + saving: boolean + running: boolean + moveKeepsSeeding: boolean + onRefresh: () => void + onSave: () => void + onRunNow: () => void + onTabChange: (tab: AutoOrganizeTab) => void + onConfigChange: (key: keyof AutoOrganizeConfig, value: string) => void +} + +const AUTO_ORGANIZE_TABS: Array<[AutoOrganizeTab, string]> = [ + ['basic', '基础设置'], + ['naming', '命名规则'], + ['scrape', '刮削联动'], +] + +export function AutoOrganizeSettingsPanel({ + config, + currentDir, + activeTab, + dirty, + loading, + saving, + running, + moveKeepsSeeding, + onRefresh, + onSave, + onRunNow, + onTabChange, + onConfigChange, +}: AutoOrganizeSettingsPanelProps) { + return ( +
+
+
+

自动整理设置

+

+ 设置后可自动递归扫描下载/待整理目录,整理到媒体库目录;也可以在这里立即执行一次。 +

+
+
+ + + +
+
+ +
+ {AUTO_ORGANIZE_TABS.map(([key, label]) => ( + + ))} + + {dirty ? '有未保存设置' : '设置已同步'} · 定时任务名:organize_source + +
+ + {activeTab === 'basic' && ( + + )} + {activeTab === 'naming' && ( + + )} + {activeTab === 'scrape' && ( + + )} +
+ ) +} diff --git a/web/src/pages/AutoOrganizeSettingsTabs.tsx b/web/src/pages/AutoOrganizeSettingsTabs.tsx new file mode 100644 index 0000000..613c228 --- /dev/null +++ b/web/src/pages/AutoOrganizeSettingsTabs.tsx @@ -0,0 +1,210 @@ +import { + type AutoOrganizeConfig, + settingOn, +} from './autoOrganizeModel' + +type ConfigChangeHandler = (key: keyof AutoOrganizeConfig, value: string) => void + +type AutoOrganizeTabProps = { + config: AutoOrganizeConfig + onConfigChange: ConfigChangeHandler +} + +export function AutoOrganizeBasicTab({ + config, + currentDir, + moveKeepsSeeding, + onConfigChange, +}: AutoOrganizeTabProps & { + currentDir: string + moveKeepsSeeding: boolean +}) { + return ( + <> +
+ + + + +
+ + {moveKeepsSeeding && ( +
+ 当前同时选择了“移动”和“保种”。为避免 qB 做种源文件被删除,后端会实际使用硬链接;Docker / NAS + 多挂载或不同子卷下可能报 invalid cross-device link。需要真正移动时请关闭“保种”,需要保种但硬链接失败时请选择“复制”。 +
+ )} + +
+ + + + + +
+ + ) +} + +export function AutoOrganizeNamingTab({ config, onConfigChange }: AutoOrganizeTabProps) { + return ( +
+ + + +

+ 可用占位符:{'{title}'} {'{year}'} {'{season}'} {'{season:02}'} {'{episode}'} {'{episode:02}'} {'{category}'}。扩展名会自动补齐。 +

+
+ ) +} + +export function AutoOrganizeScrapeTab({ config, onConfigChange }: AutoOrganizeTabProps) { + return ( + <> +
+ + +
+
+ + + + +
+ + ) +} + +function BooleanSetting({ + config, + settingKey, + label, + onConfigChange, +}: AutoOrganizeTabProps & { + settingKey: keyof AutoOrganizeConfig + label: string +}) { + return ( + + ) +} + +function TextSetting({ + config, + settingKey, + label, + className = 'input-base w-full', + placeholder, + onConfigChange, +}: AutoOrganizeTabProps & { + settingKey: keyof AutoOrganizeConfig + label: string + className?: string + placeholder?: string +}) { + return ( + + ) +} + +function NumberSetting({ + config, + settingKey, + label, + onConfigChange, +}: AutoOrganizeTabProps & { + settingKey: keyof AutoOrganizeConfig + label: string +}) { + return ( + + ) +} diff --git a/web/src/pages/CloudBrowser.tsx b/web/src/pages/CloudBrowser.tsx new file mode 100644 index 0000000..5ff409b --- /dev/null +++ b/web/src/pages/CloudBrowser.tsx @@ -0,0 +1,272 @@ +import { useEffect, useState } from 'react' +import toast from 'react-hot-toast' + +import { libraryAPI } from '../api/library' +import { + cloudAPI, + storageAPI, + type CloudEntry, + type CloudScanStatus, + type StorageType, +} from '../api/storage_config' +import { confirmAction } from '../components/confirmAction' +import type { Library } from '../types' +import { CloudBrowserToolbar } from './CloudBrowserToolbar' +import { CloudEntryList } from './CloudEntryList' +import { CloudMountList } from './CloudMountList' +import { CloudScanPanel } from './CloudScanPanel' +import { + TYPE_LABEL, + cloudLibraryProvider, + cloudMountDisplayPath, +} from './storageConfigModel' + +// Lists cloud directories and imports a file as a 302-backed media. +export function CloudBrowser({ type }: { type: StorageType }) { + const [stack, setStack] = useState<{ id: string; name: string }[]>([{ id: '', name: '根目录' }]) + const [items, setItems] = useState([]) + const [mounts, setMounts] = useState([]) + const [loading, setLoading] = useState(false) + const [mounting, setMounting] = useState(false) + const [batchMounting, setBatchMounting] = useState(false) + const [scanBusy, setScanBusy] = useState(false) + const [cancelBusy, setCancelBusy] = useState(false) + const [scanStatuses, setScanStatuses] = useState([]) + const [mountMediaType, setMountMediaType] = useState('auto') + const [error, setError] = useState('') + + const cur = stack[stack.length - 1] + const load = async (dir: string) => { + setLoading(true) + setError('') + try { + const r = await cloudAPI.list(type, dir) + setItems(r.items ?? []) + if (r.error) setError(r.error) + } catch (err: unknown) { + setError((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '加载失败') + setItems([]) + } finally { + setLoading(false) + } + } + + const loadMounts = async () => { + const libs = await libraryAPI.list({ includeHidden: true }) + setMounts(libs.filter((lib) => cloudLibraryProvider(lib.path) === type)) + } + + const loadScanStatus = async () => { + const r = await storageAPI.cloudScanStatus() + setScanStatuses((r.items ?? []).filter((item) => !type || item.provider === type)) + } + + useEffect(() => { + load(cur.id).catch(() => undefined) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [stack.length, type]) + + useEffect(() => { + loadMounts().catch(() => undefined) + loadScanStatus().catch(() => undefined) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [type]) + + useEffect(() => { + const timer = window.setInterval(() => { + loadScanStatus().catch(() => undefined) + }, 3000) + return () => window.clearInterval(timer) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [type]) + + const enter = (entry: CloudEntry) => setStack((current) => [...current, { id: entry.id, name: entry.name }]) + const goTo = (index: number) => setStack((current) => current.slice(0, index + 1)) + const goUp = () => setStack((current) => (current.length > 1 ? current.slice(0, -1) : current)) + const currentMountPath = () => cloudMountDisplayPath(type, stack) + const childMountPath = (child: CloudEntry) => cloudMountDisplayPath(type, stack, child) + const currentDir = () => stack[stack.length - 1]?.id ?? '' + + const handleMountResult = (res: unknown, label: string) => { + const out = res as { already_mounted?: boolean; skipped?: boolean; reason?: string; library?: Library; message?: string; estimate_message?: string } + if (out.skipped) { + toast(`已跳过「${label}」:和已挂载目录重叠`) + return 'skipped' + } + if (out.already_mounted) { + toast(`「${label}」已经挂载,后台会刷新扫描并自动入库`) + return 'mounted' + } + toast.success(`已挂载「${label}」,${out.message ?? '后台会递归扫描并自动加入媒体库'}。${out.estimate_message ?? ''}`) + return 'mounted' + } + + const doImport = async (entry: CloudEntry) => { + const ref = type === 'cloud115' ? entry.pick_code || entry.id : entry.id + try { + await cloudAPI.import(type, ref, entry.name, entry.size) + toast.success(`已导入「${entry.name}」,可在媒体库中 302 播放`) + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '导入失败') + } + } + + const normalizeFolderInput = (value: string | null) => { + const name = (value ?? '').trim() + if (!name) return '' + if (name === '.' || name === '..' || /[\\/]/.test(name)) { + toast.error('文件夹名称不能包含路径分隔符') + return '' + } + return name + } + + const createFolder = async () => { + const name = normalizeFolderInput(window.prompt('新建文件夹名称') ?? '') + if (!name) return + setLoading(true) + try { + await cloudAPI.mkdir(type, currentDir(), name) + toast.success(`已新建文件夹「${name}」`) + await load(currentDir()) + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '新建文件夹失败') + } finally { + setLoading(false) + } + } + + const renameFolder = async (entry: CloudEntry) => { + if (!entry.is_dir) return + const name = normalizeFolderInput(window.prompt('重命名文件夹', entry.name) ?? '') + if (!name || name === entry.name) return + setLoading(true) + try { + await cloudAPI.rename(type, entry.id, name) + toast.success(`已重命名为「${name}」`) + await load(currentDir()) + await loadMounts() + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '重命名失败') + } finally { + setLoading(false) + } + } + + const mountCurrent = async () => { + setMounting(true) + try { + const label = TYPE_LABEL[type] ?? type + const name = cur.id ? cur.name : label + const res = await cloudAPI.mount(type, cur.id, name, mountMediaType, currentMountPath()) + handleMountResult(res, cur.name) + await loadMounts() + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '挂载失败') + } finally { + setMounting(false) + } + } + + const mountVisibleDirectories = async () => { + const dirs = items.filter((item) => item.is_dir) + if (dirs.length === 0) { + toast.error('当前目录下没有可挂载的子目录') + return + } + setBatchMounting(true) + let ok = 0 + let skipped = 0 + let failed = 0 + for (const dir of dirs) { + try { + const result = await cloudAPI.mount(type, dir.id, dir.name, 'auto', childMountPath(dir)) + const state = handleMountResult(result, dir.name) + if (state === 'skipped') skipped += 1 + else ok += 1 + } catch { + failed += 1 + } + } + if (failed > 0) { + toast.error(`已挂载 ${ok} 个目录,跳过 ${skipped} 个重叠目录,失败 ${failed} 个`) + } else { + toast.success(`已挂载 ${ok} 个目录,跳过 ${skipped} 个重叠目录,后台会自动生成 302/STRM 播放入口`) + } + await loadMounts() + setBatchMounting(false) + } + + const removeMount = async (lib: Library) => { + const ok = await confirmAction({ + title: '移除网盘挂载', + message: `仅移除「${lib.name}」在本项目中的媒体库和媒体记录,不会删除网盘文件。`, + confirmText: '移除', + }) + if (!ok) return + await libraryAPI.remove(lib.id) + toast.success('已移除挂载') + await loadMounts() + } + + const scanAllCloudLibraries = async () => { + setScanBusy(true) + try { + const r = await storageAPI.scanAllCloud() + setScanStatuses(r.items ?? []) + toast.success(r.message ?? '已开始扫描所有启用的网盘媒体库') + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '启动扫描失败') + } finally { + setScanBusy(false) + } + } + + const cancelCloudScans = async () => { + setCancelBusy(true) + try { + const r = await storageAPI.cancelCloudScan('', type) + toast.success(r.message ?? `已中断 ${r.cancelled} 个扫描任务`) + await loadScanStatus() + } catch (err: unknown) { + toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '中断扫描失败') + } finally { + setCancelBusy(false) + } + } + + return ( +
+ + + item.is_dir)} + onGoTo={goTo} + onGoUp={goUp} + onCreateFolder={createFolder} + onMediaTypeChange={setMountMediaType} + onMountCurrent={mountCurrent} + onMountVisibleDirectories={mountVisibleDirectories} + /> + +
+ ) +} diff --git a/web/src/pages/CloudBrowserToolbar.tsx b/web/src/pages/CloudBrowserToolbar.tsx new file mode 100644 index 0000000..3a3c77f --- /dev/null +++ b/web/src/pages/CloudBrowserToolbar.tsx @@ -0,0 +1,99 @@ +import { ArrowUp, FolderPlus } from 'lucide-react' + +interface CloudBrowserToolbarProps { + stack: { id: string; name: string }[] + mountMediaType: string + mounting: boolean + batchMounting: boolean + loading: boolean + hasDirectories: boolean + onGoTo: (index: number) => void + onGoUp: () => void + onCreateFolder: () => void + onMediaTypeChange: (value: string) => void + onMountCurrent: () => void + onMountVisibleDirectories: () => void +} + +export function CloudBrowserToolbar({ + stack, + mountMediaType, + mounting, + batchMounting, + loading, + hasDirectories, + onGoTo, + onGoUp, + onCreateFolder, + onMediaTypeChange, + onMountCurrent, + onMountVisibleDirectories, +}: CloudBrowserToolbarProps) { + return ( +
+
+ 网盘资源: + {stack.map((item, index) => ( + + + {index < stack.length - 1 && /} + + ))} +
+

+ 挂载后不会复制网盘文件;后台会递归读取该目录里的子文件夹和媒体文件,扫描到的影片会自动加入对应媒体库。小目录通常几十秒,大目录取决于网盘接口速度。 + 如果已有同名同类型媒体库,会在首页和 Emby/SenPlayer 中自动归并显示。 +

+
+ + + + + +
+
+ ) +} diff --git a/web/src/pages/CloudEntryList.tsx b/web/src/pages/CloudEntryList.tsx new file mode 100644 index 0000000..8dd35dc --- /dev/null +++ b/web/src/pages/CloudEntryList.tsx @@ -0,0 +1,61 @@ +import { FileVideo, Folder, Loader2, Pencil } from 'lucide-react' + +import type { CloudEntry } from '../api/storage_config' + +interface CloudEntryListProps { + loading: boolean + error: string + items: CloudEntry[] + onEnter: (entry: CloudEntry) => void + onImport: (entry: CloudEntry) => void + onRename: (entry: CloudEntry) => void +} + +export function CloudEntryList({ loading, error, items, onEnter, onImport, onRename }: CloudEntryListProps) { + if (loading) { + return ( +
+ +
+ ) + } + + if (error) return

{error}

+ if (items.length === 0) return

该目录为空

+ + return ( +
    + {items.map((entry) => ( +
  • + {entry.is_dir ? : } + {entry.is_dir ? ( + <> + + + + ) : ( + <> + {entry.name} + + + )} +
  • + ))} +
+ ) +} diff --git a/web/src/pages/CloudMountList.tsx b/web/src/pages/CloudMountList.tsx new file mode 100644 index 0000000..c99ee2d --- /dev/null +++ b/web/src/pages/CloudMountList.tsx @@ -0,0 +1,36 @@ +import { Trash2 } from 'lucide-react' + +import type { Library } from '../types' +import { cloudLibraryLabel } from './storageConfigModel' + +interface CloudMountListProps { + mounts: Library[] + onRemove: (library: Library) => void +} + +export function CloudMountList({ mounts, onRemove }: CloudMountListProps) { + if (mounts.length === 0) return null + + return ( +
+
已挂载目录
+
+ {mounts.map((lib) => ( +
+ + {lib.name} · {cloudLibraryLabel(lib.path)} + + +
+ ))} +
+
+ ) +} diff --git a/web/src/pages/CloudScanPanel.tsx b/web/src/pages/CloudScanPanel.tsx new file mode 100644 index 0000000..1258386 --- /dev/null +++ b/web/src/pages/CloudScanPanel.tsx @@ -0,0 +1,70 @@ +import { Loader2, PauseCircle, RefreshCw } from 'lucide-react' + +import type { CloudScanStatus } from '../api/storage_config' + +interface CloudScanPanelProps { + scanBusy: boolean + cancelBusy: boolean + scanStatuses: CloudScanStatus[] + onScanAll: () => void + onCancelScans: () => void +} + +export function CloudScanPanel({ + scanBusy, + cancelBusy, + scanStatuses, + onScanAll, + onCancelScans, +}: CloudScanPanelProps) { + return ( +
+
+
+
网盘媒体库扫描
+

+ 只需在系统设置填写公开域名,扫描会自动为网盘媒体生成 STRM/302 播放入口;中断后再次扫描会去重补齐。 +

+
+
+ + +
+
+ {scanStatuses.length > 0 && ( +
+ {scanStatuses.slice(0, 6).map((item) => ( +
+ {item.state} + {' · '} + {item.provider} + {' · 目录 '} + {item.dirs} + {' · 发现 '} + {item.discovered} + {' · 入库 '} + {item.added + item.updated} + {item.error ? · {item.error} : null} +
+ ))} +
+ )} +
+ ) +} diff --git a/web/src/pages/DiscoverContentRow.tsx b/web/src/pages/DiscoverContentRow.tsx new file mode 100644 index 0000000..a4c5560 --- /dev/null +++ b/web/src/pages/DiscoverContentRow.tsx @@ -0,0 +1,94 @@ +import { Info } from 'lucide-react' + +import type { DiscoverItem } from '../api/discover' +import { imageURL } from '../api/client' +import { discoverItemSource } from './discoverPageModel' + +export function ContentRow({ + title, + items, + onSelect, +}: { + title: string + items: DiscoverItem[] + onSelect: (item: DiscoverItem) => void +}) { + return ( +
+

{title}

+
+ {items.map((item, index) => ( + + ))} +
+
+ ) +} + +export function DiscoverSkeleton() { + return ( +
+ {[1, 2, 3].map((section) => ( +
+
+
+ {[1, 2, 3, 4, 5, 6, 7, 8].map((item) => ( +
+ ))} +
+
+ ))} +
+ ) +} + +function DiscoverCard({ item, onSelect }: { item: DiscoverItem; onSelect: (item: DiscoverItem) => void }) { + const source = discoverItemSource(item) + return ( + + ) +} + +function discoverKey(item: DiscoverItem, index: number): string { + return `${item.source || 'source'}:${item.tmdb_id || item.douban_id || item.bangumi_id || item.title}:${index}` +} diff --git a/web/src/pages/DiscoverDetailModal.tsx b/web/src/pages/DiscoverDetailModal.tsx new file mode 100644 index 0000000..773338f --- /dev/null +++ b/web/src/pages/DiscoverDetailModal.tsx @@ -0,0 +1,218 @@ +import { useState } from 'react' +import toast from 'react-hot-toast' +import { Download, Rss, X } from 'lucide-react' + +import type { DiscoverItem } from '../api/discover' +import { imageURL } from '../api/client' +import { buildSiteSearchFeedURL, buildSubscriptionAliases, subscriptionsAPI } from '../api/subscriptions' +import { buildSubscribeKeyword, discoverItemSource } from './discoverPageModel' + +export function DiscoverDetailModal({ item, onClose }: { item: DiscoverItem; onClose: () => void }) { + const source = discoverItemSource(item) + const keyword = item.subscribe_keyword || buildSubscribeKeyword(item) + const [form, setForm] = useState({ + keyword, + search_mode: 'keyword', + imdb_id: '', + media_type: item.media_type || '', + resolution: 'best', + quality: '', + effects: '', + release_groups: '', + exclude_words: 'cam,ts,tc,枪版', + wash_enabled: false, + wash_priority: 'balanced', + save_path: '', + media_category: '', + priority: 50, + run_now: true, + }) + const [busy, setBusy] = useState(false) + + const submit = async () => { + const finalKeyword = form.keyword.trim() || keyword + const feed = buildSiteSearchFeedURL(finalKeyword, source, buildSubscriptionAliases(item)) + setBusy(true) + try { + const sub = await subscriptionsAPI.create({ + name: `${item.title} 自动订阅`, + feed_url: feed, + filter: finalKeyword, + media_type: form.media_type || undefined, + media_category: form.media_category || undefined, + save_path: form.save_path || undefined, + search_mode: form.search_mode, + imdb_id: form.imdb_id || undefined, + source, + poster_url: item.poster_url || undefined, + backdrop_url: item.backdrop_url || undefined, + overview: item.overview || undefined, + original_name: item.original_name || undefined, + year: item.year || undefined, + resolution: form.resolution === 'best' ? 'best' : form.resolution, + quality: form.quality || undefined, + effects: form.effects || undefined, + release_groups: form.release_groups || undefined, + exclude_words: form.exclude_words || undefined, + wash_enabled: form.wash_enabled, + wash_priority: form.wash_priority, + priority: form.priority, + enabled: true, + }) + if (form.run_now) { + const run = await subscriptionsAPI.runNow(sub.id) + toast.success(run.queued > 0 ? `已订阅并加入 ${run.queued} 个下载` : '已订阅,暂未命中可下载资源') + } else { + toast.success('已创建订阅') + } + onClose() + } catch (err) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + '订阅失败' + toast.error(msg) + } finally { + setBusy(false) + } + } + + return ( +
+
+
+
+

{source}

+

{item.title}

+

+ {[item.media_type, item.year && item.year > 0 ? item.year : '', item.rating ? `★ ${item.rating.toFixed(1)}` : ''] + .filter(Boolean) + .join(' · ')} +

+
+ +
+ +
+
+
+ {item.poster_url ? ( + {item.title} + ) : ( +
无海报
+ )} +
+ {item.backdrop_url && ( + + )} +
+ +
+
+

简介

+

{item.overview || '当前数据源没有返回简介。'}

+
+ +
+

+ + 订阅下载规则 +

+
+ + + + + + + + + + + + + +
+
+ + +
+
+
+
+
+
+ ) +} diff --git a/web/src/pages/DiscoverPage.tsx b/web/src/pages/DiscoverPage.tsx index 4b4a14e..a0a9efd 100644 --- a/web/src/pages/DiscoverPage.tsx +++ b/web/src/pages/DiscoverPage.tsx @@ -1,24 +1,21 @@ import { useEffect, useMemo, useState } from 'react' -import toast from 'react-hot-toast' -import { AlertTriangle, Download, Info, Rss, Sparkles, X } from 'lucide-react' +import { AlertTriangle, Sparkles } from 'lucide-react' import { discoverAPI, type DiscoverItem, type DiscoverSection } from '../api/discover' -import { imageURL } from '../api/client' -import { buildSiteSearchFeedURL, subscriptionsAPI } from '../api/subscriptions' - -const defaultSections = [ - 'tmdb_trending_day', - 'douban_hot_movie', - 'douban_hot_tv', - 'bangumi_calendar', -] - -const storageKey = 'mediastation.discover.sections' +import { ContentRow, DiscoverSkeleton } from './DiscoverContentRow' +import { DiscoverDetailModal } from './DiscoverDetailModal' +import { + defaultSectionDefs, + defaultSections, + discoverStorageKey, + readSavedSections, +} from './discoverPageModel' export function DiscoverPage() { const [sections, setSections] = useState([]) const [selected, setSelected] = useState(defaultSections) const [rows, setRows] = useState>({}) + const [rowErrors, setRowErrors] = useState>({}) const [error, setError] = useState('') const [loading, setLoading] = useState(true) const [activeItem, setActiveItem] = useState(null) @@ -29,7 +26,9 @@ export function DiscoverPage() { .then((items) => { setSections(items) const saved = readSavedSections(items) - setSelected(saved.length > 0 ? saved : defaultSections) + const available = new Set(items.map((item) => item.key)) + const fallback = defaultSections.filter((key) => available.has(key)) + setSelected(saved.length > 0 ? saved : fallback) }) .catch(() => { setSections(defaultSectionDefs) @@ -40,26 +39,46 @@ export function DiscoverPage() { useEffect(() => { if (selected.length === 0) { setRows({}) + setRowErrors({}) setLoading(false) return } + let cancelled = false setLoading(true) setError('') - window.localStorage.setItem(storageKey, JSON.stringify(selected)) - discoverAPI - .feed(selected) - .then((feed) => { - const next: Record = {} - for (const key of selected) { - next[key] = feed[key] ?? [] - } - setRows(next) - }) - .catch((err) => { - setRows({}) - setError(err instanceof Error ? err.message : String(err)) - }) - .finally(() => setLoading(false)) + setRowErrors({}) + setRows((current) => { + const next: Record = {} + for (const key of selected) { + next[key] = current[key] ?? [] + } + return next + }) + window.localStorage.setItem(discoverStorageKey, JSON.stringify(selected)) + + let pending = selected.length + const markDone = () => { + pending -= 1 + if (!cancelled && pending <= 0) setLoading(false) + } + for (const key of selected) { + discoverAPI + .feed([key]) + .then((feed) => { + if (cancelled) return + setRows((current) => ({ ...current, [key]: feed[key] ?? [] })) + }) + .catch((err) => { + if (cancelled) return + const message = err instanceof Error ? err.message : String(err) + setRows((current) => ({ ...current, [key]: [] })) + setRowErrors((current) => ({ ...current, [key]: message })) + }) + .finally(markDone) + } + return () => { + cancelled = true + } }, [selected]) const sectionMap = useMemo( @@ -67,6 +86,7 @@ export function DiscoverPage() { [sections], ) const hasContent = selected.some((key) => (rows[key] ?? []).length > 0) + const hasRowErrors = Object.keys(rowErrors).length > 0 const toggleSection = (key: string) => { setSelected((current) => { @@ -116,7 +136,7 @@ export function DiscoverPage() {
- {loading && } + {loading && !hasContent && } {!loading && error && (
@@ -131,7 +151,7 @@ export function DiscoverPage() {
)} - {!loading && !error && selected.length > 0 && ( + {!error && selected.length > 0 && (hasContent || !loading) && (
{selected.map((key) => { const items = rows[key] ?? [] @@ -146,7 +166,18 @@ export function DiscoverPage() { ) })} - {!hasContent && ( + {hasRowErrors && ( +
+ +
+ {Object.entries(rowErrors).map(([key, message]) => ( +

{sectionMap.get(key)?.label ?? key}:{message}

+ ))} +
+
+ )} + + {!loading && !hasContent && (

当前选择的推荐源暂未返回内容,可切换豆瓣 / Bangumi 或检查网络代理。 @@ -165,325 +196,3 @@ export function DiscoverPage() {

) } - -function ContentRow({ - title, - items, - onSelect, -}: { - title: string - items: DiscoverItem[] - onSelect: (item: DiscoverItem) => void -}) { - return ( -
-

{title}

-
- {items.map((item, index) => ( - - ))} -
-
- ) -} - -function DiscoverCard({ item, onSelect }: { item: DiscoverItem; onSelect: (item: DiscoverItem) => void }) { - const source = item.source || (item.bangumi_id ? 'bangumi' : item.douban_id ? 'douban' : 'tmdb') - return ( - - ) -} - -function DiscoverDetailModal({ item, onClose }: { item: DiscoverItem; onClose: () => void }) { - const source = item.source || (item.bangumi_id ? 'bangumi' : item.douban_id ? 'douban' : 'tmdb') - const keyword = item.subscribe_keyword || buildSubscribeKeyword(item) - const [form, setForm] = useState({ - keyword, - search_mode: 'keyword', - imdb_id: '', - media_type: item.media_type || '', - resolution: 'best', - quality: '', - effects: '', - release_groups: '', - exclude_words: 'cam,ts,tc,枪版', - wash_enabled: false, - wash_priority: 'balanced', - save_path: '', - media_category: '', - priority: 50, - run_now: true, - }) - const [busy, setBusy] = useState(false) - - const submit = async () => { - const finalKeyword = form.keyword.trim() || keyword - const feed = buildSiteSearchFeedURL(finalKeyword, source, [item.title, item.original_name || '']) - setBusy(true) - try { - const sub = await subscriptionsAPI.create({ - name: `${item.title} 自动订阅`, - feed_url: feed, - filter: finalKeyword, - media_type: form.media_type || undefined, - media_category: form.media_category || undefined, - save_path: form.save_path || undefined, - search_mode: form.search_mode, - imdb_id: form.imdb_id || undefined, - source, - poster_url: item.poster_url || undefined, - backdrop_url: item.backdrop_url || undefined, - overview: item.overview || undefined, - resolution: form.resolution === 'best' ? 'best' : form.resolution, - quality: form.quality || undefined, - effects: form.effects || undefined, - release_groups: form.release_groups || undefined, - exclude_words: form.exclude_words || undefined, - wash_enabled: form.wash_enabled, - wash_priority: form.wash_priority, - priority: form.priority, - enabled: true, - }) - if (form.run_now) { - const run = await subscriptionsAPI.runNow(sub.id) - toast.success(run.queued > 0 ? `已订阅并加入 ${run.queued} 个下载` : '已订阅,暂未命中可下载资源') - } else { - toast.success('已创建订阅') - } - onClose() - } catch (err) { - const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? - '订阅失败' - toast.error(msg) - } finally { - setBusy(false) - } - } - - return ( -
-
-
-
-

{source}

-

{item.title}

-

- {[item.media_type, item.year && item.year > 0 ? item.year : '', item.rating ? `★ ${item.rating.toFixed(1)}` : ''] - .filter(Boolean) - .join(' · ')} -

-
- -
- -
-
-
- {item.poster_url ? ( - {item.title} - ) : ( -
无海报
- )} -
- {item.backdrop_url && ( - - )} -
- -
-
-

简介

-

{item.overview || '当前数据源没有返回简介。'}

-
- -
-

- - 订阅下载规则 -

-
- - - - - - - - - - - - - -
-
- - -
-
-
-
-
-
- ) -} - -function buildSubscribeKeyword(item: DiscoverItem): string { - return [item.title, item.year && item.year > 0 ? item.year : ''].filter(Boolean).join(' ') -} - -function DiscoverSkeleton() { - return ( -
- {[1, 2, 3].map((section) => ( -
-
-
- {[1, 2, 3, 4, 5, 6, 7, 8].map((item) => ( -
- ))} -
-
- ))} -
- ) -} - -function discoverKey(item: DiscoverItem, index: number): string { - return `${item.source || 'source'}:${item.tmdb_id || item.douban_id || item.bangumi_id || item.title}:${index}` -} - -function readSavedSections(sections: DiscoverSection[]): string[] { - try { - const raw = window.localStorage.getItem(storageKey) - if (!raw) return [] - const parsed = JSON.parse(raw) - if (!Array.isArray(parsed)) return [] - const allowed = new Set(sections.map((section) => section.key)) - return parsed.filter((key) => typeof key === 'string' && allowed.has(key)) - } catch { - return [] - } -} - -const defaultSectionDefs: DiscoverSection[] = [ - { key: 'tmdb_trending_day', label: 'TMDb 今日趋势', provider: 'tmdb' }, - { key: 'tmdb_popular_movie', label: 'TMDb 热门电影', provider: 'tmdb' }, - { key: 'douban_hot_movie', label: '豆瓣热门电影', provider: 'douban' }, - { key: 'douban_hot_tv', label: '豆瓣热门剧集', provider: 'douban' }, - { key: 'bangumi_calendar', label: 'Bangumi 每日放送', provider: 'bangumi' }, -] diff --git a/web/src/pages/DownloadClientCard.tsx b/web/src/pages/DownloadClientCard.tsx new file mode 100644 index 0000000..86afd13 --- /dev/null +++ b/web/src/pages/DownloadClientCard.tsx @@ -0,0 +1,70 @@ +import { Loader2, Pencil, Send, Trash2 } from 'lucide-react' + +import type { DownloadClient } from '../api/download_clients' + +export function DownloadClientCard({ + client, + testing, + onDelete, + onEdit, + onTest, +}: { + client: DownloadClient + testing: boolean + onDelete: (client: DownloadClient) => void + onEdit: (client: DownloadClient) => void + onTest: (id: string) => void +}) { + return ( +
+
+
+ {client.name} + + {client.type} + + {client.is_default && ( + + 默认 + + )} + {!client.enabled && ( + + 已禁用 + + )} +
+
+ {client.host} + {client.username && ` · ${client.username}`} +
+
+
+ + + +
+
+ ) +} diff --git a/web/src/pages/DownloadClientFormModal.tsx b/web/src/pages/DownloadClientFormModal.tsx new file mode 100644 index 0000000..55b6cd9 --- /dev/null +++ b/web/src/pages/DownloadClientFormModal.tsx @@ -0,0 +1,170 @@ +import { type FormEvent, type ReactNode, useState } from 'react' +import { Loader2 } from 'lucide-react' +import toast from 'react-hot-toast' + +import { + downloadClientsAPI, + type DownloadClient, + type DownloadClientInput, + type DownloadClientType, +} from '../api/download_clients' +import { apiErrorMessage } from './downloadClientPageModel' + +export function ClientFormModal({ + editing, + onClose, + onSaved, +}: { + editing: DownloadClient | null + onClose: () => void + onSaved: () => void | Promise +}) { + const [form, setForm] = useState(() => ({ + name: editing?.name ?? '', + type: editing?.type ?? 'qbittorrent', + host: editing?.host ?? '', + username: editing?.username ?? '', + password: '', + is_default: editing?.is_default ?? false, + enabled: editing?.enabled ?? true, + })) + const [saving, setSaving] = useState(false) + + const onSubmit = async (e: FormEvent) => { + e.preventDefault() + if (saving) return + setSaving(true) + try { + if (editing) await downloadClientsAPI.update(editing.id, form) + else await downloadClientsAPI.create(form) + toast.success('已保存') + setSaving(false) + onSaved() + } catch (err: unknown) { + const msg = apiErrorMessage(err, '保存失败') + toast.error(msg) + setSaving(false) + } + } + + const update = (k: K, v: DownloadClientInput[K]) => + setForm((f) => ({ ...f, [k]: v })) + + const placeholder = ( + { + qbittorrent: 'http://127.0.0.1:8080', + aria2: 'http://127.0.0.1:6800/jsonrpc', + transmission: 'http://127.0.0.1:9091/transmission/rpc', + } as Record + )[form.type] + + return ( +
+
+

+ {editing ? '编辑下载器' : '添加下载器'} +

+
+ + update('name', e.target.value)} + /> + + + + + + update('host', e.target.value)} + /> + + {form.type !== 'aria2' && ( + <> + + update('username', e.target.value)} + /> + + + update('password', e.target.value)} + /> + + + )} + {form.type === 'aria2' && ( + + update('password', e.target.value)} + /> + + )} +
+ + +
+
+ + +
+
+
+
+ ) +} + +function Field({ label, children }: { label: string; children: ReactNode }) { + return ( + + ) +} diff --git a/web/src/pages/DownloadClientsPage.tsx b/web/src/pages/DownloadClientsPage.tsx index 439024a..7f2dde6 100644 --- a/web/src/pages/DownloadClientsPage.tsx +++ b/web/src/pages/DownloadClientsPage.tsx @@ -1,14 +1,15 @@ -import { FormEvent, useEffect, useState } from 'react' -import { Loader2, Pencil, Plus, Send, Server, Trash2 } from 'lucide-react' +import { useEffect, useState } from 'react' +import { Loader2, Plus, Server } from 'lucide-react' import toast from 'react-hot-toast' import { downloadClientsAPI, type DownloadClient, - type DownloadClientInput, - type DownloadClientType, } from '../api/download_clients' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' +import { DownloadClientCard } from './DownloadClientCard' +import { ClientFormModal } from './DownloadClientFormModal' +import { apiErrorMessage } from './downloadClientPageModel' // DownloadClientsPage manages multiple downloader integrations. // Replaces the Vue UI's DownloadView "clients" tab with a typed CRUD @@ -98,62 +99,17 @@ export function DownloadClientsPage() { {!loading && clients.length > 0 && (
{clients.map((c) => ( -
-
-
- {c.name} - - {c.type} - - {c.is_default && ( - - 默认 - - )} - {!c.enabled && ( - - 已禁用 - - )} -
-
- {c.host} - {c.username && ` · ${c.username}`} -
-
-
- - - -
-
+ client={c} + testing={Boolean(testing[c.id])} + onDelete={onDelete} + onEdit={(client) => { + setEditing(client) + setShowForm(true) + }} + onTest={onTest} + /> ))}
)} @@ -171,170 +127,3 @@ export function DownloadClientsPage() {
) } - -function ClientFormModal({ - editing, - onClose, - onSaved, -}: { - editing: DownloadClient | null - onClose: () => void - onSaved: () => void | Promise -}) { - const [form, setForm] = useState(() => ({ - name: editing?.name ?? '', - type: editing?.type ?? 'qbittorrent', - host: editing?.host ?? '', - username: editing?.username ?? '', - password: '', - is_default: editing?.is_default ?? false, - enabled: editing?.enabled ?? true, - })) - const [saving, setSaving] = useState(false) - - const onSubmit = async (e: FormEvent) => { - e.preventDefault() - if (saving) return - setSaving(true) - try { - if (editing) await downloadClientsAPI.update(editing.id, form) - else await downloadClientsAPI.create(form) - toast.success('已保存') - setSaving(false) - onSaved() - } catch (err: unknown) { - const msg = apiErrorMessage(err, '保存失败') - toast.error(msg) - setSaving(false) - } - } - - const update = (k: K, v: DownloadClientInput[K]) => - setForm((f) => ({ ...f, [k]: v })) - - const placeholder = ( - { - qbittorrent: 'http://127.0.0.1:8080', - aria2: 'http://127.0.0.1:6800/jsonrpc', - transmission: 'http://127.0.0.1:9091/transmission/rpc', - } as Record - )[form.type] - - return ( -
-
-

- {editing ? '编辑下载器' : '添加下载器'} -

-
- - update('name', e.target.value)} - /> - - - - - - update('host', e.target.value)} - /> - - {form.type !== 'aria2' && ( - <> - - update('username', e.target.value)} - /> - - - update('password', e.target.value)} - /> - - - )} - {form.type === 'aria2' && ( - - update('password', e.target.value)} - /> - - )} -
- - -
-
- - -
-
-
-
- ) -} - -function apiErrorMessage(err: unknown, fallback: string): string { - const data = (err as { response?: { data?: { error?: string; message?: string } } })?.response?.data - if (data?.error) return data.error - if (data?.message) return data.message - if ((err as { code?: string })?.code === 'ECONNABORTED') return '请求超时,请检查服务或网络' - return fallback -} - -function Field({ label, children }: { label: string; children: React.ReactNode }) { - return ( - - ) -} diff --git a/web/src/pages/DownloadTaskCard.tsx b/web/src/pages/DownloadTaskCard.tsx new file mode 100644 index 0000000..cccab82 --- /dev/null +++ b/web/src/pages/DownloadTaskCard.tsx @@ -0,0 +1,205 @@ +import type { ReactNode } from 'react' +import { ArrowDown, ArrowUp, Film, HardDrive, Rss, Trash2 } from 'lucide-react' + +import { imageURL } from '../api/client' +import type { DownloadTask, QBitTorrent } from '../types' + +export type DownloadCardItem = { + id?: string + hash?: string + title: string + poster_url?: string + backdrop_url?: string + overview?: string + save_path?: string + status?: string + state?: string + progress: number + dlspeed?: number + upspeed?: number + num_seeds?: number + num_leechs?: number + size?: number + downloaded?: number + created_at?: string + updated_at?: string +} + +export function toLiveCard(t: QBitTorrent): DownloadCardItem { + return { + hash: t.hash, + title: t.title || t.name || '下载任务', + poster_url: t.poster_url, + backdrop_url: t.backdrop_url, + overview: t.overview, + save_path: t.save_path, + state: t.state, + progress: t.progress, + dlspeed: t.dlspeed, + upspeed: t.upspeed, + num_seeds: t.num_seeds, + num_leechs: t.num_leechs, + size: t.size, + downloaded: t.downloaded, + } +} + +export function toTaskCard(t: DownloadTask): DownloadCardItem { + return { + id: t.id, + title: t.title || '下载任务', + poster_url: t.poster_url, + backdrop_url: t.backdrop_url, + overview: t.overview, + save_path: t.save_path, + status: t.status, + state: t.state, + progress: t.progress, + dlspeed: t.dlspeed, + upspeed: t.upspeed, + num_seeds: t.num_seeds, + num_leechs: t.num_leechs, + size: t.size, + downloaded: t.downloaded, + created_at: t.created_at, + updated_at: t.updated_at, + } +} + +export function DownloadTaskCard({ + item, + removable, + onRemove, +}: { + item: DownloadCardItem + removable?: boolean + onRemove?: () => Promise +}) { + const progress = pct(item.progress) + const visual = item.poster_url || item.backdrop_url + const downloaded = item.downloaded || (item.size ? Math.round(item.size * (item.progress || 0)) : 0) + + return ( +
+
+
+ {visual ? ( + {item.title} + ) : ( +
+ + {item.title} +
+ )} + + {stateLabel(item)} + +
+ +
+
+

+ {item.title} +

+

+ {item.overview || item.save_path || '已隐藏原始种子 URL,避免泄露私有 Token。'} +

+
+ +
+
+ 进度 {progress.toFixed(1)}% + {fmtBytes(downloaded)} / {fmtBytes(item.size)} +
+
+
+
+
+ +
+ } value={fmtSpeed(item.dlspeed)} /> + } value={fmtSpeed(item.upspeed)} /> + } value={`${item.num_seeds ?? 0} / ${item.num_leechs ?? 0}`} /> + } value={fmtBytes(item.size)} /> +
+ +
+ + {item.save_path || '默认下载目录'} + + {item.created_at && {new Date(item.created_at).toLocaleString()}} +
+
+
+ + {removable && onRemove && ( +
+ +
+ )} +
+ ) +} + +function DownloadMetric({ icon, value }: { icon: ReactNode; value: string }) { + return ( +
+ {icon} + {value} +
+ ) +} + +function fmtBytes(n?: number): string { + if (!n || n <= 0) return '0 B' + const u = ['B', 'KB', 'MB', 'GB', 'TB'] + let v = n + let i = 0 + while (v >= 1024 && i < u.length - 1) { + v /= 1024 + i++ + } + return `${v.toFixed(v >= 100 ? 0 : 1)} ${u[i]}` +} + +function fmtSpeed(n?: number): string { + return `${fmtBytes(n)}/s` +} + +function pct(progress?: number): number { + if (!Number.isFinite(progress)) return 0 + return Math.min(100, Math.max(0, Math.round((progress ?? 0) * 1000) / 10)) +} + +function stateLabel(item: DownloadCardItem): string { + const state = (item.state || item.status || 'queued').toLowerCase() + if (state.includes('down') || state.includes('meta')) return '下载中' + if (state.includes('up') || state.includes('seed')) return '做种中' + if (state.includes('pause')) return '已暂停' + if (state.includes('error')) return '出错' + if (state.includes('complete') || pct(item.progress) >= 100) return '已完成' + if (state.includes('queue')) return '排队中' + return item.state || item.status || '等待中' +} + +function statusTone(item: DownloadCardItem): string { + const state = stateLabel(item) + if (state === '已完成' || state === '做种中') return 'bg-emerald-50 text-emerald-600' + if (state === '出错') return 'bg-red-50 text-red-500' + if (state === '已暂停') return 'bg-amber-50 text-amber-600' + return 'bg-primary-400/10 text-brand-500' +} diff --git a/web/src/pages/DownloadsPage.tsx b/web/src/pages/DownloadsPage.tsx index 1c58d9e..9f04a8e 100644 --- a/web/src/pages/DownloadsPage.tsx +++ b/web/src/pages/DownloadsPage.tsx @@ -1,214 +1,13 @@ import { FormEvent, useEffect, useState } from 'react' import { Link } from 'react-router-dom' import toast from 'react-hot-toast' -import { ArrowDown, ArrowUp, Download, Film, HardDrive, Rss, ShieldCheck, Trash2 } from 'lucide-react' +import { Download, ShieldCheck } from 'lucide-react' -import { imageURL } from '../api/client' import { downloadsAPI } from '../api/downloads' import { useAuthStore } from '../stores/auth' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' import type { DownloadTask, QBitTorrent } from '../types' - -type DownloadCardItem = { - id?: string - hash?: string - title: string - poster_url?: string - backdrop_url?: string - overview?: string - save_path?: string - status?: string - state?: string - progress: number - dlspeed?: number - upspeed?: number - num_seeds?: number - num_leechs?: number - size?: number - downloaded?: number - created_at?: string -} - -function fmtBytes(n?: number): string { - if (!n || n <= 0) return '0 B' - const u = ['B', 'KB', 'MB', 'GB', 'TB'] - let v = n - let i = 0 - while (v >= 1024 && i < u.length - 1) { - v /= 1024 - i++ - } - return `${v.toFixed(v >= 100 ? 0 : 1)} ${u[i]}` -} - -function fmtSpeed(n?: number): string { - return `${fmtBytes(n)}/s` -} - -function pct(progress?: number): number { - if (!Number.isFinite(progress)) return 0 - return Math.min(100, Math.max(0, Math.round((progress ?? 0) * 1000) / 10)) -} - -function stateLabel(item: DownloadCardItem): string { - const state = (item.state || item.status || 'queued').toLowerCase() - if (state.includes('down') || state.includes('meta')) return '下载中' - if (state.includes('up') || state.includes('seed')) return '做种中' - if (state.includes('pause')) return '已暂停' - if (state.includes('error')) return '出错' - if (state.includes('complete') || pct(item.progress) >= 100) return '已完成' - if (state.includes('queue')) return '排队中' - return item.state || item.status || '等待中' -} - -function statusTone(item: DownloadCardItem): string { - const state = stateLabel(item) - if (state === '已完成' || state === '做种中') return 'bg-emerald-50 text-emerald-600' - if (state === '出错') return 'bg-red-50 text-red-500' - if (state === '已暂停') return 'bg-amber-50 text-amber-600' - return 'bg-primary-400/10 text-brand-500' -} - -function toLiveCard(t: QBitTorrent): DownloadCardItem { - return { - hash: t.hash, - title: t.title || t.name || '下载任务', - poster_url: t.poster_url, - backdrop_url: t.backdrop_url, - overview: t.overview, - save_path: t.save_path, - state: t.state, - progress: t.progress, - dlspeed: t.dlspeed, - upspeed: t.upspeed, - num_seeds: t.num_seeds, - num_leechs: t.num_leechs, - size: t.size, - downloaded: t.downloaded, - } -} - -function toTaskCard(t: DownloadTask): DownloadCardItem { - return { - id: t.id, - title: t.title || '下载任务', - poster_url: t.poster_url, - backdrop_url: t.backdrop_url, - overview: t.overview, - save_path: t.save_path, - status: t.status, - state: t.state, - progress: t.progress, - dlspeed: t.dlspeed, - upspeed: t.upspeed, - num_seeds: t.num_seeds, - num_leechs: t.num_leechs, - size: t.size, - downloaded: t.downloaded, - created_at: t.created_at, - } -} - -function DownloadCard({ - item, - removable, - onRemove, -}: { - item: DownloadCardItem - removable?: boolean - onRemove?: () => Promise -}) { - const progress = pct(item.progress) - const visual = item.poster_url || item.backdrop_url - const downloaded = item.downloaded || (item.size ? Math.round(item.size * (item.progress || 0)) : 0) - - return ( -
-
-
- {visual ? ( - {item.title} - ) : ( -
- - {item.title} -
- )} - - {stateLabel(item)} - -
- -
-
-

- {item.title} -

-

- {item.overview || item.save_path || '已隐藏原始种子 URL,避免泄露私有 Token。'} -

-
- -
-
- 进度 {progress.toFixed(1)}% - {fmtBytes(downloaded)} / {fmtBytes(item.size)} -
-
-
-
-
- -
-
- - {fmtSpeed(item.dlspeed)} -
-
- - {fmtSpeed(item.upspeed)} -
-
- - {item.num_seeds ?? 0} / {item.num_leechs ?? 0} -
-
- - {fmtBytes(item.size)} -
-
- -
- - {item.save_path || '默认下载目录'} - - {item.created_at && {new Date(item.created_at).toLocaleString()}} -
-
-
- - {removable && onRemove && ( -
- -
- )} -
- ) -} +import { DownloadTaskCard, toLiveCard, toTaskCard } from './DownloadTaskCard' export function DownloadsPage() { const role = useAuthStore((s) => s.user?.role) @@ -301,7 +100,7 @@ export function DownloadsPage() { {torrents && torrents.length > 0 && (
{torrents.map((torrent) => ( - {tasks.map((task) => ( - + ))}
)} diff --git a/web/src/pages/DuplicatesPage.tsx b/web/src/pages/DuplicatesPage.tsx index f3e4efa..f39e385 100644 --- a/web/src/pages/DuplicatesPage.tsx +++ b/web/src/pages/DuplicatesPage.tsx @@ -4,7 +4,7 @@ import { Copy, Trash2 } from 'lucide-react' import { duplicatesAPI, type DuplicateReport } from '../api/duplicates' import { libraryAPI } from '../api/library' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' import type { Library } from '../types' function fmtBytes(n: number): string { diff --git a/web/src/pages/FileBrowserRoots.tsx b/web/src/pages/FileBrowserRoots.tsx new file mode 100644 index 0000000..0554d36 --- /dev/null +++ b/web/src/pages/FileBrowserRoots.tsx @@ -0,0 +1,31 @@ +import { FolderOpen } from 'lucide-react' + +type FileRoot = { + label: string + path: string +} + +type FileBrowserRootsProps = { + roots: FileRoot[] + onOpen: (path: string) => void +} + +export function FileBrowserRoots({ roots, onOpen }: FileBrowserRootsProps) { + return ( +
+ {roots.map((root) => ( + + ))} +
+ ) +} diff --git a/web/src/pages/FileEntriesTable.tsx b/web/src/pages/FileEntriesTable.tsx new file mode 100644 index 0000000..c49506f --- /dev/null +++ b/web/src/pages/FileEntriesTable.tsx @@ -0,0 +1,104 @@ +import { FileVideo, Folder, HardDrive } from 'lucide-react' + +import type { FileEntry } from '../api/files' + +type FileEntriesTableProps = { + basePath: string + entries: FileEntry[] + recursive: boolean + selectedPath?: string + selectedPaths: string[] + onSelectAll: (checked: boolean) => void + onToggleSelectedPath: (entry: FileEntry, checked: boolean) => void + onEnter: (entry: FileEntry) => void + onChoose: (entry: FileEntry) => void +} + +export function FileEntriesTable({ + basePath, + entries, + recursive, + selectedPath, + selectedPaths, + onSelectAll, + onToggleSelectedPath, + onEnter, + onChoose, +}: FileEntriesTableProps) { + return ( +
+ + + + + + + + + + + + {entries.map((entry) => ( + + + + + + + + ))} + +
+ 0 && entries.every((entry) => selectedPaths.includes(entry.path))} + onChange={(event) => onSelectAll(event.target.checked)} + /> + 名称大小修改时间选择
+ onToggleSelectedPath(entry, event.target.checked)} + /> + + + {entry.is_dir ? '—' : fmtBytes(entry.size)}{new Date(entry.modified * 1000).toLocaleString()} + +
+
+ ) +} + +function entryLabel(entry: FileEntry, basePath: string, recursive: boolean): string { + if (!recursive) return entry.name + return entry.path.replace((basePath || '') + '\\', '').replace((basePath || '') + '/', '') +} + +function fmtBytes(n: number): string { + if (!n) return '0 B' + const u = ['B', 'KB', 'MB', 'GB', 'TB'] + let v = n + let i = 0 + while (v >= 1024 && i < u.length - 1) { + v /= 1024 + i++ + } + return `${v.toFixed(1)} ${u[i]}` +} diff --git a/web/src/pages/FileManagerPage.tsx b/web/src/pages/FileManagerPage.tsx index 4dc17f3..ebbd37a 100644 --- a/web/src/pages/FileManagerPage.tsx +++ b/web/src/pages/FileManagerPage.tsx @@ -1,161 +1,27 @@ import { useCallback, useEffect, useMemo, useState } from 'react' import toast from 'react-hot-toast' -import { - ChevronUp, - Copy, - FileVideo, - Folder, - FolderOpen, - GitBranch, - HardDrive, - Home, - Move, - Pencil, - Plus, - RefreshCw, - Trash2, -} from 'lucide-react' import { filesAPI, type FileEntry, type FileListing } from '../api/files' -import { adminAPI } from '../api/admin' import { libraryAPI } from '../api/library' -import { schedulerAPI } from '../api/scheduler' import { toolsAPI } from '../api/tools' -import { confirmAction } from '../components/ConfirmDialog' -import type { Library, Setting } from '../types' - -function fmtBytes(n: number): string { - if (!n) return '0 B' - const u = ['B', 'KB', 'MB', 'GB', 'TB'] - let v = n - let i = 0 - while (v >= 1024 && i < u.length - 1) { - v /= 1024 - i++ - } - return `${v.toFixed(1)} ${u[i]}` -} - -function formatScanSummary(scans: Array<{ name: string; added: number; updated: number; visited: number; error?: string }>): string { - if (scans.length === 0) return ' · 未扫描:没有匹配的媒体库' - const ok = scans.filter((scan) => !scan.error) - const added = ok.reduce((sum, scan) => sum + (scan.added ?? 0), 0) - const updated = ok.reduce((sum, scan) => sum + (scan.updated ?? 0), 0) - const visited = ok.reduce((sum, scan) => sum + (scan.visited ?? 0), 0) - return ` · 扫描 ${ok.length}/${scans.length} 个库 · 新入库 ${added} · 更新 ${updated} · 访问 ${visited}` -} - -function formatScrapeSummary(scrapes: Array<{ name: string; matched: number; skipped?: boolean; reason?: string; error?: string }>): string { - if (scrapes.length === 0) return '' - const ok = scrapes.filter((scrape) => !scrape.error && !scrape.skipped) - const skipped = scrapes.filter((scrape) => scrape.skipped).length - const matched = ok.reduce((sum, scrape) => sum + (scrape.matched ?? 0), 0) - if (ok.length === 0 && skipped > 0) return ` · 刮削跳过 ${skipped} 个库` - return ` · 刮削 ${ok.length}/${scrapes.length} 个库 · 匹配 ${matched}${skipped ? ` · 跳过 ${skipped}` : ''}` -} - -type AutoOrganizeConfig = { - enabled: string - afterDownload: string - scrapeAfter: string - downloadSmartClassify: string - smartClassify: string - sourceDir: string - targetDir: string - transferMode: string - intervalSeconds: string - keepSeeding: string - movieFormat: string - tvFormat: string - animeFormat: string - scrapeAutoOnScan: string - scrapeProviders: string - scrapeLanguage: string - scrapeDelayMinMs: string - scrapeDelayMaxMs: string -} - -const AUTO_ORGANIZE_DEFAULTS: AutoOrganizeConfig = { - enabled: 'false', - afterDownload: 'false', - scrapeAfter: 'true', - downloadSmartClassify: 'true', - smartClassify: 'true', - sourceDir: '', - targetDir: '', - transferMode: 'hardlink', - intervalSeconds: '300', - keepSeeding: 'true', - movieFormat: '{title} ({year})/{title} ({year})', - tvFormat: '{title} ({year})/Season {season:02}/{title} S{season:02}E{episode:02}', - animeFormat: '{title}/Season {season:02}/{title} S{season:02}E{episode:02}', - scrapeAutoOnScan: 'false', - scrapeProviders: 'tmdb,douban,bangumi,thetvdb,fanart', - scrapeLanguage: 'zh-CN', - scrapeDelayMinMs: '250', - scrapeDelayMaxMs: '500', -} - -const AUTO_ORGANIZE_KEYS: Record = { - enabled: 'organize.auto', - afterDownload: 'organizer.auto_after_download', - scrapeAfter: 'organize.scrape_after', - downloadSmartClassify: 'downloads.smart_classify', - smartClassify: 'organizer.smart_classify', - sourceDir: 'organize.source_dir', - targetDir: 'organize.target_dir', - transferMode: 'organize.transfer_mode', - intervalSeconds: 'organize.interval_seconds', - keepSeeding: 'organize.keep_seeding', - movieFormat: 'organize.movie_format', - tvFormat: 'organize.tv_format', - animeFormat: 'organize.anime_format', - scrapeAutoOnScan: 'scrape.auto_on_scan', - scrapeProviders: 'scrape.providers', - scrapeLanguage: 'scrape.language', - scrapeDelayMinMs: 'scrape.delay_min_ms', - scrapeDelayMaxMs: 'scrape.delay_max_ms', -} - -type AutoOrganizeTab = 'basic' | 'naming' | 'scrape' - -function settingIndex(rows: Setting[]): Record { - const out: Record = {} - for (const row of rows) out[row.key] = row.value - return out -} - -function mergeAutoOrganizeSettings(rows: Setting[]): AutoOrganizeConfig { - const idx = settingIndex(rows) - return { - enabled: idx[AUTO_ORGANIZE_KEYS.enabled] ?? AUTO_ORGANIZE_DEFAULTS.enabled, - afterDownload: idx[AUTO_ORGANIZE_KEYS.afterDownload] ?? AUTO_ORGANIZE_DEFAULTS.afterDownload, - scrapeAfter: idx[AUTO_ORGANIZE_KEYS.scrapeAfter] ?? AUTO_ORGANIZE_DEFAULTS.scrapeAfter, - downloadSmartClassify: idx[AUTO_ORGANIZE_KEYS.downloadSmartClassify] ?? AUTO_ORGANIZE_DEFAULTS.downloadSmartClassify, - smartClassify: idx[AUTO_ORGANIZE_KEYS.smartClassify] ?? AUTO_ORGANIZE_DEFAULTS.smartClassify, - sourceDir: idx[AUTO_ORGANIZE_KEYS.sourceDir] ?? AUTO_ORGANIZE_DEFAULTS.sourceDir, - targetDir: idx[AUTO_ORGANIZE_KEYS.targetDir] ?? AUTO_ORGANIZE_DEFAULTS.targetDir, - transferMode: idx[AUTO_ORGANIZE_KEYS.transferMode] ?? AUTO_ORGANIZE_DEFAULTS.transferMode, - intervalSeconds: idx[AUTO_ORGANIZE_KEYS.intervalSeconds] ?? AUTO_ORGANIZE_DEFAULTS.intervalSeconds, - keepSeeding: idx[AUTO_ORGANIZE_KEYS.keepSeeding] ?? AUTO_ORGANIZE_DEFAULTS.keepSeeding, - movieFormat: idx[AUTO_ORGANIZE_KEYS.movieFormat] ?? AUTO_ORGANIZE_DEFAULTS.movieFormat, - tvFormat: idx[AUTO_ORGANIZE_KEYS.tvFormat] ?? AUTO_ORGANIZE_DEFAULTS.tvFormat, - animeFormat: idx[AUTO_ORGANIZE_KEYS.animeFormat] ?? AUTO_ORGANIZE_DEFAULTS.animeFormat, - scrapeAutoOnScan: idx[AUTO_ORGANIZE_KEYS.scrapeAutoOnScan] ?? AUTO_ORGANIZE_DEFAULTS.scrapeAutoOnScan, - scrapeProviders: idx[AUTO_ORGANIZE_KEYS.scrapeProviders] ?? AUTO_ORGANIZE_DEFAULTS.scrapeProviders, - scrapeLanguage: idx[AUTO_ORGANIZE_KEYS.scrapeLanguage] ?? AUTO_ORGANIZE_DEFAULTS.scrapeLanguage, - scrapeDelayMinMs: idx[AUTO_ORGANIZE_KEYS.scrapeDelayMinMs] ?? AUTO_ORGANIZE_DEFAULTS.scrapeDelayMinMs, - scrapeDelayMaxMs: idx[AUTO_ORGANIZE_KEYS.scrapeDelayMaxMs] ?? AUTO_ORGANIZE_DEFAULTS.scrapeDelayMaxMs, - } -} - -function settingOn(value: string): boolean { - return ['1', 'true', 'yes', 'on', 'enabled', '启用', '开启'].includes(value.trim().toLowerCase()) -} - -function isCloudLibraryPath(value: string): boolean { - return value.trim().toLowerCase().startsWith('cloud://') -} +import { confirmAction } from '../components/confirmAction' +import type { Library } from '../types' +import { settingOn } from './autoOrganizeModel' +import { AutoOrganizeSettingsPanel } from './AutoOrganizeSettingsPanel' +import { FileBrowserRoots } from './FileBrowserRoots' +import { FileEntriesTable } from './FileEntriesTable' +import { FileManagerToolbar } from './FileManagerToolbar' +import { FileOperationsPanel } from './FileOperationsPanel' +import { ManualOrganizePanel } from './ManualOrganizePanel' +import { useAutoOrganizeSettings } from './useAutoOrganizeSettings' +import { useFileOperations } from './useFileOperations' +import { + formatScanSummary, + formatScrapeSummary, + isCloudLibraryPath, + summarizeOrganizeResults, + type OrganizePreviewItem, +} from './fileManagerModel' // FileManagerPage provides a focused local storage view: // browse allowed roots, optionally recurse, and perform safe local operations. @@ -166,13 +32,6 @@ export function FileManagerPage() { const [error, setError] = useState('') const [loading, setLoading] = useState(true) const [recursive, setRecursive] = useState(false) - const [selected, setSelected] = useState(null) - const [selectedPaths, setSelectedPaths] = useState([]) - const [folderName, setFolderName] = useState('') - const [renameTo, setRenameTo] = useState('') - const [destPath, setDestPath] = useState('') - const [transferMode, setTransferMode] = useState('copy') - const [busy, setBusy] = useState('') const [organizeLibraryID, setOrganizeLibraryID] = useState('') const [organizeDestPath, setOrganizeDestPath] = useState('') const [organizeTransferMode, setOrganizeTransferMode] = useState('hardlink') @@ -180,18 +39,8 @@ export function FileManagerPage() { const [scanAfter, setScanAfter] = useState(true) const [scrapeAfter, setScrapeAfter] = useState(true) const [organizeBusy, setOrganizeBusy] = useState('') - const [previewItems, setPreviewItems] = useState>([]) - const [autoConfig, setAutoConfig] = useState(AUTO_ORGANIZE_DEFAULTS) - const [autoDirty, setAutoDirty] = useState(false) - const [autoSaving, setAutoSaving] = useState(false) - const [autoRunning, setAutoRunning] = useState(false) - const [autoLoading, setAutoLoading] = useState(true) - const [autoTab, setAutoTab] = useState('basic') + const [previewItems, setPreviewItems] = useState([]) + const autoOrganize = useAutoOrganizeSettings({ onScrapeAfterChange: setScrapeAfter }) const currentDir = useMemo(() => { if (data?.path) return data.path @@ -201,8 +50,8 @@ export function FileManagerPage() { () => libraries.filter((library) => !isCloudLibraryPath(library.path)), [libraries], ) - const autoMoveKeepsSeeding = autoConfig.transferMode === 'move' && settingOn(autoConfig.keepSeeding) - const manualMoveKeepsSeeding = organizeTransferMode === 'move' && settingOn(autoConfig.keepSeeding) + const autoMoveKeepsSeeding = autoOrganize.moveKeepsSeeding + const manualMoveKeepsSeeding = organizeTransferMode === 'move' && settingOn(autoOrganize.config.keepSeeding) const refresh = useCallback(() => { setLoading(true) @@ -227,29 +76,7 @@ export function FileManagerPage() { libraryAPI.list({ includeHidden: true }).then(setLibraries).catch(() => undefined) }, []) - const refreshAutoConfig = useCallback(() => { - setAutoLoading(true) - adminAPI - .listSettings() - .then((rows) => { - const nextConfig = mergeAutoOrganizeSettings(rows) - setAutoConfig(nextConfig) - setScrapeAfter(settingOn(nextConfig.scrapeAfter)) - setAutoDirty(false) - }) - .catch(() => undefined) - .finally(() => setAutoLoading(false)) - }, []) - - useEffect(() => { - refreshAutoConfig() - }, [refreshAutoConfig]) - - useEffect(() => { - setSelected(null) - setSelectedPaths([]) - setRenameTo('') - }, [path]) + const fileOperations = useFileOperations({ currentDir, path, refresh }) useEffect(() => { if (!organizeLibraryID) return @@ -262,132 +89,11 @@ export function FileManagerPage() { setOrganizeMediaType(lib.type || 'auto') }, [localLibraries, organizeLibraryID]) - const changeAutoConfig = (key: keyof AutoOrganizeConfig, value: string) => { - setAutoConfig((current) => ({ ...current, [key]: value })) - if (key === 'scrapeAfter') setScrapeAfter(settingOn(value)) - setAutoDirty(true) - } - - const saveAutoConfig = async (): Promise => { - setAutoSaving(true) - try { - for (const key of Object.keys(AUTO_ORGANIZE_KEYS) as Array) { - await adminAPI.updateSetting(AUTO_ORGANIZE_KEYS[key], autoConfig[key] ?? '') - } - setAutoDirty(false) - toast.success('整理入库设置已保存') - return true - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '保存整理入库设置失败') - return false - } finally { - setAutoSaving(false) - } - } - - const runAutoOrganizeNow = async () => { - if (autoDirty) { - const saved = await saveAutoConfig() - if (!saved) return - } - setAutoRunning(true) - try { - await schedulerAPI.run('organize_source') - toast.success('已触发自动整理任务,请稍后刷新媒体库查看入库结果') - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '触发自动整理失败') - } finally { - setAutoRunning(false) - } - } - const enter = (e: FileEntry) => { if (e.is_dir) setPath(e.path) } - const choose = (e: FileEntry) => { - setSelected(e) - setRenameTo(e.name) - if (e.is_dir) setDestPath(e.path) - } - - const toggleSelectedPath = (entry: FileEntry, checked: boolean) => { - setSelectedPaths((current) => { - if (checked) { - return current.includes(entry.path) ? current : [...current, entry.path] - } - return current.filter((item) => item !== entry.path) - }) - } - - const createFolder = async () => { - if (!currentDir || !folderName.trim()) return - setBusy('mkdir') - try { - const res = await filesAPI.createFolder(currentDir, folderName.trim()) - toast.success(`已创建目录:${res.path}`) - setFolderName('') - refresh() - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '创建失败') - } finally { - setBusy('') - } - } - - const renameSelected = async () => { - if (!selected || !renameTo.trim()) return - setBusy('rename') - try { - const res = await filesAPI.rename(selected.path, renameTo.trim()) - toast.success(`已重命名:${res.path}`) - setSelected(null) - refresh() - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '重命名失败') - } finally { - setBusy('') - } - } - - const deleteSelected = async () => { - if (!selected) return - const ok = await confirmAction({ - title: '确认删除文件', - message: `将删除:${selected.path}。此操作不可恢复,请确认不是媒体库根目录。`, - confirmText: '确认删除', - danger: true, - }) - if (!ok) return - setBusy('delete') - try { - await filesAPI.remove(selected.path) - toast.success('已删除') - setSelected(null) - refresh() - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '删除失败') - } finally { - setBusy('') - } - } - - const transferSelected = async () => { - if (!selected || selected.is_dir || !destPath.trim()) return - setBusy('transfer') - try { - const res = await filesAPI.transfer(selected.path, destPath.trim(), transferMode) - toast.success(`已完成转移:${res.path}`) - if (transferMode === 'move') setSelected(null) - refresh() - } catch (err: unknown) { - toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '转移失败') - } finally { - setBusy('') - } - } - - const organizeSources = selectedPaths.length > 0 ? selectedPaths : [selected?.path || currentDir].filter(Boolean) + const organizeSources = fileOperations.selectedPaths.length > 0 ? fileOperations.selectedPaths : [fileOperations.selected?.path || currentDir].filter(Boolean) const organizeSource = organizeSources.length === 1 ? organizeSources[0] : `${organizeSources.length} 个已选项目` const organizeReady = organizeSources.length > 0 && Boolean(organizeDestPath.trim()) @@ -419,16 +125,9 @@ export function FileManagerPage() { dry_run: dryRun, })) } - const preview = results.flatMap((result) => result.items ?? []) - setPreviewItems(preview) - const organized = results.reduce((sum, result) => sum + (result.organized ?? 0), 0) - const replaced = results.reduce((sum, result) => sum + (result.replaced ?? 0), 0) - const skipped = results.reduce((sum, result) => sum + (result.skipped ?? 0), 0) - const errors = results.flatMap((result) => result.errors ?? []) - const scans = results.flatMap((result) => result.scans ?? []) - const scrapes = results.flatMap((result) => result.scrapes ?? []) - const total = organized + replaced + skipped + errors.length - if (total === 0) { + const summary = summarizeOrganizeResults(results) + setPreviewItems(summary.preview) + if (summary.total === 0) { toast(`未发现可整理视频:${organizeSource}`, { icon: '!', duration: 6000, @@ -436,13 +135,13 @@ export function FileManagerPage() { return } if (dryRun) { - toast.success(`预览完成:新增 ${organized} · 替换 ${replaced} · 跳过 ${skipped}`) + toast.success(`预览完成:新增 ${summary.organized} · 替换 ${summary.replaced} · 纠偏 ${summary.reclassified} · 跳过 ${summary.skipped}`) return } - const scanText = scanAfter ? formatScanSummary(scans) : '' - const scrapeText = scanAfter && scrapeAfter ? formatScrapeSummary(scrapes) : '' - toast.success(`整理完成:新增 ${organized} · 替换 ${replaced} · 跳过 ${skipped}${scanText}${scrapeText}`) - setSelectedPaths([]) + const scanText = scanAfter ? formatScanSummary(summary.scans) : '' + const scrapeText = scanAfter && scrapeAfter ? formatScrapeSummary(summary.scrapes) : '' + toast.success(`整理完成:新增 ${summary.organized} · 替换 ${summary.replaced} · 纠偏 ${summary.reclassified} · 跳过 ${summary.skipped}${scanText}${scrapeText}`) + fileOperations.setSelectedPaths([]) refresh() } catch (err: unknown) { toast.error((err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '整理失败') @@ -460,566 +159,97 @@ export function FileManagerPage() {

-
- - {data?.parent && ( - - )} - - - {data?.path && ( - - {data.path} - - )} -
+ setPath('')} + onParent={setPath} + onRefresh={refresh} + onRecursiveChange={setRecursive} + /> -
-
-
-

自动整理设置

-

- 设置后可自动递归扫描下载/待整理目录,整理到媒体库目录;也可以在这里立即执行一次。 -

-
-
- - - -
-
- -
- {[ - ['basic', '基础设置'], - ['naming', '命名规则'], - ['scrape', '刮削联动'], - ].map(([key, label]) => ( - - ))} - - {autoDirty ? '有未保存设置' : '设置已同步'} · 定时任务名:organize_source - -
- - {autoTab === 'basic' && ( - <> -
- - - - -
- - {autoMoveKeepsSeeding && ( -
- 当前同时选择了“移动”和“保种”。为避免 qB 做种源文件被删除,后端会实际使用硬链接;Docker / NAS - 多挂载或不同子卷下可能报 invalid cross-device link。需要真正移动时请关闭“保种”,需要保种但硬链接失败时请选择“复制”。 -
- )} - -
- - - - - -
- - )} - - {autoTab === 'naming' && ( -
- - - -

- 可用占位符:{'{title}'} {'{year}'} {'{season}'} {'{season:02}'} {'{episode}'} {'{episode:02}'} {'{category}'}。扩展名会自动补齐。 -

-
- )} - - {autoTab === 'scrape' && ( - <> -
- - -
-
- - - - -
- - )} -
+ void autoOrganize.save()} + onRunNow={autoOrganize.runNow} + onTabChange={autoOrganize.setActiveTab} + onConfigChange={autoOrganize.changeConfig} + /> {data?.path && ( -
-
-
-

手动整理入库

-

来源优先使用选中项;未选中时使用当前目录。

-
-
- 来源:{organizeSource || '未选择'} -
-
- {selectedPaths.length > 0 && ( -
- 已选择 {selectedPaths.length} 个项目用于整理。 - -
- )} - -
- - - - -
- - {manualMoveKeepsSeeding && ( -
- “保种”已开启,选择“移动”时后端会改用硬链接以保留下载源。要执行真正移动,请先在上方自动整理设置里关闭“保种”并保存。 -
- )} - -
- - - - -
- - {previewItems.length > 0 && ( -
- - - - - - - - - - - {previewItems.map((item, index) => ( - - - - - - - ))} - -
动作来源目标原因
{item.action}{item.source}{item.target || '—'}{item.reason || '—'}
-
- )} -
+ fileOperations.setSelectedPaths([])} + onLibraryChange={setOrganizeLibraryID} + onDestPathChange={setOrganizeDestPath} + onMediaTypeChange={setOrganizeMediaType} + onTransferModeChange={setOrganizeTransferMode} + onScanAfterChange={setScanAfter} + onScrapeAfterChange={setScrapeAfter} + onPreview={() => runManualOrganize(true)} + onRun={() => runManualOrganize(false)} + /> )} {data?.path && ( -
- - 文件操作 - - 新建目录 / 重命名 / 删除 / 转移 - - -
-
-

新建目录

-
- setFolderName(e.target.value)} /> - -
-
- -
-

选中项

- {selected ? ( -
-

{selected.path}

-
- setRenameTo(e.target.value)} /> - - -
- {!selected.is_dir && ( -
- setDestPath(e.target.value)} /> - - -
- )} -
- ) : ( -

先在下方列表点击“操作”选择文件或目录。

- )} -
-
-
+ )} {loading &&

加载中…

} {error &&
{error}
} {!loading && data && !data.path && data.roots && ( -
- {data.roots.map((r) => ( - - ))} -
+ )} {!loading && data?.entries && data.entries.length > 0 && ( -
- - - - - - - - - - - - {data.entries.map((entry) => ( - - - - - - - - ))} - -
- 0 && data.entries.every((entry) => selectedPaths.includes(entry.path))} - onChange={(event) => { - const entries = data.entries ?? [] - if (event.target.checked) { - setSelectedPaths(entries.map((entry) => entry.path)) - } else { - setSelectedPaths([]) - } - }} - /> - 名称大小修改时间选择
- toggleSelectedPath(entry, event.target.checked)} - /> - - - {entry.is_dir ? '—' : fmtBytes(entry.size)}{new Date(entry.modified * 1000).toLocaleString()} - -
-
+ fileOperations.setSelectedPaths(checked ? data.entries?.map((entry) => entry.path) ?? [] : [])} + onToggleSelectedPath={fileOperations.toggleSelectedPath} + onEnter={enter} + onChoose={fileOperations.choose} + /> )} {!loading && data?.entries && data.entries.length === 0 &&

空目录。

} diff --git a/web/src/pages/FileManagerToolbar.tsx b/web/src/pages/FileManagerToolbar.tsx new file mode 100644 index 0000000..1d287c8 --- /dev/null +++ b/web/src/pages/FileManagerToolbar.tsx @@ -0,0 +1,46 @@ +import { ChevronUp, Home, RefreshCw } from 'lucide-react' + +type FileManagerToolbarProps = { + currentPath?: string + parentPath?: string + recursive: boolean + onRoot: () => void + onParent: (path: string) => void + onRefresh: () => void + onRecursiveChange: (value: boolean) => void +} + +export function FileManagerToolbar({ + currentPath, + parentPath, + recursive, + onRoot, + onParent, + onRefresh, + onRecursiveChange, +}: FileManagerToolbarProps) { + return ( +
+ + {parentPath && ( + + )} + + + {currentPath && ( + + {currentPath} + + )} +
+ ) +} diff --git a/web/src/pages/FileOperationsPanel.tsx b/web/src/pages/FileOperationsPanel.tsx new file mode 100644 index 0000000..438a3f4 --- /dev/null +++ b/web/src/pages/FileOperationsPanel.tsx @@ -0,0 +1,122 @@ +import { Copy, GitBranch, Move, Pencil, Plus, Trash2 } from 'lucide-react' + +import type { FileEntry } from '../api/files' + +type FileOperationsPanelProps = { + selected: FileEntry | null + folderName: string + renameTo: string + destPath: string + transferMode: string + busy: string + onFolderNameChange: (value: string) => void + onRenameToChange: (value: string) => void + onDestPathChange: (value: string) => void + onTransferModeChange: (value: string) => void + onCreateFolder: () => void + onRenameSelected: () => void + onDeleteSelected: () => void + onTransferSelected: () => void +} + +export function FileOperationsPanel({ + selected, + folderName, + renameTo, + destPath, + transferMode, + busy, + onFolderNameChange, + onRenameToChange, + onDestPathChange, + onTransferModeChange, + onCreateFolder, + onRenameSelected, + onDeleteSelected, + onTransferSelected, +}: FileOperationsPanelProps) { + return ( +
+ + 文件操作 + + 新建目录 / 重命名 / 删除 / 转移 + + +
+
+

新建目录

+
+ onFolderNameChange(event.target.value)} + /> + +
+
+ +
+

选中项

+ {selected ? ( +
+

{selected.path}

+
+ onRenameToChange(event.target.value)} + /> + + +
+ {!selected.is_dir && ( +
+ onDestPathChange(event.target.value)} + /> + + +
+ )} +
+ ) : ( +

先在下方列表点击“操作”选择文件或目录。

+ )} +
+
+
+ ) +} + +function transferIcon(mode: string) { + if (mode === 'move') return + if (mode === 'copy') return + return +} diff --git a/web/src/pages/HomePage.tsx b/web/src/pages/HomePage.tsx index d9ce52f..45b3804 100644 --- a/web/src/pages/HomePage.tsx +++ b/web/src/pages/HomePage.tsx @@ -60,10 +60,10 @@ export function HomePage() {
-
+
- 首页内容准备中… + 首页内容准备中…
) @@ -72,11 +72,11 @@ export function HomePage() { if (empty) { return (
-
- +
+
-

您的家庭影视站暂无内容

-

+

您的家庭影视站暂无内容

+

前往管理后台添加媒体目录,扫描后首页将展示本周力荐、继续观看和最近入库。

@@ -89,65 +89,65 @@ export function HomePage() { return (
{featuredItem && ( -
+
-
+
{featuredVisual && ( { event.currentTarget.style.display = 'none' }} /> )} -
-
+
+
-
+
本周力荐 / Featured
-
+
{featuredMark}
-

+

{featuredItem.title}

-

+

{featuredItem.overview || '家庭私人媒体中心收藏。支持多端播放、外部播放器、智能刮削与订阅下载。'}

-
+
{featuredItem.year > 0 && ( - {featuredItem.year} 年 + {featuredItem.year} 年 )} {featuredItem.video_codec && ( - + {featuredItem.video_codec} )} {featuredItem.container && ( - + {featuredItem.container} )}
- + 立即播放 - + 发现更多精彩 @@ -156,14 +156,14 @@ export function HomePage() {
-
-
+
+
- {featuredItem.title} + {featuredItem.title}
{featuredPoster && ( {featuredItem.title} 0 && (
- + -

继续观看

- {history.length} 个记录 +

继续观看

+ {history.length} 个记录
@@ -198,17 +198,17 @@ export function HomePage() { {recentCards.length > 0 && (
-
+
- +
-

最近入库

-

按整部电影、剧集、番剧和综艺合集展示新增内容。

+

最近入库

+

按整部电影、剧集、番剧和综艺合集展示新增内容。

- + 海报墙 @@ -232,18 +232,18 @@ export function HomePage() { function ContinueCard({ media, progress }: { media: Media; progress: number }) { return ( - -
+ +
{media.poster_url ? ( ) : ( -
+
)} @@ -252,11 +252,11 @@ function ContinueCard({ media, progress }: { media: Media; progress: number }) {
-

+

{media.title}

-
+
-

+

已观看到 {Math.round(progress * 100)}%

diff --git a/web/src/pages/LibrariesPage.tsx b/web/src/pages/LibrariesPage.tsx index ffa6d18..d62e6ef 100644 --- a/web/src/pages/LibrariesPage.tsx +++ b/web/src/pages/LibrariesPage.tsx @@ -6,6 +6,7 @@ import { ArrowRight, Film, FolderOpen, Library as LibraryIcon, Music, PlayCircle import { imageURL } from '../api/client' import { libraryAPI } from '../api/library' import { toolsAPI } from '../api/tools' +import { EpisodeArtworkToggle } from '../components/EpisodeArtworkToggle' import { MediaCard } from '../components/MediaCard' import type { Library, Media } from '../types' import { artworkScore, groupSeries, seriesCardLink, type SeriesCard } from '../utils/groupSeries' @@ -37,6 +38,7 @@ export function LibrariesPage() { const [previews, setPreviews] = useState([]) const [loading, setLoading] = useState(true) const [repairing, setRepairing] = useState(false) + const [repairEpisodeArtwork, setRepairEpisodeArtwork] = useState(false) const [repairMsg, setRepairMsg] = useState('') async function handleRepairRescrape() { @@ -44,7 +46,7 @@ export function LibrariesPage() { setRepairing(true) setRepairMsg('') try { - await toolsAPI.repairAndRescrapeAll() + await toolsAPI.repairAndRescrapeAll({ episode_images: repairEpisodeArtwork, refresh_matched: true }) setRepairMsg('已开始全库修复+重刮,进度可在任务中查看。') } catch { setRepairMsg('启动失败,请稍后重试。') @@ -61,7 +63,7 @@ export function LibrariesPage() { const libs = await libraryAPI.list() const rows = await Promise.all(libs.map(async (library) => { try { - const page = await libraryAPI.listMedia(library.id, 1, 160) + const page = await libraryAPI.listMedia(library.id, 1, 160, { groupVersions: false }) const cards = latestCards(page.items) return { library, items: page.items, total: page.total, cards } satisfies LibraryPreview } catch { @@ -94,6 +96,12 @@ export function LibrariesPage() {
{repairMsg && {repairMsg}} + + + + + + + ) +} diff --git a/web/src/pages/LibraryPage.tsx b/web/src/pages/LibraryPage.tsx index ecf7270..8c71649 100644 --- a/web/src/pages/LibraryPage.tsx +++ b/web/src/pages/LibraryPage.tsx @@ -1,24 +1,20 @@ -import { useCallback, useEffect, useMemo, useState } from 'react' -import { Link, useLocation, useParams, useSearchParams } from 'react-router-dom' +import { useState } from 'react' +import { useLocation, useParams, useSearchParams } from 'react-router-dom' import { motion, AnimatePresence } from 'framer-motion' -import toast from 'react-hot-toast' -import { ArrowLeft, Play, Film, Database, FileText, Search, Sparkles, Trash2, Pencil, FolderInput } from 'lucide-react' -import { libraryAPI } from '../api/library' -import { toolsAPI } from '../api/tools' -import { storageAPI, type CloudScanStatus } from '../api/storage_config' -import { api } from '../api/client' -import { recycleAPI } from '../api/recycle' -import type { Library, Media } from '../types' -import { MediaCard } from '../components/MediaCard' -import { ExternalPlayerButton } from '../components/ExternalPlayerButton' -import { imageURL } from '../api/client' +import type { Media } from '../types' import { useAuthStore } from '../stores/auth' -import { getSeriesKey, groupSeries, isEpisodeLike, seriesTitle, type SeriesCard } from '../utils/groupSeries' -import { useWebSocket } from '../hooks/useWebSocket' -import { confirmAction } from '../components/ConfirmDialog' +import { seriesTitle, type SeriesCard } from '../utils/groupSeries' import { ManualScrapeDialog } from '../components/ManualScrapeDialog' import { MetadataEditDialog } from '../components/MetadataEditDialog' +import { LibrarySeriesEpisodes } from './LibrarySeriesEpisodes' +import { LibrarySeriesDetailHeader } from './LibrarySeriesDetailHeader' +import { LibraryPageHeader } from './LibraryPageHeader' +import { LibraryMediaSections } from './LibraryMediaSections' +import { useLibraryData } from './useLibraryData' +import { useLibraryScanStatus } from './useLibraryScanStatus' +import { useLibrarySeriesSelection } from './useLibrarySeriesSelection' +import { useLibraryAdminActions } from './useLibraryAdminActions' export function LibraryPage() { const { id = '' } = useParams() @@ -26,422 +22,85 @@ export function LibraryPage() { const location = useLocation() const role = useAuthStore((s) => s.user?.role) - const [library, setLibrary] = useState(null) - const [items, setItems] = useState([]) - const [total, setTotal] = useState(0) - const [loading, setLoading] = useState(true) - const [loadingAll, setLoadingAll] = useState(false) - const [scanning, setScanning] = useState(false) - const [scanProgress, setScanProgress] = useState('') - const [scraping, setScraping] = useState(false) - const [repairing, setRepairing] = useState(false) - const [seriesToolBusy, setSeriesToolBusy] = useState('') const [manualSeriesScrapeOpen, setManualSeriesScrapeOpen] = useState(false) const [seriesMetadataEditOpen, setSeriesMetadataEditOpen] = useState(false) const [manualMovie, setManualMovie] = useState(null) - const [movieToolBusy, setMovieToolBusy] = useState('') // 剧集模式:选中某个剧集后展开详情 const [selectedSeries, setSelectedSeries] = useState(null) const [selectedSeason, setSelectedSeason] = useState(null) - const hasEpisodicItems = useMemo(() => items.some(isEpisodeLike), [items]) - const isSeries = library?.type === 'tv' || library?.type === 'anime' || library?.type === 'variety' || hasEpisodicItems + const { + library, + items, + seriesEpisodeItems, + total, + loading, + loadingSeriesEpisodes, + isSeriesLibrary, + isSeries, + seriesCards, + loadingAllText, + reloadCurrentLibrary, + } = useLibraryData(id, selectedSeries) - // 折叠后的剧集卡片 - const seriesCards = useMemo(() => { - if (!isSeries || items.length === 0) return [] - return groupSeries(items) - }, [isSeries, items]) + const { + scanning, + scanProgress, + handleScan, + } = useLibraryScanStatus({ + libraryID: id, + isAdmin: role === 'admin', + onLibraryChanged: reloadCurrentLibrary, + }) - // 选中的剧集:所有集按季分组 - const selectedEpisodes = useMemo(() => { - if (!selectedSeries || items.length === 0) return [] - const eps = items.filter((m) => getSeriesKey(m) === selectedSeries.key) - // 按季分组,按集排序 - const seasons = new Map() - for (const ep of eps) { - const s = ep.episode_num > 0 ? (ep.season_num ?? 0) : (ep.season_num || 1) - if (!seasons.has(s)) seasons.set(s, []) - seasons.get(s)!.push(ep) - } - for (const [, list] of seasons) { - list.sort((a, b) => (a.episode_num || 0) - (b.episode_num || 0)) - } - return Array.from(seasons.entries()) - .sort(([a], [b]) => a - b) - .map(([season, episodes]) => ({ season, episodes })) - }, [selectedSeries, items]) + const { + selectedEpisodes, + visibleEpisodes, + selectedSeriesEpisodes, + selectedSeriesMediaIDs, + handleSeriesClick, + clearSelectedSeries, + } = useLibrarySeriesSelection({ + items, + seriesEpisodeItems, + isSeriesLibrary, + isSeries, + loading, + seriesCards, + searchParams, + setSearchParams, + selectedSeries, + setSelectedSeries, + selectedSeason, + setSelectedSeason, + onClearSeriesState: () => setSeriesMetadataEditOpen(false), + }) - const visibleEpisodes = useMemo(() => { - if (selectedSeason == null) return selectedEpisodes[0]?.episodes ?? [] - return selectedEpisodes.find((s) => s.season === selectedSeason)?.episodes ?? [] - }, [selectedEpisodes, selectedSeason]) - - const selectedSeriesEpisodes = useMemo( - () => selectedEpisodes.flatMap((season) => season.episodes), - [selectedEpisodes], - ) - - const selectedSeriesMediaIDs = useMemo( - () => selectedSeriesEpisodes.map((ep) => ep.id), - [selectedSeriesEpisodes], - ) - - useEffect(() => { - if (!id) return - libraryAPI.list().then((all) => { - const lib = all.find((l) => l.id === id) ?? null - setLibrary(lib) - }) - }, [id]) - - useEffect(() => { - if (!id || !library) return - let cancelled = false - setLoading(true) - setLoadingAll(true) - setItems([]) - const loadAll = async () => { - const pageSize = 2000 - let page = 1 - let collected: Media[] = [] - try { - for (;;) { - const d = await libraryAPI.listMedia(id, page, pageSize) - if (cancelled) return - collected = collected.concat(d.items) - setItems(collected) - setTotal(d.total) - if (page === 1) setLoading(false) - if (collected.length >= d.total || d.items.length < pageSize) break - page += 1 - } - } finally { - if (!cancelled) { - setLoading(false) - setLoadingAll(false) - } - } - } - loadAll().catch(() => { - if (!cancelled) { - toast.error('媒体库加载失败') - setLoading(false) - } - }) - return () => { cancelled = true } - }, [id, library]) - - const reloadCurrentLibrary = useCallback(() => { - setLibrary((l) => (l ? { ...l } : l)) - }, []) - - const onRealtimeEvent = useCallback((topic: string, payload: unknown) => { - if (role !== 'admin') return - if (topic !== 'scan' || !payload || typeof payload !== 'object') return - const p = payload as Record - if (p.library_id !== id) return - if (p.error) { - setScanning(false) - setScanProgress(`扫描失败:${String(p.error)}`) - return - } - if (p.finished) { - setScanning(false) - const elapsed = Number(p.elapsed_seconds ?? p.elapsed ?? 0) - const elapsedText = elapsed > 0 ? ` · 耗时 ${formatDuration(elapsed)}` : '' - setScanProgress(`扫描完成:发现 ${p.discovered ?? p.visited ?? 0} · 新增 ${p.added ?? 0} · 更新 ${p.updated ?? 0} · 跳过 ${p.skipped ?? 0}${elapsedText}`) - reloadCurrentLibrary() - return - } - if (p.queued) { - setScanning(true) - setScanProgress(String(p.message ?? '扫描已排队,后台会自动入库')) - return - } - if (p.cloud && p.stage) { - const stage = p.stage === 'importing' ? '正在入库' : '正在遍历目录' - const speed = Number(p.files_per_second ?? 0) - const speedText = speed > 0 ? ` · ${speed.toFixed(speed >= 10 ? 0 : 1)} 个/秒` : '' - setScanning(true) - setScanProgress(`${stage}:目录 ${p.dirs ?? 0} · 已发现 ${p.discovered ?? 0} · 已入库 ${p.visited ?? 0}${speedText}`) - } - }, [id, reloadCurrentLibrary, role]) - - useWebSocket(onRealtimeEvent) - - useEffect(() => { - if (role !== 'admin' || !id) return - let cancelled = false - let terminal = false - let timer: number | undefined - const restoreCloudScanStatus = async () => { - const r = await storageAPI.cloudScanStatus() - if (cancelled) return - const status = (r.items ?? []).find((item) => item.library_id === id) - if (!status) return - if (status.state === 'running' || status.state === 'queued' || status.state === 'canceling') { - setScanning(true) - setScanProgress(formatCloudScanStatus(status)) - return - } - if (status.state === 'finished') { - setScanning(false) - setScanProgress(formatCloudScanStatus(status)) - reloadCurrentLibrary() - terminal = true - if (timer) window.clearInterval(timer) - return - } - if (status.state === 'error' && status.error) { - setScanning(false) - setScanProgress(`扫描失败:${status.error}`) - terminal = true - if (timer) window.clearInterval(timer) - } - } - restoreCloudScanStatus() - .catch(() => undefined) - .finally(() => { - if (!cancelled && !terminal) { - timer = window.setInterval(() => { - restoreCloudScanStatus().catch(() => undefined) - }, 5000) - } - }) - return () => { - cancelled = true - if (timer) window.clearInterval(timer) - } - }, [id, reloadCurrentLibrary, role]) - - useEffect(() => { - if (loading) return - if (!isSeries) { - setSelectedSeries(null) - setSelectedSeason(null) - return - } - - const key = searchParams.get('series') - if (!key) { - setSelectedSeries(null) - return - } - - const next = seriesCards.find((card) => card.key === key) - if (next) { - setSelectedSeries(next) - } else { - setSelectedSeries(null) - } - }, [isSeries, loading, searchParams, seriesCards]) - - useEffect(() => { - if (!selectedSeries || selectedEpisodes.length === 0) { - setSelectedSeason(null) - return - } - if (selectedSeason == null || !selectedEpisodes.some((s) => s.season === selectedSeason)) { - setSelectedSeason(selectedEpisodes[0].season) - } - }, [selectedSeries, selectedEpisodes, selectedSeason]) - - const handleScan = async () => { - setScanning(true) - setScanProgress('正在提交扫描任务…') - let keepScanning = false - try { - const r = await libraryAPI.scan(id) - if (r.queued) { - keepScanning = true - setScanProgress(`${r.message ?? '云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库'};${r.estimate_message ?? '大目录耗时取决于网盘接口速度'}`) - toast.success('云盘扫描已加入后台队列') - } else { - toast.success(`扫描完成:新增 ${r.added} 项,更新 ${r.updated ?? 0} 项`) - setScanProgress(`扫描完成:新增 ${r.added} · 更新 ${r.updated ?? 0}`) - reloadCurrentLibrary() - setScanning(false) - } - } catch { - toast.error('扫描失败') - setScanProgress('扫描失败,请查看日志或稍后重试') - setScanning(false) - } finally { - if (!keepScanning) setScanning(false) - } - } - - const handleScrape = async () => { - setScraping(true) - try { - await libraryAPI.scrape(id) - toast.success('刮削已加入后台队列') - } catch { toast.error('刮削失败') } - finally { setScraping(false) } - } - - const handleRepairRescrape = async () => { - if (repairing) return - setRepairing(true) - try { - await toolsAPI.repairAndRescrapeLibrary(id) - toast.success('本库修复+重刮已加入后台队列,进度可在任务中查看') - } catch { toast.error('修复+重刮启动失败') } - finally { setRepairing(false) } - } - - const runSeriesTool = async (key: string, label: string, action: (media: Media) => Promise) => { - if (selectedSeriesEpisodes.length === 0) return - setSeriesToolBusy(key) - try { - for (const ep of selectedSeriesEpisodes) { - await action(ep) - } - toast.success(`${label}完成:${selectedSeriesEpisodes.length} 个媒体`) - reloadCurrentLibrary() - } catch (err: unknown) { - const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error || `${label}失败` - toast.error(msg) - } finally { - setSeriesToolBusy('') - } - } - - const handleSeriesSmartScrape = () => { - runSeriesTool('scrape', '整剧智能刮削', (media) => api.post(`/media/${media.id}/scrape`)) - } - - const handleSeriesProbe = () => { - runSeriesTool('probe', '整剧媒体轨探测', (media) => api.post(`/media/${media.id}/probe`)) - } - - const handleSeriesNFO = () => { - runSeriesTool('nfo', '整剧 NFO 写出', (media) => recycleAPI.exportNFO(media.id)) - } - - const handleSeriesOrganize = async () => { - if (!selectedSeries || selectedSeriesEpisodes.length === 0 || !library) return - const source = seriesSourceRoot(selectedSeriesEpisodes) - if (!source || source.toLowerCase().startsWith('cloud://')) { - toast.error('当前合集不是本地文件夹,无法使用本地整理入库') - return - } - if (!(await confirmAction({ - title: '整理当前合集', - message: `来源:${source}\n目标:${library.path}\n将按当前媒体库类型整理整个文件夹。`, - confirmText: '开始整理', - }))) return - setSeriesToolBusy('organize') - try { - const result = await toolsAPI.organizeDirectory({ - source_path: source, - dest_path: library.path, - media_type: library.type || 'auto', - library_id: library.id, - scan_after: true, - scrape_after: true, - }) - const replaced = result.replaced ?? 0 - toast.success(`合集整理完成:新增 ${result.organized ?? 0} · 替换 ${replaced} · 跳过 ${result.skipped ?? 0}`) - reloadCurrentLibrary() - } catch (err: unknown) { - const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error || '合集整理失败' - toast.error(msg) - } finally { - setSeriesToolBusy('') - } - } - - const handleSeriesSoftDelete = async () => { - if (!selectedSeries || selectedSeriesEpisodes.length === 0) return - if (!(await confirmAction({ - title: '移入回收站', - message: `将「${seriesTitle(selectedSeries.rep)}」的 ${selectedSeriesEpisodes.length} 个媒体移至回收站? (磁盘文件保留)`, - confirmText: '移入回收站', - }))) return - await runSeriesTool('delete', '整剧移入回收站', (media) => recycleAPI.softDelete(media.id)) - clearSelectedSeries() - } - - const runMovieTool = async (media: Media, key: string, label: string, action: (media: Media) => Promise) => { - const busyKey = `${key}:${media.id}` - setMovieToolBusy(busyKey) - try { - await action(media) - toast.success(`${label}完成:${media.title}`) - reloadCurrentLibrary() - } catch (err: unknown) { - const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error || `${label}失败` - toast.error(msg) - } finally { - setMovieToolBusy('') - } - } - - const handleMovieSmartScrape = (media: Media) => { - runMovieTool(media, 'scrape', '智能刮削', (item) => api.post(`/media/${item.id}/scrape`)) - } - - const handleMovieProbe = (media: Media) => { - runMovieTool(media, 'probe', '媒体轨探测', (item) => api.post(`/media/${item.id}/probe`)) - } - - const handleMovieNFO = (media: Media) => { - runMovieTool(media, 'nfo', 'NFO 写出', (item) => recycleAPI.exportNFO(item.id)) - } - - const handleMovieSoftDelete = async (media: Media) => { - if (!(await confirmAction({ - title: '移入回收站', - message: `将「${media.title}」移至回收站? (磁盘文件保留)`, - confirmText: '移入回收站', - }))) return - await runMovieTool(media, 'delete', '移入回收站', (item) => recycleAPI.softDelete(item.id)) - } - - const movieActions = (media: Media) => { - if (role !== 'admin') return undefined - const busy = movieToolBusy.endsWith(`:${media.id}`) - const buttonClass = 'flex h-8 w-8 items-center justify-center rounded-lg border border-white/70 bg-white/90 text-gray-700 shadow-sm backdrop-blur transition hover:bg-brand-50 hover:text-brand-600 disabled:opacity-50' - return ( - <> - - - - - - - ) - } - - const handleSeriesClick = (card: SeriesCard) => { - setSelectedSeries(card) - const next = new URLSearchParams(searchParams) - next.set('series', card.key) - setSearchParams(next) - window.scrollTo({ top: 0, behavior: 'smooth' }) - } - - const clearSelectedSeries = () => { - setSelectedSeries(null) - setSelectedSeason(null) - setSeriesMetadataEditOpen(false) - const next = new URLSearchParams(searchParams) - next.delete('series') - setSearchParams(next) - } + const { + scraping, + scrapeEpisodeArtwork, + repairing, + seriesToolBusy, + setScrapeEpisodeArtwork, + handleScrape, + handleRepairRescrape, + handleSeriesSmartScrape, + handleSeriesProbe, + handleSeriesNFO, + handleSeriesOrganize, + handleSeriesSoftDelete, + movieActions, + } = useLibraryAdminActions({ + libraryID: id, + role, + library, + selectedSeries, + selectedSeriesEpisodes, + reloadCurrentLibrary, + clearSelectedSeries, + setManualMovie, + }) if (loading) { return ( @@ -456,52 +115,31 @@ export function LibraryPage() { return (
- {/* Header */} -
-
-

- {library?.name ?? '媒体库'} - ({isSeries ? seriesCards.length : total}) -

- {library &&

{library.type} · {library.path}

} - {loadingAll && !loading && total > items.length && ( -

正在继续加载全部条目:{items.length} / {total}

- )} - {scanProgress &&

{scanProgress}

} -
- {role === 'admin' && ( -
- - - -
- )} -
+ - {/* 非剧集:直接展示海报网格 */} - {!isSeries && items.length > 0 && ( -
- {items.map((m) => ( - - ))} -
- )} - - {!isSeries && items.length === 0 && ( -
- -

该媒体库暂无内容,触发一次扫描后再来看看

-
- )} - - {/* 剧集模式:折叠卡片网格 */} - {isSeries && seriesCards.length > 0 && !selectedSeries && ( -
- {seriesCards.map((s) => ( - handleSeriesClick(s)} /> - ))} -
- )} + {/* 剧集详情:季/集选择器 */} @@ -512,173 +150,46 @@ export function LibraryPage() { exit={{ opacity: 0 }} className="space-y-6" > - {/* 返回按钮 + 标题 */} -
- -

- {seriesTitle(selectedSeries.rep)} -

- 共 {selectedSeries.count} 集 -
- - {/* 海报 + 从第一集开始播放 */} -
-
- {selectedSeries.rep.poster_url ? ( - {selectedSeries.rep.title} - ) : ( -
- )} -
-
-

- {selectedSeries.rep.overview || '暂无简介'} -

- - {/* 从第一集开始 */} - {(() => { - const firstEps = [...(visibleEpisodes.length > 0 ? visibleEpisodes : selectedEpisodes.flatMap((s) => s.episodes))] - firstEps.sort((a, b) => - (a.season_num || 0) - (b.season_num || 0) - || (a.episode_num || 0) - (b.episode_num || 0), - ) - const first = firstEps.length > 0 ? firstEps[0] : null - return first ? ( -
- - - 从第一集开始播放 - - -
- ) : null - })()} - - {role === 'admin' && selectedSeriesEpisodes.length > 0 && ( -
-

系统后台高级控制面板

-
- - - - - - - -
-
- )} -
-
+ setManualSeriesScrapeOpen(true)} + onMetadataEdit={() => setSeriesMetadataEditOpen(true)} + onProbe={handleSeriesProbe} + onNFO={handleSeriesNFO} + onOrganize={handleSeriesOrganize} + onSoftDelete={handleSeriesSoftDelete} + /> {/* 季 / 集列表 */}
-
- {selectedEpisodes.map(({ season, episodes }) => ( - - ))} -
- -
-

- {(selectedSeason ?? selectedEpisodes[0]?.season ?? 1) === 0 - ? '特别篇' - : `第 ${selectedSeason ?? selectedEpisodes[0]?.season ?? 1} 季`} -

-
- {visibleEpisodes.map((ep) => ( -
- -
- {ep.backdrop_url || ep.poster_url ? ( - - ) : ( - ep.episode_num || '—' - )} -
-
-

- {ep.original_name || (ep.episode_num > 0 ? `第 ${ep.episode_num} 集` : ep.title)} -

-

- {ep.duration_sec > 0 - ? `${Math.floor(ep.duration_sec / 60)} 分钟` - : formatSize(ep.size_bytes)} -

-
- - - -
- ))} -
-
+
)}
- {/* 剧集为空 */} - {isSeries && seriesCards.length === 0 && !loading && ( -
- -

该库尚未发现任何剧集,触发一次扫描后再来看看

-
- )} - setManualSeriesScrapeOpen(false)} onApplied={reloadCurrentLibrary} /> @@ -695,8 +206,9 @@ export function LibraryPage() { open={!!manualMovie} media={manualMovie} defaultQuery={manualMovie?.title ?? ''} - mediaType={library?.type || 'movie'} + mediaType={manualMovie ? scrapeMediaType(library?.type, manualMovie) : library?.type || 'movie'} scopeLabel={manualMovie?.title ?? '当前电影'} + episodeArtwork={scrapeEpisodeArtwork} onClose={() => setManualMovie(null)} onApplied={reloadCurrentLibrary} /> @@ -704,55 +216,9 @@ export function LibraryPage() { ) } -function formatDuration(seconds: number): string { - if (!Number.isFinite(seconds) || seconds <= 0) return '' - if (seconds < 60) return `${Math.round(seconds)}秒` - const minutes = Math.floor(seconds / 60) - const rest = Math.round(seconds % 60) - if (minutes < 60) return `${minutes}分${rest}秒` - const hours = Math.floor(minutes / 60) - return `${hours}小时${minutes % 60}分` -} - -function formatCloudScanStatus(status: CloudScanStatus): string { - const stage = - status.state === 'queued' ? '扫描已排队' - : status.state === 'canceling' ? '正在中断扫描' - : status.state === 'finished' ? '扫描完成' - : status.stage === 'importing' ? '正在入库' - : '正在遍历目录' - const speed = Number(status.files_per_second ?? 0) - const speedText = speed > 0 && status.state !== 'finished' - ? ` · ${speed.toFixed(speed >= 10 ? 0 : 1)} 个/秒` - : '' - return `${stage}:目录 ${status.dirs ?? 0} · 已发现 ${status.discovered ?? 0} · 已入库 ${status.visited ?? 0} · 新增 ${status.added ?? 0} · 更新 ${status.updated ?? 0}${speedText}` -} - -function seriesSourceRoot(episodes: Media[]): string { - const firstPath = episodes.find((item) => item.path)?.path ?? '' - if (!firstPath) return '' - const dir = dirname(firstPath) - const base = basename(dir) - if (/^(?:s\d{1,2}|season[\s._-]*\d{1,2}|第\s*\d{1,2}\s*季|specials?|sp|ova|oad|特别篇|特別篇)$/i.test(base)) { - return dirname(dir) +function scrapeMediaType(libraryType: string | undefined, media: Media): string { + if ((media.season_num ?? 0) > 0 || (media.episode_num ?? 0) > 0) { + return 'tv' } - return dir -} - -function dirname(value: string): string { - const index = Math.max(value.lastIndexOf('/'), value.lastIndexOf('\\')) - return index > 0 ? value.slice(0, index) : '' -} - -function basename(value: string): string { - const index = Math.max(value.lastIndexOf('/'), value.lastIndexOf('\\')) - return index >= 0 ? value.slice(index + 1) : value -} - -function formatSize(bytes: number): string { - if (!bytes || bytes <= 0) return '—' - const units = ['B', 'KB', 'MB', 'GB', 'TB'] - let v = bytes, i = 0 - while (v >= 1024 && i < units.length - 1) { v /= 1024; i++ } - return `${v.toFixed(1)} ${units[i]}` + return libraryType || 'movie' } diff --git a/web/src/pages/LibraryPageHeader.tsx b/web/src/pages/LibraryPageHeader.tsx new file mode 100644 index 0000000..765f200 --- /dev/null +++ b/web/src/pages/LibraryPageHeader.tsx @@ -0,0 +1,72 @@ +import type { Library } from '../types' +import { EpisodeArtworkToggle } from '../components/EpisodeArtworkToggle' + +type LibraryPageHeaderProps = { + library: Library | null + itemCount: number + loadingAllText: string + scanProgress: string + isAdmin: boolean + scrapeEpisodeArtwork: boolean + scanning: boolean + scraping: boolean + repairing: boolean + onScrapeEpisodeArtworkChange: (checked: boolean) => void + onScan: () => void + onScrape: () => void + onRepairRescrape: () => void +} + +export function LibraryPageHeader({ + library, + itemCount, + loadingAllText, + scanProgress, + isAdmin, + scrapeEpisodeArtwork, + scanning, + scraping, + repairing, + onScrapeEpisodeArtworkChange, + onScan, + onScrape, + onRepairRescrape, +}: LibraryPageHeaderProps) { + return ( +
+
+

+ {library?.name ?? '媒体库'} + ({itemCount}) +

+ {library &&

{library.type} · {library.path}

} + {loadingAllText &&

{loadingAllText}

} + {scanProgress &&

{scanProgress}

} +
+ {isAdmin && ( +
+ + + + +
+ )} +
+ ) +} diff --git a/web/src/pages/LibrarySeriesDetailHeader.tsx b/web/src/pages/LibrarySeriesDetailHeader.tsx new file mode 100644 index 0000000..a0f75a0 --- /dev/null +++ b/web/src/pages/LibrarySeriesDetailHeader.tsx @@ -0,0 +1,135 @@ +import { Link } from 'react-router-dom' +import { ArrowLeft, Database, FileText, Film, FolderInput, Pencil, Play, Search, Sparkles, Trash2 } from 'lucide-react' + +import { imageURL } from '../api/client' +import { ExternalPlayerButton } from '../components/ExternalPlayerButton' +import type { Media } from '../types' +import { seriesTitle, type SeriesCard } from '../utils/groupSeries' + +type LibrarySeriesDetailHeaderProps = { + series: SeriesCard + visibleEpisodes: Media[] + allEpisodes: Media[] + playbackFrom: string + isAdmin: boolean + seriesToolBusy: string + onBack: () => void + onSmartScrape: () => void + onManualScrape: () => void + onMetadataEdit: () => void + onProbe: () => void + onNFO: () => void + onOrganize: () => void + onSoftDelete: () => void +} + +export function LibrarySeriesDetailHeader({ + series, + visibleEpisodes, + allEpisodes, + playbackFrom, + isAdmin, + seriesToolBusy, + onBack, + onSmartScrape, + onManualScrape, + onMetadataEdit, + onProbe, + onNFO, + onOrganize, + onSoftDelete, +}: LibrarySeriesDetailHeaderProps) { + const firstEpisode = firstPlayableEpisode(visibleEpisodes.length > 0 ? visibleEpisodes : allEpisodes) + + return ( + <> +
+ +

+ {seriesTitle(series.rep)} +

+ 共 {series.count} 集 +
+ +
+
+ {series.rep.poster_url ? ( + {series.rep.title} + ) : ( +
+ +
+ )} +
+
+

+ {series.rep.overview || '暂无简介'} +

+ + {firstEpisode && ( +
+ + + 从第一集开始播放 + + +
+ )} + + {isAdmin && allEpisodes.length > 0 && ( +
+

系统后台高级控制面板

+
+ + + + + + + +
+
+ )} +
+
+ + ) +} + +function firstPlayableEpisode(episodes: Media[]): Media | null { + const sorted = [...episodes] + sorted.sort((a, b) => + (a.season_num || 0) - (b.season_num || 0) + || (a.episode_num || 0) - (b.episode_num || 0), + ) + return sorted[0] ?? null +} diff --git a/web/src/pages/LibrarySeriesEpisodes.tsx b/web/src/pages/LibrarySeriesEpisodes.tsx new file mode 100644 index 0000000..8a73957 --- /dev/null +++ b/web/src/pages/LibrarySeriesEpisodes.tsx @@ -0,0 +1,141 @@ +import { Link } from 'react-router-dom' +import { Play } from 'lucide-react' + +import { imageURL } from '../api/client' +import { ExternalPlayerButton } from '../components/ExternalPlayerButton' +import type { Media } from '../types' +import { seriesTitleFromPath } from '../utils/groupSeries' +import { formatSize } from './libraryPageModel' + +type SeasonGroup = { + season: number + episodes: Media[] +} + +type LibrarySeriesEpisodesProps = { + loading: boolean + selectedEpisodes: SeasonGroup[] + selectedSeason: number | null + visibleEpisodes: Media[] + playbackFrom: string + onSeasonChange: (season: number) => void +} + +export function LibrarySeriesEpisodes({ + loading, + selectedEpisodes, + selectedSeason, + visibleEpisodes, + playbackFrom, + onSeasonChange, +}: LibrarySeriesEpisodesProps) { + if (loading) { + return ( +
+ 正在加载剧集… +
+ ) + } + + const displaySeason = selectedSeason ?? selectedEpisodes[0]?.season ?? 1 + + return ( + <> +
+ {selectedEpisodes.map(({ season, episodes }) => ( + + ))} +
+ +
+

+ {displaySeason === 0 ? '特别篇' : `第 ${displaySeason} 季`} +

+
+ {visibleEpisodes.map((ep) => ( +
+ +
+ {ep.backdrop_url || ep.poster_url ? ( + + ) : ( + ep.episode_num || '—' + )} +
+
+

+ {episodeDisplayTitle(ep, visibleEpisodes)} +

+

+ {ep.duration_sec > 0 + ? `${Math.floor(ep.duration_sec / 60)} 分钟` + : formatSize(ep.size_bytes)} +

+
+ + + +
+ ))} +
+
+ + ) +} + +function episodeDisplayTitle(ep: Media, siblings: Media[]): string { + const title = ep.episode_title?.trim() + if (title && !looksLikeSeriesTitle(ep, title, siblings)) { + return title + } + + const mediaTitle = ep.title?.trim() + if (mediaTitle && !looksLikeSeriesTitle(ep, mediaTitle, siblings)) { + return mediaTitle + } + + return ep.episode_num > 0 ? `第 ${ep.episode_num} 集` : mediaTitle || title || '未命名' +} + +function looksLikeSeriesTitle(ep: Media, title: string, siblings: Media[]): boolean { + const normalized = normalizeEpisodeTitle(title) + if (!normalized) return true + if (ep.original_name && normalizeEpisodeTitle(ep.original_name) === normalized) return true + const pathTitle = seriesTitleFromPath(ep.path) + if (pathTitle && normalizeEpisodeTitle(pathTitle) === normalized) return true + + const siblingTitles = new Set( + siblings + .map((item) => normalizeEpisodeTitle(item.title)) + .filter(Boolean), + ) + return siblingTitles.size === 1 && siblingTitles.has(normalized) && siblings.length > 1 +} + +function normalizeEpisodeTitle(value?: string): string { + return (value ?? '') + .toLowerCase() + .replace(/\s*\((?:19|20)\d{2}\)\s*/g, ' ') + .replace(/\s*\{(?:tmdb|tmdbid|douban|bangumi|bgm|thetvdb|tvdb)[\s:=#-]*[a-z0-9_-]+\}\s*/g, ' ') + .replace(/[\s._-]+/g, ' ') + .trim() +} diff --git a/web/src/pages/ManualOrganizePanel.tsx b/web/src/pages/ManualOrganizePanel.tsx new file mode 100644 index 0000000..d1b6b4c --- /dev/null +++ b/web/src/pages/ManualOrganizePanel.tsx @@ -0,0 +1,197 @@ +import type { Library } from '../types' + +type PreviewItem = { + source: string + target?: string + action: string + reason?: string +} + +type ManualOrganizePanelProps = { + organizeSource: string + selectedCount: number + localLibraries: Library[] + organizeLibraryID: string + organizeDestPath: string + organizeMediaType: string + organizeTransferMode: string + manualMoveKeepsSeeding: boolean + scanAfter: boolean + scrapeAfter: boolean + organizeReady: boolean + organizeBusy: string + previewItems: PreviewItem[] + onClearSelected: () => void + onLibraryChange: (value: string) => void + onDestPathChange: (value: string) => void + onMediaTypeChange: (value: string) => void + onTransferModeChange: (value: string) => void + onScanAfterChange: (value: boolean) => void + onScrapeAfterChange: (value: boolean) => void + onPreview: () => void + onRun: () => void +} + +export function ManualOrganizePanel({ + organizeSource, + selectedCount, + localLibraries, + organizeLibraryID, + organizeDestPath, + organizeMediaType, + organizeTransferMode, + manualMoveKeepsSeeding, + scanAfter, + scrapeAfter, + organizeReady, + organizeBusy, + previewItems, + onClearSelected, + onLibraryChange, + onDestPathChange, + onMediaTypeChange, + onTransferModeChange, + onScanAfterChange, + onScrapeAfterChange, + onPreview, + onRun, +}: ManualOrganizePanelProps) { + return ( +
+
+
+

手动整理入库

+

来源优先使用选中项;未选中时使用当前目录。

+
+
+ 来源:{organizeSource || '未选择'} +
+
+ {selectedCount > 0 && ( +
+ 已选择 {selectedCount} 个项目用于整理。 + +
+ )} + +
+ + + + +
+ + {manualMoveKeepsSeeding && ( +
+ “保种”已开启,选择“移动”时后端会改用硬链接以保留下载源。要执行真正移动,请先在上方自动整理设置里关闭“保种”并保存。 +
+ )} + +
+ + + + +
+ + {previewItems.length > 0 && } +
+ ) +} + +function ManualOrganizePreviewTable({ items }: { items: PreviewItem[] }) { + return ( +
+ + + + + + + + + + + {items.map((item, index) => ( + + + + + + + ))} + +
动作来源目标原因
{item.action}{item.source}{item.target || '—'}{item.reason || '—'}
+
+ ) +} diff --git a/web/src/pages/MediaDetailAdminPanel.tsx b/web/src/pages/MediaDetailAdminPanel.tsx new file mode 100644 index 0000000..4b68f93 --- /dev/null +++ b/web/src/pages/MediaDetailAdminPanel.tsx @@ -0,0 +1,81 @@ +import { Database, FileText, FolderInput, Pencil, Search, Sparkles, Trash2 } from 'lucide-react' + +import { EpisodeArtworkToggle } from '../components/EpisodeArtworkToggle' +import type { Media } from '../types' + +type MediaDetailAdminPanelProps = { + media: Media + scrapeEpisodeArtwork: boolean + onScrapeEpisodeArtworkChange: (checked: boolean) => void + onSmartScrape: () => void + onManualScrape: () => void + onMetadataEdit: () => void + onOrganize: () => void + onProbe: () => void + onExportNFO: () => void + onSoftDelete: () => void +} + +export function MediaDetailAdminPanel({ + media, + scrapeEpisodeArtwork, + onScrapeEpisodeArtworkChange, + onSmartScrape, + onManualScrape, + onMetadataEdit, + onOrganize, + onProbe, + onExportNFO, + onSoftDelete, +}: MediaDetailAdminPanelProps) { + return ( +
+

系统后台高级控制面板

+ {isEpisodeArtworkTarget(media) && ( + + )} +
+ + + + + + + +
+
+ ) +} + +function isEpisodeArtworkTarget(media: Media): boolean { + return media.season_num > 0 || media.episode_num > 0 +} diff --git a/web/src/pages/MediaDetailArtwork.tsx b/web/src/pages/MediaDetailArtwork.tsx new file mode 100644 index 0000000..113fd75 --- /dev/null +++ b/web/src/pages/MediaDetailArtwork.tsx @@ -0,0 +1,62 @@ +import { motion } from 'framer-motion' +import { FileText, Play } from 'lucide-react' +import { Link } from 'react-router-dom' + +import { imageURL } from '../api/client' +import type { Media } from '../types' + +type MediaDetailArtworkProps = { + media: Media +} + +export function MediaDetailBackdrop({ media }: MediaDetailArtworkProps) { + return ( +
+ {media.backdrop_url || media.poster_url ? ( + + ) : ( +
+ )} +
+
+ ) +} + +export function MediaDetailPoster({ media }: MediaDetailArtworkProps) { + return ( +
+ + {media.poster_url ? ( + {media.title} + ) : ( +
+ + 无海报 +
+ )} + + +
+ +
+ +
+
+ ) +} diff --git a/web/src/pages/MediaDetailMetadata.tsx b/web/src/pages/MediaDetailMetadata.tsx new file mode 100644 index 0000000..b2cd965 --- /dev/null +++ b/web/src/pages/MediaDetailMetadata.tsx @@ -0,0 +1,109 @@ +import { Calendar } from 'lucide-react' + +import type { Media } from '../types' + +type MediaDetailMetadataProps = { + media: Media +} + +export function MediaDetailMetadata({ media }: MediaDetailMetadataProps) { + const heading = media.episode_title?.trim() || media.title + const showTitleContext = Boolean(media.episode_title?.trim() && media.title && media.title !== heading) + + return ( + <> +
+

+ {heading} +

+ {showTitleContext && ( +

+ {media.title} +

+ )} +
+ {media.year > 0 && ( + + + {media.year} 年 + + )} + {media.width > 0 && ( + + {media.width} × {media.height} + + )} + + {fmtSize(media.size_bytes)} + + + {fmtDuration(media.duration_sec)} + + {media.container && ( + + {media.container} + + )} +
+
+ + {media.overview && ( +
+

剧情简介

+

+ {media.overview} +

+
+ )} + +
+ + + +
+ + ) +} + +function MetadataTags({ label, values, primary = false }: { label: string; values: string[]; primary?: boolean }) { + if (values.length === 0) return null + const tagClass = primary + ? 'rounded-full bg-brand-50 text-brand-700 border border-brand-100/30 px-3 py-1 text-2xs font-bold uppercase tracking-wider' + : 'rounded-xl bg-gray-100 text-gray-600 border border-gray-200/40 px-2.5 py-1 text-2xs font-semibold' + return ( +
+ {label} +
+ {values.map((value) => ( + + {value} + + ))} +
+
+ ) +} + +function fmtDuration(sec: number): string { + if (!sec || sec <= 0) return '—' + const h = Math.floor(sec / 3600) + const m = Math.floor((sec % 3600) / 60) + return h > 0 ? `${h}h ${m}m` : `${m}m` +} + +function fmtSize(bytes: number): string { + if (!bytes || bytes <= 0) return '—' + const units = ['B', 'KB', 'MB', 'GB', 'TB'] + let v = bytes + let i = 0 + while (v >= 1024 && i < units.length - 1) { + v /= 1024 + i++ + } + return `${v.toFixed(2)} ${units[i]}` +} + +function parseCSV(s?: string): string[] { + if (!s) return [] + return s.split(',').map((x) => x.trim()).filter(Boolean) +} diff --git a/web/src/pages/MediaDetailPage.tsx b/web/src/pages/MediaDetailPage.tsx index 3781d0c..42739fd 100644 --- a/web/src/pages/MediaDetailPage.tsx +++ b/web/src/pages/MediaDetailPage.tsx @@ -1,44 +1,32 @@ -import { motion } from 'framer-motion' -import { useEffect, useState } from 'react' +import { useCallback, useEffect, useState } from 'react' import { Link, useNavigate, useParams } from 'react-router-dom' -import { FileText, Heart, Play, RefreshCw, Sparkles, Trash2, Calendar, Database, Search, Pencil, FolderInput } from 'lucide-react' +import { ArrowLeft, Heart, Play, RefreshCw } from 'lucide-react' import toast from 'react-hot-toast' import { mediaAPI } from '../api/library' import { playbackAPI } from '../api/playback' import { recycleAPI } from '../api/recycle' -import { imageURL } from '../api/client' import { useAuthStore } from '../stores/auth' import { api } from '../api/client' import type { Media } from '../types' -import { confirmAction } from '../components/ConfirmDialog' +import { confirmAction } from '../components/confirmAction' import { ExternalPlayerButton } from '../components/ExternalPlayerButton' import { ManualScrapeDialog } from '../components/ManualScrapeDialog' import { MetadataEditDialog } from '../components/MetadataEditDialog' import { OrganizeMediaDialog } from '../components/OrganizeMediaDialog' +import { getSeriesKey, isEpisodeLike } from '../utils/groupSeries' +import { MediaDetailAdminPanel } from './MediaDetailAdminPanel' +import { MediaDetailBackdrop, MediaDetailPoster } from './MediaDetailArtwork' +import { MediaDetailMetadata } from './MediaDetailMetadata' -function fmtDuration(sec: number): string { - if (!sec || sec <= 0) return '—' - const h = Math.floor(sec / 3600) - const m = Math.floor((sec % 3600) / 60) - return h > 0 ? `${h}h ${m}m` : `${m}m` -} +function mediaLibraryBackTarget(media: Media): string { + const libraryID = media.display_library_id || media.library_id + if (!libraryID) return '' + if (!isEpisodeLike(media)) return `/library/${encodeURIComponent(libraryID)}` -function fmtSize(bytes: number): string { - if (!bytes || bytes <= 0) return '—' - const units = ['B', 'KB', 'MB', 'GB', 'TB'] - let v = bytes - let i = 0 - while (v >= 1024 && i < units.length - 1) { - v /= 1024 - i++ - } - return `${v.toFixed(2)} ${units[i]}` -} - -function parseCSV(s?: string): string[] { - if (!s) return [] - return s.split(',').map(x => x.trim()).filter(Boolean) + const seriesKey = getSeriesKey(media) + const target = `/library/${encodeURIComponent(libraryID)}` + return seriesKey ? `${target}?series=${encodeURIComponent(seriesKey)}` : target } export function MediaDetailPage() { @@ -51,8 +39,9 @@ export function MediaDetailPage() { const [manualScrapeOpen, setManualScrapeOpen] = useState(false) const [metadataEditOpen, setMetadataEditOpen] = useState(false) const [organizeOpen, setOrganizeOpen] = useState(false) + const [scrapeEpisodeArtwork, setScrapeEpisodeArtwork] = useState(false) - const refresh = async () => { + const refresh = useCallback(async () => { if (!id) return setLoading(true) try { @@ -63,10 +52,11 @@ export function MediaDetailPage() { } finally { setLoading(false) } - } + }, [id]) + useEffect(() => { refresh().catch(() => undefined) - }, [id]) + }, [refresh]) const toggleFav = async () => { if (!media) return @@ -77,7 +67,11 @@ export function MediaDetailPage() { const rescrape = async () => { if (!media) return - await api.post(`/media/${media.id}/scrape`) + await api.post(`/media/${media.id}/scrape`, { + episode_images: scrapeEpisodeArtwork, + refresh_matched: true, + include_matched: true, + }) toast.success('已触发重新刮削') await refresh() } @@ -118,7 +112,9 @@ export function MediaDetailPage() { if (!(await confirmAction({ title: '移入回收站', message: `将「${media.title}」移至回收站? (磁盘文件保留)`, confirmText: '移入回收站' }))) return await recycleAPI.softDelete(media.id) toast.success('已移至回收站') - navigate(-1) + const backTarget = mediaLibraryBackTarget(media) + if (backTarget) navigate(backTarget, { replace: true }) + else navigate(-1) } if (loading) { @@ -138,155 +134,38 @@ export function MediaDetailPage() { return (
- {/* ── Cinematic Blurred Backdrop Glow ── */} -
- {(media.backdrop_url || media.poster_url) ? ( - - ) : ( -
- )} -
+ + +
+
- {/* ── Main Details Layout Container ── */}
- {/* Poster Card */} -
- - {media.poster_url ? ( - {media.title} - ) : ( -
- - 无海报 -
- )} - - {/* Quick Play overlay button */} - -
- -
- -
-
+ - {/* Detailed Metadata Body */}
- {/* Title and Year Header */} -
-

- {media.title} -

-
- {media.year > 0 && ( - - - {media.year} 年 - - )} - {media.width > 0 && ( - - {media.width} × {media.height} - - )} - - {fmtSize(media.size_bytes)} - - - {fmtDuration(media.duration_sec)} - - {media.container && ( - - {media.container} - - )} -
-
- - {/* Description Card */} - {media.overview && ( -
-

剧情简介

-

- {media.overview} -

-
- )} - - {/* Tag Rows (Genres, Languages, Countries) */} -
- {/* Genres */} - {parseCSV(media.genres).length > 0 && ( -
- 类型流派 -
- {parseCSV(media.genres).map((g) => ( - - {g} - - ))} -
-
- )} - - {/* Languages */} - {parseCSV(media.languages).length > 0 && ( -
- 语言 -
- {parseCSV(media.languages).map((l) => ( - - {l} - - ))} -
-
- )} - - {/* Countries */} - {parseCSV(media.countries).length > 0 && ( -
- 国家/地区 -
- {parseCSV(media.countries).map((c) => ( - - {c} - - ))} -
-
- )} -
+
- {/* Action Buttons Panel */}
- {/* Primary Direct Play */} 立即播放 - {/* Transcode Playback */} - {/* Toggle Favourites */}
- {/* Admin Management Toolbar */} {role === 'admin' && ( -
-

系统后台高级控制面板

-
- - - - - - - -
-
+ setManualScrapeOpen(true)} + onMetadataEdit={() => setMetadataEditOpen(true)} + onOrganize={() => setOrganizeOpen(true)} + onProbe={reprobe} + onExportNFO={exportNFO} + onSoftDelete={softDelete} + /> )}
@@ -360,6 +213,7 @@ export function MediaDetailPage() { defaultQuery={media.title} mediaType={media.season_num > 0 || media.episode_num > 0 ? 'tv' : undefined} scopeLabel={media.title} + episodeArtwork={scrapeEpisodeArtwork} onClose={() => setManualScrapeOpen(false)} onApplied={refresh} /> diff --git a/web/src/pages/NotifyChannelCard.tsx b/web/src/pages/NotifyChannelCard.tsx new file mode 100644 index 0000000..dff8c2e --- /dev/null +++ b/web/src/pages/NotifyChannelCard.tsx @@ -0,0 +1,60 @@ +import { Loader2, Pencil, Send, Trash2 } from 'lucide-react' + +import type { NotifyChannel } from '../types' +import { channelSummary, eventSummary, TYPE_LABELS } from './notifyChannelsModel' + +type NotifyChannelCardProps = { + channel: NotifyChannel + onTest: () => void + testing?: boolean + onEdit: () => void + onDelete: () => void +} + +export function NotifyChannelCard({ + channel, + onTest, + testing, + onEdit, + onDelete, +}: NotifyChannelCardProps) { + const summary = channelSummary(channel) + return ( +
+
+
+ {channel.name} + + {TYPE_LABELS[channel.type] ?? channel.type} + + {!channel.enabled && ( + 已禁用 + )} +
+
{summary}
+
{eventSummary(channel.events)}
+
+
+ + + +
+
+ ) +} diff --git a/web/src/pages/NotifyChannelEventFields.tsx b/web/src/pages/NotifyChannelEventFields.tsx new file mode 100644 index 0000000..435f8b4 --- /dev/null +++ b/web/src/pages/NotifyChannelEventFields.tsx @@ -0,0 +1,67 @@ +import { EVENT_OPTIONS, type EventMode } from './notifyChannelsModel' +import { Field } from './NotifyChannelFormField' + +type NotifyChannelEventFieldsProps = { + events: string[] + eventMode: EventMode + setEventMode: (mode: EventMode) => void + toggleEvent: (event: string) => void +} + +export function NotifyChannelEventFields({ + events, + eventMode, + setEventMode, + toggleEvent, +}: NotifyChannelEventFieldsProps) { + return ( + +
+ + + +
+ {EVENT_OPTIONS.map((event) => ( + + ))} +
+
+
+ ) +} diff --git a/web/src/pages/NotifyChannelFormField.tsx b/web/src/pages/NotifyChannelFormField.tsx new file mode 100644 index 0000000..59fc264 --- /dev/null +++ b/web/src/pages/NotifyChannelFormField.tsx @@ -0,0 +1,10 @@ +import type { ReactNode } from 'react' + +export function Field({ label, children }: { label: string; children: ReactNode }) { + return ( + + ) +} diff --git a/web/src/pages/NotifyChannelFormModal.tsx b/web/src/pages/NotifyChannelFormModal.tsx new file mode 100644 index 0000000..ad8514b --- /dev/null +++ b/web/src/pages/NotifyChannelFormModal.tsx @@ -0,0 +1,185 @@ +import type { FormEvent } from 'react' +import { useState } from 'react' +import { Loader2 } from 'lucide-react' +import toast from 'react-hot-toast' + +import { + notifyChannelsAPI, + type NotifyChannelInput, +} from '../api/notify_channels' +import type { NotifyChannel } from '../types' +import { + EMPTY_CONFIG, + EVENT_ALL, + EVENT_NONE, + EVENT_OPTIONS, + type EventMode, + initialEventMode, + normalizeInitialConfig, +} from './notifyChannelsModel' +import { Field } from './NotifyChannelFormField' +import { NotifyChannelEventFields } from './NotifyChannelEventFields' +import { + BarkFields, + EmailFields, + TelegramFields, + WebhookFields, + WechatFields, +} from './NotifyChannelProviderFields' + +type NotifyChannelFormModalProps = { + editing: NotifyChannel | null + onClose: () => void + onSaved: () => void | Promise +} + +export function NotifyChannelFormModal({ + editing, + onClose, + onSaved, +}: NotifyChannelFormModalProps) { + const [name, setName] = useState(editing?.name ?? '') + const [type, setType] = useState( + editing?.type ?? 'telegram', + ) + const [config, setConfig] = useState>( + normalizeInitialConfig(editing?.type ?? 'telegram', editing?.config ?? {}), + ) + const [events, setEvents] = useState(editing?.events ?? []) + const [eventMode, setEventMode] = useState(initialEventMode(editing?.events)) + const [enabled, setEnabled] = useState(editing?.enabled ?? true) + const [saving, setSaving] = useState(false) + + const onTypeChange = (t: NotifyChannel['type']) => { + setType(t) + setConfig({ ...EMPTY_CONFIG[t] }) + } + + const onSubmit = async (e: FormEvent) => { + e.preventDefault() + if (type === 'telegram') { + if (!String(config.bot_token ?? '').trim()) { + toast.error('请填写 Telegram Bot Token') + return + } + if (!String(config.admin_user_ids ?? '').trim()) { + toast.error('请填写管理员 Telegram ID') + return + } + } + const selectedEvents = events.filter((event) => EVENT_OPTIONS.some((item) => item.value === event)) + if (eventMode === 'custom' && selectedEvents.length === 0) { + toast.error('请至少选择一个推送事件,或选择关闭全部推送事件') + return + } + setSaving(true) + try { + const cleanedConfig = Object.fromEntries( + Object.entries(config).map(([key, value]) => [key, String(value ?? '').trim()]), + ) + delete cleanedConfig.chat_id + const input: NotifyChannelInput = { + name: name.trim(), + type: type, + config: cleanedConfig, + events: eventMode === 'all' ? [EVENT_ALL] : eventMode === 'none' ? [EVENT_NONE] : selectedEvents, + enabled, + } + if (editing) { + await notifyChannelsAPI.update(editing.id, input) + } else { + await notifyChannelsAPI.create(input) + } + toast.success('已保存') + await onSaved() + } catch (err: unknown) { + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '保存失败' + toast.error(msg) + } finally { + setSaving(false) + } + } + + const updateConfig = (k: string, v: string) => + setConfig((c) => ({ ...c, [k]: v })) + + const toggleEvent = (event: string) => { + setEvents((current) => + current.includes(event) + ? current.filter((item) => item !== event) + : [...current.filter((item) => item !== EVENT_ALL && item !== EVENT_NONE), event], + ) + } + + return ( +
+
+

+ {editing ? '编辑通知渠道' : '添加通知渠道'} +

+
+ + setName(e.target.value)} + /> + + + + + + + {type === 'telegram' && } + {type === 'wechat' && } + {type === 'bark' && } + {type === 'webhook' && } + {type === 'email' && } + + + + + +
+ + +
+ +
+
+ ) +} diff --git a/web/src/pages/NotifyChannelProviderFields.tsx b/web/src/pages/NotifyChannelProviderFields.tsx new file mode 100644 index 0000000..2b61a3d --- /dev/null +++ b/web/src/pages/NotifyChannelProviderFields.tsx @@ -0,0 +1,221 @@ +import { Field } from './NotifyChannelFormField' + +type ProviderFieldsProps = { + config: Record + updateConfig: (key: string, value: string) => void +} + +export function TelegramFields({ config, updateConfig }: ProviderFieldsProps) { + return ( + <> + + updateConfig('bot_token', e.target.value)} + /> + + + updateConfig('admin_user_ids', e.target.value)} + /> + + + updateConfig('group_chat_id', e.target.value)} + /> + + + updateConfig('channel_chat_id', e.target.value)} + /> + + + updateConfig('api_base_url', e.target.value)} + /> + + + updateConfig('proxy_url', e.target.value)} + /> + +
+ 群组 ID、频道 ID 均为选填,可填一个、两个都填,也可以不填。管理功能始终仅管理员 Telegram ID 或已绑定的本地管理员可用;普通用户需要在已配置的群组或频道中,才能使用 /start 用户名 密码 绑定账号、切换隐藏成人媒体库和目录。不配置群组/频道时,普通用户不会被放行。若测试通知超时,可填写反代 API 地址或代理地址。 +
+ + ) +} + +export function WechatFields({ config, updateConfig }: ProviderFieldsProps) { + return ( + + updateConfig('sendkey', e.target.value)} + /> + + ) +} + +export function BarkFields({ config, updateConfig }: ProviderFieldsProps) { + return ( + <> + + updateConfig('device_key', e.target.value)} + /> + + + updateConfig('server', e.target.value)} + /> + + + ) +} + +export function WebhookFields({ config, updateConfig }: ProviderFieldsProps) { + return ( + <> + + updateConfig('url', e.target.value)} + /> + + + + + +