From 1407b9b5c45c4fba3db880be8d3684fa16622afc Mon Sep 17 00:00:00 2001 From: truewhile <62226914+truewhile@users.noreply.github.com> Date: Sat, 5 Sep 2026 12:34:17 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=EF=BC=8C=E6=8E=92=E6=9F=A5?= =?UTF-8?q?=E9=A1=B9=E7=9B=AE=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/server/server_manager.go | 21 ++- go.mod | 2 +- go.sum | 4 +- internal/config/config.go | 16 +- internal/config/normalize.go | 10 +- internal/config/save.go | 11 +- internal/database/database.go | 9 ++ internal/database/schema_media_search.go | 2 +- internal/database/schema_migration.go | 14 +- internal/database/sqlite_migration_copy.go | 47 ++++-- internal/database/sqlite_runtime.go | 100 ++++++++++--- internal/handler/admin_settings.go | 24 ++- internal/handler/dlna.go | 39 +++++ internal/handler/emby_users.go | 29 ++-- internal/handler/media.go | 6 +- internal/handler/playback.go | 39 ++++- internal/handler/playlist_extra.go | 9 +- internal/handler/stats_user.go | 21 +++ internal/handler/system_extra.go | 7 +- internal/handler/ws.go | 21 ++- internal/middleware/emby_auth.go | 13 ++ internal/middleware/rate_limiter.go | 17 ++- internal/model/api_config.go | 23 ++- internal/model/api_config_legacy.go | 27 +--- internal/model/library_media.go | 4 +- internal/model/model.go | 1 - internal/repository/api_config_repository.go | 51 ++++--- .../repository/media_repository_upsert.go | 91 ++++++++---- .../repository/media_search_repository.go | 14 +- internal/repository/permission_repository.go | 38 ++++- internal/repository/scrape_task_repository.go | 34 +++++ internal/repository/strm_repository.go | 63 ++++++++ .../service/adult_scraper_routing_test.go | 2 +- internal/service/api_config_connection.go | 12 +- internal/service/api_config_helpers.go | 6 +- internal/service/api_config_svc.go | 10 +- internal/service/cloud/clouddrive2.go | 11 +- internal/service/cloud/clouddrive2_dav.go | 5 +- .../service/cloud/clouddrive2_mutation.go | 64 +++++--- .../service/cloud/clouddrive2_openlist.go | 81 +++++++++- internal/service/cloud/pan115_openapi.go | 13 +- internal/service/cloud115/client.go | 111 ++++++++++++-- internal/service/cloud115/oauth.go | 7 +- internal/service/cloud115/open.go | 43 ++++-- internal/service/cloud115/oss_multipart.go | 113 +++++++++----- internal/service/cloud115/upload.go | 8 +- internal/service/cloud115/utils.go | 18 ++- internal/service/dlna.go | 8 +- internal/service/emby_compat.go | 100 +++++++++---- internal/service/emby_items_detail.go | 23 ++- internal/service/emby_remote.go | 85 +++++++---- internal/service/emby_remote_lines.go | 24 ++- internal/service/emby_remote_web.go | 76 +++++++--- internal/service/emby_user_data.go | 16 ++ internal/service/image_proxy.go | 40 +++++ internal/service/image_proxy_paths.go | 11 +- .../service/organizer_directory_versions.go | 51 +++++-- internal/service/playback.go | 33 ++++- internal/service/runtime_settings.go | 2 + internal/service/scanner_prune.go | 27 +++- internal/service/scraper_queue.go | 44 +++++- internal/service/strm_115_oauth.go | 43 ++++++ internal/service/strm_queue.go | 88 ++++++++++- internal/service/strm_service.go | 27 ++++ internal/service/strm_sync.go | 138 +++++++++++++----- internal/service/watcher.go | 21 ++- internal/service/ws_hub.go | 7 + web/src/components/LayoutHeaderSections.tsx | 9 +- web/src/hooks/useSSE.ts | 123 ---------------- web/src/hooks/useWebSocket.ts | 17 ++- web/src/pages/AdminLibraryPanel.tsx | 1 + web/src/pages/AdminLibraryPanelSections.tsx | 9 +- web/src/pages/AdminUsersForm.tsx | 14 +- web/src/pages/AdminUsersPanel.tsx | 8 +- web/src/pages/AdultSettingsPanel.tsx | 21 +++ web/src/pages/PlayerPage.tsx | 46 ++++-- web/src/pages/PlayerVideoStage.tsx | 5 + web/src/pages/ScraperQueuePage.tsx | 30 +++- web/src/pages/StrmDialogs.tsx | 7 +- web/src/pages/StrmQueuePage.tsx | 30 +++- web/src/pages/WatchHistoryPage.tsx | 10 ++ .../pages/strm-dialogs/Strm115AuthPanel.tsx | 18 ++- web/src/pages/useAdminLibraryPanel.ts | 9 +- web/src/pages/useLibraryData.ts | 12 +- web/src/pages/useMediaDetailPageState.ts | 34 +++-- 85 files changed, 1940 insertions(+), 638 deletions(-) delete mode 100644 web/src/hooks/useSSE.ts diff --git a/cmd/server/server_manager.go b/cmd/server/server_manager.go index 02ff1da..abe2526 100644 --- a/cmd/server/server_manager.go +++ b/cmd/server/server_manager.go @@ -130,26 +130,35 @@ func (m *serverManager) Shutdown(ctx context.Context) error { // desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。 // 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。 func (m *serverManager) desiredPair() (*tlsPair, error) { - if m.cfg == nil || !m.cfg.App.HTTPSEnabled { + // 与 ApplyRuntimeSetting 的写锁配对:HTTPS 相关字段可能被运行时设置 + // 热更新,无锁读存在数据竞争(string 撕裂)。 + config.RuntimeMu.RLock() + httpsEnabled := m.cfg != nil && m.cfg.App.HTTPSEnabled + cert := m.cfg.App.SSLCert + certPath := m.cfg.App.SSLCertPath + key := m.cfg.App.SSLKey + keyPath := m.cfg.App.SSLKeyPath + config.RuntimeMu.RUnlock() + if !httpsEnabled { return nil, nil } - certPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLCert, m.cfg.App.SSLCertPath, "证书") + certPEM, err := service.ResolveSSLMaterial(cert, certPath, "证书") if err != nil { return nil, err } - keyPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLKey, m.cfg.App.SSLKeyPath, "私钥") + keyPEM, err := service.ResolveSSLMaterial(key, keyPath, "私钥") if err != nil { return nil, err } if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil { return nil, err } - cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + pairCert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) if err != nil { return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err) } return &tlsPair{ - cert: cert, + cert: pairCert, certPEM: certPEM, keyPEM: keyPEM, version: certPEM + "\x00" + keyPEM, @@ -171,6 +180,8 @@ func (m *serverManager) maybeStartAutoReloadLocked() { // pathBased 是否至少有一侧证书/私钥通过文件路径配置。 func (m *serverManager) pathBased() bool { + config.RuntimeMu.RLock() + defer config.RuntimeMu.RUnlock() return strings.TrimSpace(m.cfg.App.SSLCertPath) != "" || strings.TrimSpace(m.cfg.App.SSLKeyPath) != "" } diff --git a/go.mod b/go.mod index 23b85d9..6c2025d 100644 --- a/go.mod +++ b/go.mod @@ -8,7 +8,7 @@ require ( github.com/gin-contrib/gzip v1.2.6 github.com/gin-gonic/gin v1.12.0 github.com/glebarez/sqlite v1.11.0 - github.com/golang-jwt/jwt/v5 v5.2.0 + github.com/golang-jwt/jwt/v5 v5.2.2 github.com/google/uuid v1.6.0 github.com/gorilla/websocket v1.5.3 github.com/redis/go-redis/v9 v9.7.0 diff --git a/go.sum b/go.sum index 1d6b93c..b1210d6 100644 --- a/go.sum +++ b/go.sum @@ -52,8 +52,8 @@ github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= -github.com/golang-jwt/jwt/v5 v5.2.0 h1:d/ix8ftRUorsN+5eMIlF4T6J8CAt9rch3My2winC1Jw= -github.com/golang-jwt/jwt/v5 v5.2.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= +github.com/golang-jwt/jwt/v5 v5.2.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8= +github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= diff --git a/internal/config/config.go b/internal/config/config.go index 263c876..24eb7d3 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -12,6 +12,7 @@ import ( "os" "path/filepath" "strings" + "sync" "github.com/spf13/viper" ) @@ -19,6 +20,12 @@ import ( // EnvPrefix 是所有环境变量驱动的覆盖使用的前缀。 const EnvPrefix = "MeBox" +// RuntimeMu 保护运行时热更新配置字段的并发读写:ApplyRuntimeSetting 在 +// HTTP goroutine 中写字段,serverManager 的证书轮询等后台协程在无锁读取 +// 同一批字段。string 是双字结构,无锁并发读写可读到撕裂的 header。 +// 写方在 ApplyRuntimeSetting 内 Lock,读方(cmd/server)在轮询处 RLock。 +var RuntimeMu sync.RWMutex + // Load 从默认值 / 文件 / 环境读取配置。 // // 即使没有文件也始终返回可用的 Config。 @@ -45,8 +52,13 @@ func Load() (*Config, error) { } s := viper.New() s.SetConfigFile(filepath.Join("config", e.Name())) - if err := s.ReadInConfig(); err == nil { - _ = v.MergeConfigMap(s.AllSettings()) + if err := s.ReadInConfig(); err != nil { + // 分片解析失败不能静默吞掉:database.yaml 语法错误会让 + // database.dsn 缺失 → type=auto 静默回退 SQLite,新数据 + // 全部写进一个空库而用户无感知。 + fmt.Fprintf(os.Stderr, "warning: parse config/%s failed: %v\n", e.Name(), err) + } else { + v.MergeConfigMap(s.AllSettings()) } } } diff --git a/internal/config/normalize.go b/internal/config/normalize.go index d811c9b..f971768 100644 --- a/internal/config/normalize.go +++ b/internal/config/normalize.go @@ -68,8 +68,14 @@ func (c *Config) normalize() error { return fmt.Errorf("generate jwt secret: %w", err) } c.Secrets.JWTSecret = hex.EncodeToString(buf) - _ = os.MkdirAll(c.App.DataDir, 0o750) - _ = os.WriteFile(path, []byte(c.Secrets.JWTSecret), 0o600) + // 持久化失败(DataDir 只读/权限异常)会导致每次重启重新生成 + // 密钥、全部会话静默失效、多实例各持不同 secret——必须让 + // 操作员感知。 + if mkErr := os.MkdirAll(c.App.DataDir, 0o750); mkErr != nil { + fmt.Fprintf(os.Stderr, "warning: persist jwt secret failed (mkdir): %v\n", mkErr) + } else if wErr := os.WriteFile(path, []byte(c.Secrets.JWTSecret), 0o600); wErr != nil { + fmt.Fprintf(os.Stderr, "warning: persist jwt secret failed (write): %v\n", wErr) + } } } return nil diff --git a/internal/config/save.go b/internal/config/save.go index cbedb10..3b84617 100644 --- a/internal/config/save.go +++ b/internal/config/save.go @@ -34,8 +34,15 @@ func SaveDatabaseConfig(dbType, dsn string) error { return fmt.Errorf("marshal config.yaml: %w", err) } - if err := os.WriteFile(configPath, out, 0644); err != nil { - return fmt.Errorf("write config.yaml: %w", err) + // 原子写:临时文件 + rename,避免进程崩溃/断电留下截断的 config.yaml + // (下次启动会硬失败);DSN 含数据库密码,权限收窄到 0600。 + tmp := configPath + ".tmp" + if err := os.WriteFile(tmp, out, 0o600); err != nil { + return fmt.Errorf("write config.yaml.tmp: %w", err) + } + if err := os.Rename(tmp, configPath); err != nil { + _ = os.Remove(tmp) + return fmt.Errorf("replace config.yaml: %w", err) } return nil } diff --git a/internal/database/database.go b/internal/database/database.go index 39103ec..c7aa75d 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "strings" + "time" "github.com/glebarez/sqlite" "go.uber.org/zap" @@ -73,6 +74,14 @@ func configureConnectionPool(db *gorm.DB, cfg *config.Config) error { if cfg.Database.MaxIdleConns > 0 { sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns) } + // 连接生命周期:默认 0 意味着 Postgres 重启/故障切换后的陈旧连接 + // 永不过期,首次复用才报错,运行期断连恢复慢且可能批量报错。 + if isPostgres(db) { + sqlDB.SetConnMaxLifetime(time.Hour) + sqlDB.SetConnMaxIdleTime(10 * time.Minute) + } else if isSQLite(db) { + sqlDB.SetConnMaxLifetime(24 * time.Hour) + } return nil } diff --git a/internal/database/schema_media_search.go b/internal/database/schema_media_search.go index f38ba45..641fa8a 100644 --- a/internal/database/schema_media_search.go +++ b/internal/database/schema_media_search.go @@ -9,7 +9,7 @@ const mediaSearchIndexSchemaVersion = 2 func ensureMediaSearchIndex(db *gorm.DB) error { if err := ensureMediaSearchMetaTable(db); err != nil { - return nil + return err // meta 表创建失败必须上抛,不能静默掩盖 } version := currentMediaSearchIndexVersion(db) if version != mediaSearchIndexSchemaVersion { diff --git a/internal/database/schema_migration.go b/internal/database/schema_migration.go index ff6a393..3c2048e 100644 --- a/internal/database/schema_migration.go +++ b/internal/database/schema_migration.go @@ -132,13 +132,19 @@ func ensureEmbyMountsCompatibility(db *gorm.DB) error { return err } } - // 针对已有数据:如果存在多个 sort_order=0/NULL 的记录,按创建时间顺序赋予稳定递增的序号 + // 针对已有数据:只给 sort_order=0/NULL 的行按创建时间补号(从现有 + // 最大值之后递增),不能整表重排——此前无条件按 created_at 从 0 重新 + // 编号,会把用户自定义的顺序覆盖掉。 var zeroCount int64 - if err := db.Model(&model.EmbyMount{}).Where("sort_order = 0 OR sort_order IS NULL").Count(&zeroCount).Error; err == nil && zeroCount > 1 { + if err := db.Model(&model.EmbyMount{}).Where("sort_order = 0 OR sort_order IS NULL").Count(&zeroCount).Error; err == nil && zeroCount > 0 { + // max 只统计非 0 行:sort_order=0 与 NULL 同样视为“未分配”, + // 全部为 0 时从 0 开始编号(与迁移前的初始化语义一致)。 + var maxOrder int + _ = db.Raw("SELECT COALESCE(MAX(sort_order), -1) FROM emby_mounts WHERE sort_order > 0").Scan(&maxOrder).Error var mounts []model.EmbyMount - if err := db.Order("created_at asc, id asc").Find(&mounts).Error; err == nil { + if err := db.Where("sort_order = 0 OR sort_order IS NULL").Order("created_at asc, id asc").Find(&mounts).Error; err == nil { for i, m := range mounts { - _ = db.Exec("UPDATE emby_mounts SET sort_order = ? WHERE id = ?", i, m.ID).Error + _ = db.Exec("UPDATE emby_mounts SET sort_order = ? WHERE id = ?", maxOrder+1+i, m.ID).Error } } } diff --git a/internal/database/sqlite_migration_copy.go b/internal/database/sqlite_migration_copy.go index 748af90..038ec20 100644 --- a/internal/database/sqlite_migration_copy.go +++ b/internal/database/sqlite_migration_copy.go @@ -50,28 +50,43 @@ func copyModelTables(src, target *gorm.DB, batchSize int) (map[string]int64, int if modelType.Kind() != reflect.Ptr { return tableCounts, totalCopied, 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 tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err) - } - filtered := slicePtr.Elem() + var primaryKeySet map[string]struct{} if targetCount > 0 { - primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns) + primaryKeySet, err = targetPrimaryKeySet(target, table, primaryColumns) if err != nil { return tableCounts, totalCopied, err } - filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet) } - if filtered.Len() == 0 { - continue + // 分页流式读取:此前整表一次性 Find 进内存,media 表几十万行、 + // 每行含 overview/genres 等长文本时可达数百 MB,迁移过程有 OOM + // 风险。源库在迁移期间是静态的,offset 分页安全。 + const readBatch = 1000 + copiedForTable := int64(0) + for offset := 0; ; offset += readBatch { + batchPtr := reflect.New(reflect.SliceOf(modelType.Elem())) + if err := src.Unscoped().Limit(readBatch).Offset(offset).Find(batchPtr.Interface()).Error; err != nil { + return tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err) + } + batch := batchPtr.Elem() + if batch.Len() == 0 { + break + } + filtered := batch + if primaryKeySet != nil { + filtered = filterRowsMissingInTarget(target, table, primaryColumns, batch, primaryKeySet) + } + if filtered.Len() > 0 { + 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 tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err) + } + copiedForTable += int64(filtered.Len()) + } + if batch.Len() < readBatch { + break + } } - 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 tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err) - } - copiedForTable := int64(filtered.Len()) tableCounts[table] = copiedForTable totalCopied += copiedForTable } diff --git a/internal/database/sqlite_runtime.go b/internal/database/sqlite_runtime.go index df0c21f..f06d3df 100644 --- a/internal/database/sqlite_runtime.go +++ b/internal/database/sqlite_runtime.go @@ -5,12 +5,21 @@ import ( "fmt" "path/filepath" "strings" + "sync" + "sync/atomic" + "time" "gorm.io/gorm" "github.com/truewhile/MeBox/internal/config" ) +// sqliteGateHoldLimit 是写闸持有者的最长合法持有时长。语句级写闸在 SQL +// 执行 panic 时 After 回调不会运行,令牌会泄漏并让后续所有写入永久等锁; +// 超过该时长的持有者按泄漏强制回收(60s 内单条写语句远未到,正常写路径 +// 不受影响)。 +const sqliteGateHoldLimit = 60 * time.Second + func installSQLiteWriteGate(db *gorm.DB) { if db == nil { return @@ -22,15 +31,18 @@ func installSQLiteWriteGate(db *gorm.DB) { if tx.Statement != nil && tx.Statement.Context != nil { ctx = tx.Statement.Context } - if err := gate.Lock(ctx); err != nil { + holder, err := gate.Lock(ctx) + if err != nil { _ = tx.AddError(err) return } - tx.InstanceSet(lockedKey, struct{}{}) + tx.InstanceSet(lockedKey, holder) } unlock := func(tx *gorm.DB) { - if _, ok := tx.InstanceGet(lockedKey); ok { - gate.Unlock() + if holder, ok := tx.InstanceGet(lockedKey); ok { + if h, ok := holder.(*sqliteGateHolder); ok { + gate.Unlock(h) + } } } rawLock := func(tx *gorm.DB) { @@ -64,38 +76,76 @@ func isReadOnlySQL(sql string) bool { return false } -// sqliteWriteGate serializes in-process SQLite writes while respecting the -// statement context, so request cancellation can break out of a queued write. +// sqliteWriteGate serializes in-process SQLite writes. 所有权令牌(而非裸 +// 信号量)保证只有持有者本人能释放;持有超时按泄漏自动回收,避免一次 +// panic 让进程的 SQLite 写入半永久性瘫痪。 type sqliteWriteGate struct { - ch chan struct{} + mu sync.Mutex + cond *sync.Cond + owner *sqliteGateHolder } +type sqliteGateHolder struct { + id uint64 + acquired time.Time +} + +var sqliteGateHolderSeq atomic.Uint64 + func newSQLiteWriteGate() *sqliteWriteGate { - return &sqliteWriteGate{ch: make(chan struct{}, 1)} + g := &sqliteWriteGate{} + g.cond = sync.NewCond(&g.mu) + return g } -func (g *sqliteWriteGate) Lock(ctx context.Context) error { - select { - case g.ch <- struct{}{}: - return nil - default: - } +func (g *sqliteWriteGate) Lock(ctx context.Context) (*sqliteGateHolder, error) { + g.mu.Lock() + defer g.mu.Unlock() if ctx == nil { ctx = context.Background() } - select { - case g.ch <- struct{}{}: - return nil - case <-ctx.Done(): - return ctx.Err() + // ctx 取消时唤醒等待者(cond 无法感知 ctx,用旁路 goroutine 广播)。 + if done := ctx.Done(); done != nil { + stop := make(chan struct{}) + defer close(stop) + go func() { + select { + case <-done: + g.cond.Broadcast() + case <-stop: + } + }() + } + for { + if g.owner == nil { + holder := &sqliteGateHolder{ + id: sqliteGateHolderSeq.Add(1), + acquired: time.Now(), + } + g.owner = holder + return holder, nil + } + if ctx.Err() != nil { + return nil, ctx.Err() + } + if time.Since(g.owner.acquired) > sqliteGateHoldLimit { + // 持有者疑似 panic 泄漏(After 回调未执行):强制回收。 + g.owner = nil + g.cond.Broadcast() + continue + } + g.cond.Wait() } } -func (g *sqliteWriteGate) Unlock() { - select { - case <-g.ch: - default: +func (g *sqliteWriteGate) Unlock(h *sqliteGateHolder) { + g.mu.Lock() + defer g.mu.Unlock() + if h == nil || g.owner != h { + return } + g.owner = nil + g.cond.Broadcast() } func buildSQLiteDSN(cfg *config.Config) string { @@ -104,7 +154,9 @@ func buildSQLiteDSN(cfg *config.Config) string { // keep as-is to respect user-provided relative paths. dbPath = filepath.Clean(dbPath) } - dsn := dbPath + "?_pragma=foreign_keys(1)" + // _txlock=immediate:事务以写锁开始。此前 deferred BEGIN 在并发事务 + // 升级写锁时会绕过 busy_timeout 直接报 SQLITE_BUSY。 + dsn := dbPath + "?_txlock=immediate&_pragma=foreign_keys(1)" if cfg.Database.WALMode { dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)" } diff --git a/internal/handler/admin_settings.go b/internal/handler/admin_settings.go index 7498ebc..f7ce6d9 100644 --- a/internal/handler/admin_settings.go +++ b/internal/handler/admin_settings.go @@ -10,6 +10,7 @@ import ( "github.com/gin-gonic/gin" "go.uber.org/zap" + "github.com/truewhile/MeBox/internal/config" "github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/service" ) @@ -86,10 +87,19 @@ func applyHTTPSetting(svc *service.Container, key, value string) error { svc.Log.Warn("https setting saved but not applied yet", zap.String("key", key), zap.String("reason", reason)) } } + + config.RuntimeMu.RLock() + httpsEnabled := svc.Cfg.App.HTTPSEnabled + cert := svc.Cfg.App.SSLCert + certPath := svc.Cfg.App.SSLCertPath + keyMaterial := svc.Cfg.App.SSLKey + keyPath := svc.Cfg.App.SSLKeyPath + config.RuntimeMu.RUnlock() + switch key { case "https.enabled": - if svc.Cfg.App.HTTPSEnabled { - if _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath); err != nil { + if httpsEnabled { + if _, err := service.ResolveSSLKeyPair(cert, certPath, keyMaterial, keyPath); err != nil { return fmt.Errorf("启用 HTTPS 失败:%v", err) } } @@ -97,7 +107,7 @@ func applyHTTPSetting(svc *service.Container, key, value string) error { if err := validateSSLMaterialSource(key, value); err != nil { return err } - if !svc.Cfg.App.HTTPSEnabled { + if !httpsEnabled { return nil } if !httpsPairReady(svc) { @@ -144,7 +154,13 @@ func validateSSLMaterialSource(key, value string) error { // httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。 func httpsPairReady(svc *service.Container) bool { - _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath) + config.RuntimeMu.RLock() + cert := svc.Cfg.App.SSLCert + certPath := svc.Cfg.App.SSLCertPath + keyMaterial := svc.Cfg.App.SSLKey + keyPath := svc.Cfg.App.SSLKeyPath + config.RuntimeMu.RUnlock() + _, err := service.ResolveSSLKeyPair(cert, certPath, keyMaterial, keyPath) return err == nil } diff --git a/internal/handler/dlna.go b/internal/handler/dlna.go index 08b2e3e..733ae93 100644 --- a/internal/handler/dlna.go +++ b/internal/handler/dlna.go @@ -3,6 +3,8 @@ package handler import ( "net/http" + "net/url" + "strings" "github.com/gin-gonic/gin" @@ -33,6 +35,21 @@ func dlnaCastHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } + // SSRF 防护:control_url 必须命中本服务发现到的真实渲染设备, + // 防止登录用户借 cast 接口向任意内网地址发起 POST。 + // 优先用 30s 缓存;未命中时强制重扫一次再校验(设备可能刚上线)。 + devices, err := svc.DLNA.Discover(c.Request.Context(), false) + if err == nil && !dlnaControlURLKnown(devices, req.ControlURL) { + devices, err = svc.DLNA.Discover(c.Request.Context(), true) + } + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + if !dlnaControlURLKnown(devices, req.ControlURL) { + c.JSON(http.StatusBadRequest, gin.H{"error": "unknown DLNA device: control_url must come from /api/dlna discovery"}) + return + } if err := svc.DLNA.Cast(c.Request.Context(), req.ControlURL, req.MediaURL); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return @@ -40,3 +57,25 @@ func dlnaCastHandler(svc *service.Container) gin.HandlerFunc { c.Status(http.StatusNoContent) } } + +// dlnaControlURLKnown 判断 control_url 是否属于发现列表中的设备。 +// 按解析后的 host:port+path 精确比对,容忍大小写与尾斜杠差异。 +func dlnaControlURLKnown(devices []service.DLNADevice, controlURL string) bool { + want, err := url.Parse(strings.TrimSpace(controlURL)) + if err != nil || want.Host == "" { + return false + } + for _, dev := range devices { + for _, candidate := range []string{dev.ControlURL, dev.Location} { + u, err := url.Parse(strings.TrimSpace(candidate)) + if err != nil || u.Host == "" { + continue + } + if strings.EqualFold(u.Host, want.Host) && + strings.EqualFold(strings.TrimRight(u.Path, "/"), strings.TrimRight(want.Path, "/")) { + return true + } + } + } + return false +} diff --git a/internal/handler/emby_users.go b/internal/handler/emby_users.go index e810202..a0419c4 100644 --- a/internal/handler/emby_users.go +++ b/internal/handler/emby_users.go @@ -132,22 +132,25 @@ func embyMeHandler(svc *service.Container) gin.HandlerFunc { func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { - u, err := svc.Emby.FindUser(c.Request.Context(), c.Param("userId")) + uid := embyUserID(c) + if uid == "" { + embyError(c, http.StatusUnauthorized, "not authenticated") + return + } + // 只返回调用者自己的用户对象:客户端误传其他 userId 时回退到 + // 调用者自身(保留旧行为的兼容语义),但绝不返回他人数据。 + u, err := svc.Emby.FindUser(c.Request.Context(), uid) 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"))) + c.JSON(http.StatusOK, embyFallbackUser(uid)) } } +// embyFallbackUser 是查库失败时的最后兜底(保持客户端可渲染)。 +// Policy 必须是最小权限:不声明管理员/删除内容/控制他人等能力, +// 实际权限始终由服务端各路由的校验决定。 func embyFallbackUser(id string) gin.H { if strings.TrimSpace(id) == "" { id = "mebox-user" @@ -161,10 +164,10 @@ func embyFallbackUser(id string) gin.H { "HasConfiguredEasyPassword": false, "EnableAutoLogin": false, "Policy": gin.H{ - "IsAdministrator": true, - "EnableContentDeletion": true, - "EnableRemoteControlOfOtherUsers": true, - "EnableSharedDeviceControl": true, + "IsAdministrator": false, + "EnableContentDeletion": false, + "EnableRemoteControlOfOtherUsers": false, + "EnableSharedDeviceControl": false, "EnableRemoteAccess": true, "EnableAllDevices": true, "EnableAllChannels": true, diff --git a/internal/handler/media.go b/internal/handler/media.go index afefcc1..ba11043 100644 --- a/internal/handler/media.go +++ b/internal/handler/media.go @@ -368,7 +368,11 @@ func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc { } uid, _ := c.Get("ctx_user_id") svc.Audit.Record(c.Request.Context(), toString(uid), "library.delete", id, c.ClientIP(), "") - go func() { _ = svc.Watcher.Refresh(context.Background()) }() + // goroutine 内的 panic 无法被 gin.Recovery 捕获,会直接崩掉进程: + // 与其他调用点一致先判空。 + if svc.Watcher != nil { + go func() { _ = svc.Watcher.Refresh(context.Background()) }() + } c.Status(http.StatusNoContent) } } diff --git a/internal/handler/playback.go b/internal/handler/playback.go index e40ad1d..c4cec5c 100644 --- a/internal/handler/playback.go +++ b/internal/handler/playback.go @@ -2,9 +2,11 @@ package handler import ( + "errors" "net/http" "github.com/gin-gonic/gin" + "gorm.io/gorm" "github.com/truewhile/MeBox/internal/middleware" "github.com/truewhile/MeBox/internal/service" @@ -157,6 +159,25 @@ type playlistItemReq struct { MediaID string `json:"media_id" binding:"required"` } +// playlistWriteGuard 校验当前用户对播放列表的写权限(属主或 admin)。 +// 校验失败时已写入错误响应,调用方直接 return。 +func playlistWriteGuard(c *gin.Context, svc *service.Container, playlistID string) (string, bool, bool) { + uid, _ := c.Get(middleware.CtxUserID) + role, _ := c.Get(middleware.CtxUserRole) + isAdmin := role == "admin" + if err := svc.Playback.EnsurePlaylistOwner(c.Request.Context(), playlistID, uid.(string), isAdmin); err != nil { + if errors.Is(err, service.ErrPlaylistForbidden) { + c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"}) + } else if errors.Is(err, gorm.ErrRecordNotFound) { + c.JSON(http.StatusNotFound, gin.H{"error": "playlist not found"}) + } else { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + } + return "", isAdmin, false + } + return uid.(string), isAdmin, true +} + func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { var req playlistItemReq @@ -164,8 +185,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } + uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id")) + if !ok { + return + } if err := svc.Playback.AddToPlaylist( - c.Request.Context(), c.Param("id"), req.MediaID, + c.Request.Context(), c.Param("id"), uid, req.MediaID, isAdmin, ); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return @@ -176,8 +201,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc { func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { + uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id")) + if !ok { + return + } if err := svc.Playback.RemoveFromPlaylist( - c.Request.Context(), c.Param("id"), c.Param("media_id"), + c.Request.Context(), c.Param("id"), uid, c.Param("media_id"), isAdmin, ); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return @@ -188,8 +217,12 @@ func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc { func deletePlaylistHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { + uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id")) + if !ok { + return + } if err := svc.Playback.DeletePlaylist( - c.Request.Context(), c.Param("id"), + c.Request.Context(), c.Param("id"), uid, isAdmin, ); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return diff --git a/internal/handler/playlist_extra.go b/internal/handler/playlist_extra.go index 525d73a..6f3939b 100644 --- a/internal/handler/playlist_extra.go +++ b/internal/handler/playlist_extra.go @@ -27,6 +27,9 @@ func reorderPlaylistHandler(svc *service.Container) gin.HandlerFunc { return } pid := c.Param("id") + if _, _, ok := playlistWriteGuard(c, svc, pid); !ok { + return + } for i, mid := range req.Order { if err := svc.Repo.DB.WithContext(c.Request.Context()). Model(&model.PlaylistItem{}). @@ -44,8 +47,12 @@ func reorderPlaylistHandler(svc *service.Container) gin.HandlerFunc { // /playlists/:id/items/:item_id (vs. the existing /:media_id variant). func deletePlaylistItemByIDHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { + pid := c.Param("id") + if _, _, ok := playlistWriteGuard(c, svc, pid); !ok { + return + } if err := svc.Repo.DB.WithContext(c.Request.Context()). - Where("playlist_id = ? AND id = ?", c.Param("id"), c.Param("item_id")). + Where("playlist_id = ? AND id = ?", pid, c.Param("item_id")). Delete(&model.PlaylistItem{}).Error; err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return diff --git a/internal/handler/stats_user.go b/internal/handler/stats_user.go index ded5cd8..af24ab6 100644 --- a/internal/handler/stats_user.go +++ b/internal/handler/stats_user.go @@ -14,9 +14,14 @@ import ( ) // statsUserHandler returns a watch-time summary for one user. +// 观看统计是隐私数据:仅允许本人或管理员查询。 func statsUserHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { uid := c.Param("id") + if !statsCallerAllowed(c, uid) { + c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"}) + return + } var watched int64 _ = svc.Repo.DB.Model(&model.PlaybackHistory{}). Where("user_id = ?", uid). @@ -35,8 +40,14 @@ func statsUserHandler(svc *service.Container) gin.HandlerFunc { } // statsTopUsersHandler returns the most active users by play count. +// 全员排行含用户名与精确时长,仅管理员可查。 func statsTopUsersHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { + role, _ := c.Get(middleware.CtxUserRole) + if role != "admin" { + c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"}) + return + } limit, _ := strconv.Atoi(c.DefaultQuery("limit", "10")) if limit <= 0 || limit > 50 { limit = 10 @@ -109,3 +120,13 @@ func statsPlayHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusOK, gin.H{"ok": true}) } } + +// statsCallerAllowed 判断当前调用者是否允许查看 uid 的观看统计。 +func statsCallerAllowed(c *gin.Context, uid string) bool { + role, _ := c.Get(middleware.CtxUserRole) + if role == "admin" { + return true + } + caller, _ := c.Get(middleware.CtxUserID) + return toString(caller) == uid +} diff --git a/internal/handler/system_extra.go b/internal/handler/system_extra.go index 5540e63..aac77f1 100644 --- a/internal/handler/system_extra.go +++ b/internal/handler/system_extra.go @@ -39,11 +39,16 @@ func listSystemConfigHandler(svc *service.Container) gin.HandlerFunc { } func isSecretKey(k string) bool { - for _, suffix := range []string{".token", ".secret", ".password", ".api_key", ".cookie"} { + for _, suffix := range []string{".token", ".secret", ".password", ".api_key", ".cookie", ".pin"} { if endsWith(k, suffix) { return true } } + // 非后缀型敏感键:可触发服务端任意命令的更新命令等。 + switch k { + case "system.update.command": + return true + } return false } diff --git a/internal/handler/ws.go b/internal/handler/ws.go index fb1e358..5ba80ce 100644 --- a/internal/handler/ws.go +++ b/internal/handler/ws.go @@ -9,6 +9,8 @@ package handler import ( "encoding/json" "net/http" + "net/url" + "strings" "time" "github.com/gin-gonic/gin" @@ -21,10 +23,21 @@ import ( var wsUpgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, - // Allow any origin: the AuthRequired middleware already validated the - // JWT before we got here, and we never serve sensitive cross-domain - // state through the socket. - CheckOrigin: func(_ *http.Request) bool { return true }, + // 同源校验:浏览器跨站页面虽读不到 ?token=,但可能借 cookie 通道 + // (extractToken 接受 msgo_access_token cookie)发起跨站 WebSocket + // 劫持。放行同源与非浏览器客户端(不发 Origin 头的 App/脚本), + // 拒绝跨站 Origin。 + CheckOrigin: func(r *http.Request) bool { + origin := strings.TrimSpace(r.Header.Get("Origin")) + if origin == "" { + return true + } + u, err := url.Parse(origin) + if err != nil || u.Host == "" { + return false + } + return strings.EqualFold(u.Host, r.Host) + }, } func wsHandler(svc *service.Container) gin.HandlerFunc { diff --git a/internal/middleware/emby_auth.go b/internal/middleware/emby_auth.go index adeb720..c4ab8fc 100644 --- a/internal/middleware/emby_auth.go +++ b/internal/middleware/emby_auth.go @@ -51,6 +51,19 @@ func EmbyAuthRequired(secret string) gin.HandlerFunc { return } + // 用途限定令牌(如 external_play,签发给外链播放器且绑定单一 + // media)只允许走 /api/stream|/hls|/cloud/play,绝不能作为全功能 + // 凭据访问 Emby 兼容面;否则外链 URL 一旦泄漏,持有者可获得 + // 该用户最长 24h 的全部 Emby API 权限。 + if strings.TrimSpace(claims.Purpose) != "" { + c.JSON(http.StatusUnauthorized, gin.H{ + "Code": 40101, + "Message": "Invalid token", + }) + c.Abort() + return + } + c.Set(EmbyCtxUserID, claims.UserID) c.Set(CtxUserID, claims.UserID) c.Set(CtxUserRole, claims.Role) diff --git a/internal/middleware/rate_limiter.go b/internal/middleware/rate_limiter.go index f3015bd..29d9b1e 100644 --- a/internal/middleware/rate_limiter.go +++ b/internal/middleware/rate_limiter.go @@ -16,6 +16,8 @@ type RateLimiter struct { window time.Duration max int requests map[string][]time.Time + stop chan struct{} + stopped sync.Once } // NewRateLimiter creates a rate limiter allowing max requests per window @@ -25,14 +27,27 @@ func NewRateLimiter(max int, window time.Duration) *RateLimiter { window: window, max: max, requests: make(map[string][]time.Time), + stop: make(chan struct{}), } go rl.cleanup() return rl } +// Close 停止后台清理 goroutine:清理循环此前无停止机制,每建一个实例 +// 就永久滞留一条 goroutine(测试场景会随实例创建不断累积)。 +func (rl *RateLimiter) Close() { + rl.stopped.Do(func() { close(rl.stop) }) +} + func (rl *RateLimiter) cleanup() { + ticker := time.NewTicker(5 * time.Minute) + defer ticker.Stop() for { - time.Sleep(5 * time.Minute) + select { + case <-rl.stop: + return + case <-ticker.C: + } rl.mu.Lock() now := time.Now() for ip, times := range rl.requests { diff --git a/internal/model/api_config.go b/internal/model/api_config.go index b382ed6..d3f3e4c 100644 --- a/internal/model/api_config.go +++ b/internal/model/api_config.go @@ -5,15 +5,22 @@ import ( "time" ) -// ApiConfig 存储第三方 API 密钥和配置信息。 -// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储。 -type ApiConfig struct { +// APIConfig 存储第三方 API 密钥和配置信息。 +// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储(AES-GCM 密文, +// base64 后常超 512 字符,因此必须是 text 而非 varchar(512))。 +// +// NOTE: 历史上曾有 APIConfig / ApiConfig 两个结构体映射到同一张 api_configs +// 表,AutoMigrate 每次启动互相改列(provider/api_key 长度来回切换),且 +// varchar(512) 收窄会让长密文入库后下一次启动迁移直接失败。现已合并为本 +// 结构体,字段取两者并集,请勿再拆分。 +type APIConfig struct { Base - Provider string `gorm:"size:64;uniqueIndex;not null" json:"provider"` - APIKey string `gorm:"size:512" json:"-"` - BaseURL string `gorm:"size:512" json:"base_url,omitempty"` - Extra string `gorm:"type:text" json:"extra,omitempty"` - Enabled bool `gorm:"default:true" json:"enabled"` + Provider string `gorm:"uniqueIndex;size:64;not null" json:"provider"` + APIKey string `gorm:"type:text" json:"-"` // ciphertext (never serialised) + BaseURL string `gorm:"size:512" json:"base_url,omitempty"` + Extra string `gorm:"type:text" json:"extra,omitempty"` // free-form JSON + Enabled bool `gorm:"default:true" json:"enabled"` + Description string `gorm:"size:255" json:"description,omitempty"` LastTestedAt *time.Time `json:"last_tested_at,omitempty"` TestResult string `gorm:"size:32" json:"test_result,omitempty"` diff --git a/internal/model/api_config_legacy.go b/internal/model/api_config_legacy.go index ae0eec7..340e204 100644 --- a/internal/model/api_config_legacy.go +++ b/internal/model/api_config_legacy.go @@ -1,23 +1,8 @@ package model -// APIConfig stores third-party data-source configuration. The api_key -// column is encrypted with AES-GCM (see internal/service/crypto.go) so an -// SQLite leak does not expose third-party credentials. -// -// Provider values mirror the original Python project: -// -// tmdb — themoviedb.org -// bangumi — bgm.tv -// thetvdb — thetvdb.com -// fanart — fanart.tv -// douban — douban.com (cookie) -// openai — OpenAI / DeepSeek / Qwen / Ollama (compatible) -type APIConfig struct { - Base - Provider string `gorm:"uniqueIndex;size:32;not null" json:"provider"` - APIKey string `gorm:"type:text" json:"-"` // ciphertext (never serialised) - BaseURL string `gorm:"size:512" json:"base_url,omitempty"` - Extra string `gorm:"type:text" json:"extra,omitempty"` // free-form JSON - Enabled bool `gorm:"default:true" json:"enabled"` - Description string `gorm:"size:255" json:"description,omitempty"` -} +// NOTE: 原 APIConfig(provider varchar(32) / api_key text)与 api_config.go +// 里的 ApiConfig(provider varchar(64) / api_key varchar(512))映射到同一张 +// api_configs 表,AutoMigrate 每次启动互相改列;且 api_key 被收窄成 +// varchar(512) 后,成人区/豆瓣等存的长 AES-GCM Cookie 密文一旦入库,下次 +// 启动迁移即失败、服务无法启动。两者已合并为 api_config.go 中唯一的 +// APIConfig 结构体(字段取并集),此处不再定义重复模型。 diff --git a/internal/model/library_media.go b/internal/model/library_media.go index de46ec7..a7a558f 100644 --- a/internal/model/library_media.go +++ b/internal/model/library_media.go @@ -26,7 +26,7 @@ type LibraryRoot struct { // Media 是单个可播放项。剧集链接到 SeriesID;电影 SeriesID == ""。 type Media struct { Base - LibraryID string `gorm:"index;size:36" json:"library_id"` + LibraryID string `gorm:"index;size:36;index:idx_media_library_release,priority:1" json:"library_id"` LibraryRootID string `gorm:"index;size:36" json:"library_root_id,omitempty"` SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"` Title string `gorm:"size:255;not null" json:"title"` @@ -46,7 +46,7 @@ type Media struct { Overview string `gorm:"type:text" json:"overview,omitempty"` Rating float32 `json:"rating"` Year int `json:"year"` - ReleaseDate string `gorm:"size:10;index" json:"release_date,omitempty"` + ReleaseDate string `gorm:"size:10;index:idx_media_library_release,priority:2" json:"release_date,omitempty"` SeasonNum int `json:"season_num"` EpisodeNum int `json:"episode_num"` ScrapeStatus string `gorm:"size:16;default:pending" json:"scrape_status"` diff --git a/internal/model/model.go b/internal/model/model.go index dfc1c76..751bb57 100644 --- a/internal/model/model.go +++ b/internal/model/model.go @@ -46,7 +46,6 @@ func AllModels() []interface{} { &APIConfig{}, &UserPermission{}, &RefreshToken{}, - &ApiConfig{}, &PlayProfile{}, &RegistrationCode{}, &SignIn{}, diff --git a/internal/repository/api_config_repository.go b/internal/repository/api_config_repository.go index 3e4993e..347c160 100644 --- a/internal/repository/api_config_repository.go +++ b/internal/repository/api_config_repository.go @@ -10,17 +10,17 @@ import ( "github.com/truewhile/MeBox/internal/model" ) -// ApiConfigRepository persists model.ApiConfig records. +// 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 { +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 +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 @@ -32,27 +32,40 @@ func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider strin } // List returns all API configs. -func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) { - var rows []model.ApiConfig +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 +// 显式 map 更新:Assign(struct) 会跳过零值字段,导致 Enabled=false、 +// 清空 BaseURL/Extra 等撤销操作静默失效。 +func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.APIConfig) error { + return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var existing model.APIConfig + err := tx.Where("provider = ?", c.Provider).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return tx.Create(c).Error + } + if err != nil { + return err + } + c.ID = existing.ID + c.CreatedAt = existing.CreatedAt + return tx.Model(&existing).Updates(map[string]any{ + "api_key": c.APIKey, + "base_url": c.BaseURL, + "extra": c.Extra, + "enabled": c.Enabled, + "updated_at": time.Now(), + }).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{}). +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, @@ -64,13 +77,13 @@ func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) er // Delete 物理删除 API 配置。 func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error { - return r.db.WithContext(ctx).Unscoped().Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error + return r.db.WithContext(ctx).Unscoped().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{}). + return r.db.WithContext(ctx).Model(&model.APIConfig{}). Where("provider = ?", provider).Updates(map[string]any{ "test_result": result, "last_tested_at": &now, diff --git a/internal/repository/media_repository_upsert.go b/internal/repository/media_repository_upsert.go index df4ee41..e2c8ab5 100644 --- a/internal/repository/media_repository_upsert.go +++ b/internal/repository/media_repository_upsert.go @@ -22,45 +22,93 @@ import ( // 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending') // 永远捞不到数据。 func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { - return withSQLiteBusyRetry(ctx, func() error { - return r.upsertWithDB(ctx, r.db, m) + var indexIDs []string + err := withSQLiteBusyRetry(ctx, func() error { + id, uerr := r.upsertWithDB(ctx, r.db, m) + if uerr != nil { + return uerr + } + indexIDs = append(indexIDs[:0], id) + return nil }) + if err != nil { + return err + } + r.indexByIDBestEffort(ctx, indexIDs) + return nil } // UpsertBatch 在单个事务里逐条执行 Upsert:扫描一批只提交(fsync)一次, // 而不是每条一个隐式事务。任一条目落库失败不影响批内已成功的条目—— // 事务回滚后由调用方退回逐条 Upsert 兜底。 +// +// OpenSearch 索引同步(HTTP,4s 超时)必须在事务提交之后统一执行:放在 +// 事务内会把 SQLite 写锁挂起在网络 IO 上,且批内用非事务连接回读只能 +// 拿到提交前的旧版本数据,把陈旧内容写进索引。 func (r *MediaRepository) UpsertBatch(ctx context.Context, items []*model.Media) error { if len(items) == 0 { return nil } - return withSQLiteBusyRetry(ctx, func() error { + indexIDs := make([]string, 0, len(items)) + err := withSQLiteBusyRetry(ctx, func() error { + indexIDs = indexIDs[:0] return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { for _, m := range items { if m == nil { continue } - if err := r.upsertWithDB(ctx, tx, m); err != nil { + id, err := r.upsertWithDB(ctx, tx, m) + if err != nil { return err } + if id != "" { + indexIDs = append(indexIDs, id) + } } return nil }) }) -} - -func (r *MediaRepository) upsertWithDB(ctx context.Context, db *gorm.DB, m *model.Media) error { - existing, created, err := r.findOrCreateMediaByPath(ctx, db, m) if err != nil { return err } + r.indexByIDBestEffort(ctx, indexIDs) + return nil +} + +// indexByIDBestEffort 在事务提交后按 ID 回读最新行并同步搜索索引。 +func (r *MediaRepository) indexByIDBestEffort(ctx context.Context, ids []string) { + for _, id := range ids { + if id == "" { + continue + } + if fresh, err := r.FindByID(ctx, id); err == nil && fresh != nil { + r.indexMediaBestEffort(ctx, *fresh) + } + } +} + +// upsertWithDB 落库(新建或更新),返回需要重建索引的媒体 ID(无则空串)。 +func (r *MediaRepository) upsertWithDB(ctx context.Context, db *gorm.DB, m *model.Media) (string, error) { + existing, created, err := r.findOrCreateMediaByPath(ctx, db, m) + if err != nil { + return "", err + } if created { - r.indexMediaBestEffort(ctx, *m) - return nil + return m.ID, nil } updates := mediaUpsertUpdates(existing, *m) - return r.applyMediaUpsertUpdates(ctx, db, m, existing, updates) + if len(updates) == 0 { + *m = existing + return "", nil + } + if err := db.WithContext(ctx).Unscoped().Model(&model.Media{}). + Where("id = ?", existing.ID).Updates(updates).Error; err != nil { + return "", err + } + // 回写 ID / 不可变字段,让 caller 拿到完整的现有行。 + *m = existing + return existing.ID, nil } func (r *MediaRepository) findOrCreateMediaByPath(ctx context.Context, db *gorm.DB, m *model.Media) (model.Media, bool, error) { @@ -75,6 +123,9 @@ func (r *MediaRepository) findOrCreateMediaByPath(ctx context.Context, db *gorm. return *m, true, nil } else if retryErr := db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil { return model.Media{}, false, createErr + } else { + // 并发插入竞态:重查已命中既有行,直接走更新分支。 + return existing, false, nil } } if err != nil { @@ -259,24 +310,6 @@ func setNonEmptyMediaString(updates map[string]any, key, current, next string) { } } -func (r *MediaRepository) applyMediaUpsertUpdates(ctx context.Context, db *gorm.DB, m *model.Media, existing model.Media, updates map[string]any) error { - if len(updates) == 0 { - *m = existing - return nil - } - if err := 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 { - *m = *fresh - r.indexMediaBestEffort(ctx, *fresh) - } - return nil -} - func setIfChanged[T comparable](updates map[string]any, key string, current, next T) { if current != next { updates[key] = next diff --git a/internal/repository/media_search_repository.go b/internal/repository/media_search_repository.go index 85ac738..d7c94d6 100644 --- a/internal/repository/media_search_repository.go +++ b/internal/repository/media_search_repository.go @@ -105,11 +105,17 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string, var total int64 q := r.db.WithContext(ctx).Model(&model.Media{}) q = applyMediaQueryFilter(q, filter) + // SQLite 的 LIKE 对 ASCII 不区分大小写;Postgres 的 LIKE 区分大小写, + // 需用 ILIKE 保持两端搜索行为一致。 + likeOp := "LIKE" + if r.db.Dialector != nil && r.db.Dialector.Name() == "postgres" { + likeOp = "ILIKE" + } 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 '\\')", + "(title "+likeOp+" ? ESCAPE '\\' OR original_name "+likeOp+" ? ESCAPE '\\' OR path "+likeOp+" ? ESCAPE '\\' OR genres "+likeOp+" ? ESCAPE '\\')", like, like, like, like, ) } @@ -120,7 +126,7 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string, 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", + "CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title "+likeOp+" ? ESCAPE '\\' THEN 2 WHEN original_name "+likeOp+" ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc", exact, exact, prefix, prefix, )) } else { @@ -259,7 +265,9 @@ func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool { } r.searchIndexOnce.Do(func() { var count int64 - err := r.db.WithContext(ctx). + // 用 Background 探测:sync.Once 只执行一次,若借用调用方的 + // ctx 且恰好被取消,FTS 会被永久误判为不可用。 + err := r.db.WithContext(context.Background()). Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`). Scan(&count).Error r.searchIndexAvailable = err == nil && count > 0 diff --git a/internal/repository/permission_repository.go b/internal/repository/permission_repository.go index 805ddfd..13fa5c2 100644 --- a/internal/repository/permission_repository.go +++ b/internal/repository/permission_repository.go @@ -3,6 +3,7 @@ package repository import ( "context" "errors" + "time" "gorm.io/gorm" @@ -44,10 +45,43 @@ func (r *PermissionRepository) Update(ctx context.Context, userID string, update } // Upsert creates or updates a permission record. +// 显式 map 更新:Assign(struct) 会被 GORM 跳过零值字段,导致权限 +// "撤销"(false)保存后静默失效且无法重置。 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 + return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var existing model.UserPermission + err := tx.Where("user_id = ?", p.UserID).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return tx.Create(p).Error + } + if err != nil { + return err + } + p.ID = existing.ID + p.CreatedAt = existing.CreatedAt + return tx.Model(&existing).Updates(map[string]any{ + "can_view_dashboard": p.CanViewDashboard, + "can_play_media": p.CanPlayMedia, + "can_cast": p.CanCast, + "can_external_player": p.CanExternalPlayer, + "can_favorite": p.CanFavorite, + "can_view_history": p.CanViewHistory, + "can_edit_media": p.CanEditMedia, + "can_rescrape": p.CanRescrape, + "can_use_ai": p.CanUseAI, + "can_capture_frames": p.CanCaptureFrames, + "can_manage_downloads": p.CanManageDownloads, + "can_manage_subscriptions": p.CanManageSubscriptions, + "can_manage_sites": p.CanManageSites, + "can_use_ai_assistant": p.CanUseAIAssistant, + "can_manage_users": p.CanManageUsers, + "can_manage_files": p.CanManageFiles, + "can_manage_strm": p.CanManageStrm, + "can_access_settings": p.CanAccessSettings, + "updated_at": time.Now(), + }).Error + }) }) } diff --git a/internal/repository/scrape_task_repository.go b/internal/repository/scrape_task_repository.go index 4140625..44fc4f8 100644 --- a/internal/repository/scrape_task_repository.go +++ b/internal/repository/scrape_task_repository.go @@ -52,6 +52,40 @@ func (r *ScrapeTaskRepository) FindActiveByMediaID(ctx context.Context, mediaID return &t, err } +// FindActiveByMediaIDs 批量查询仍处于 pending/running 的任务媒体 ID 集合, +// 供整库入队时去重(防止同一媒体被重复入队并被并发双刮)。 +func (r *ScrapeTaskRepository) FindActiveByMediaIDs(ctx context.Context, mediaIDs []string) (map[string]bool, error) { + out := make(map[string]bool, len(mediaIDs)) + if len(mediaIDs) == 0 { + return out, nil + } + var rows []model.ScrapeTask + err := r.db.WithContext(ctx). + Select("media_id"). + Where("media_id IN ? AND status IN ?", mediaIDs, []string{model.ScrapeTaskPending, model.ScrapeTaskRunning}). + Find(&rows).Error + if err != nil { + return nil, err + } + for _, r := range rows { + out[r.MediaID] = true + } + return out, nil +} + +// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务重置为 pending, +// 否则任务永久卡死(ClaimPending 只认 pending,重试按钮也拒绝 running)。 +func (r *ScrapeTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) { + res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}). + Where("status = ?", model.ScrapeTaskRunning). + Updates(map[string]any{ + "status": model.ScrapeTaskPending, + "error": "服务重启,任务已重置", + "started_at": nil, + }) + return res.RowsAffected, res.Error +} + func (r *ScrapeTaskRepository) List(ctx context.Context, status string, page, pageSize int) ([]model.ScrapeTask, int64, error) { if page < 1 { page = 1 diff --git a/internal/repository/strm_repository.go b/internal/repository/strm_repository.go index bbf411e..3b4117b 100644 --- a/internal/repository/strm_repository.go +++ b/internal/repository/strm_repository.go @@ -302,6 +302,39 @@ func (r *StrmDownloadTaskRepository) Update(ctx context.Context, t *model.StrmDo }) } +// UpdateIfRunning 仅当任务在 DB 中仍为 running 时写入给定字段。 +// 返回 false 表示任务已被外部改变状态(如用户取消),收尾不得覆盖。 +func (r *StrmDownloadTaskRepository) UpdateIfRunning(ctx context.Context, id string, updates map[string]any) (bool, error) { + var ok bool + err := withSQLiteBusyRetry(ctx, func() error { + updates["updated_at"] = time.Now() + res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). + Where("id = ? AND status = ?", id, model.StrmTaskRunning).Updates(updates) + ok = res.RowsAffected > 0 + return res.Error + }) + return ok, err +} + +// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务全部重置为 +// pending(清空退避时间以便立即可被认领),否则任务永久卡死且会阻塞 +// 该文件的重复下载。 +func (r *StrmDownloadTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) { + var n int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). + Where("status = ?", model.StrmTaskRunning). + Updates(map[string]any{ + "status": model.StrmTaskPending, + "error": "服务重启,任务已重置", + "started_at": nil, + }) + n = res.RowsAffected + return res.Error + }) + return n, err +} + func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) error { return withSQLiteBusyRetry(ctx, func() error { return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error @@ -620,6 +653,36 @@ func (r *StrmUploadTaskRepository) Update(ctx context.Context, t *model.StrmUplo }) } +// UpdateIfRunning 仅当任务在 DB 中仍为 running 时写入给定字段。 +func (r *StrmUploadTaskRepository) UpdateIfRunning(ctx context.Context, id string, updates map[string]any) (bool, error) { + var ok bool + err := withSQLiteBusyRetry(ctx, func() error { + updates["updated_at"] = time.Now() + res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}). + Where("id = ? AND status = ?", id, model.StrmTaskRunning).Updates(updates) + ok = res.RowsAffected > 0 + return res.Error + }) + return ok, err +} + +// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务全部重置为 pending。 +func (r *StrmUploadTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) { + var n int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}). + Where("status = ?", model.StrmTaskRunning). + Updates(map[string]any{ + "status": model.StrmTaskPending, + "error": "服务重启,任务已重置", + "started_at": nil, + }) + n = res.RowsAffected + return res.Error + }) + return n, err +} + func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) error { return withSQLiteBusyRetry(ctx, func() error { return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error diff --git a/internal/service/adult_scraper_routing_test.go b/internal/service/adult_scraper_routing_test.go index 1d71b6d..6b3caee 100644 --- a/internal/service/adult_scraper_routing_test.go +++ b/internal/service/adult_scraper_routing_test.go @@ -20,7 +20,7 @@ func TestAdultProviderRouting(t *testing.T) { if err != nil { t.Fatalf("failed to open sqlite: %v", err) } - _ = db.AutoMigrate(&model.Setting{}, &model.ApiConfig{}) + _ = db.AutoMigrate(&model.Setting{}, &model.APIConfig{}) repos := repository.New(db) diff --git a/internal/service/api_config_connection.go b/internal/service/api_config_connection.go index c983bc0..feb8b53 100644 --- a/internal/service/api_config_connection.go +++ b/internal/service/api_config_connection.go @@ -40,7 +40,7 @@ func (s *ApiConfigService) TestConnection(ctx context.Context, provider string) } // testTMDb 测试 TMDb API 连接。 -func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testTMDb(cfg *model.APIConfig) (string, error) { if cfg.APIKey == "" { return "error", errors.New("API key is required") } @@ -74,7 +74,7 @@ func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) { } // testOpenAI 测试 OpenAI API 连接。 -func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testOpenAI(cfg *model.APIConfig) (string, error) { if cfg.APIKey == "" { return "error", errors.New("API key is required") } @@ -108,7 +108,7 @@ func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) { } // testDeepSeek 测试 DeepSeek API 连接。 -func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testDeepSeek(cfg *model.APIConfig) (string, error) { if cfg.APIKey == "" { return "error", errors.New("API key is required") } @@ -142,7 +142,7 @@ func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) { } // testSiliconFlow 测试 SiliconFlow API 连接。 -func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testSiliconFlow(cfg *model.APIConfig) (string, error) { if cfg.APIKey == "" { return "error", errors.New("API key is required") } @@ -176,7 +176,7 @@ func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error) } // testAdult 测试 Adult (JavDB/JavBus) 刮削数据源连接与年龄验证。 -func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.APIConfig) (string, error) { bases := []string{} if cfg.BaseURL != "" { bases = append(bases, adultConfiguredBases(cfg.BaseURL)...) @@ -245,7 +245,7 @@ func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.ApiConfig) } // testMetaTube 测试 MetaTube Server 连接与 Token。 -func (s *ApiConfigService) testMetaTube(ctx context.Context, cfg *model.ApiConfig) (string, error) { +func (s *ApiConfigService) testMetaTube(ctx context.Context, cfg *model.APIConfig) (string, error) { serverURL := strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/") if serverURL == "" { serverURL = "http://127.0.0.1:7700" diff --git a/internal/service/api_config_helpers.go b/internal/service/api_config_helpers.go index be0c36e..b0e59ce 100644 --- a/internal/service/api_config_helpers.go +++ b/internal/service/api_config_helpers.go @@ -9,7 +9,7 @@ import ( ) // GetEffectiveConfig 获取生效的 API 配置(数据库配置优先于配置文件)。 -func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider string) (*model.ApiConfig, error) { +func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider string) (*model.APIConfig, error) { // 首先尝试从数据库获取 cfg, err := s.GetByProvider(ctx, provider) if err == nil && cfg != nil { @@ -21,7 +21,7 @@ func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider stri } // getConfigFromFile 从配置文件获取 API 配置。 -func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig, error) { +func (s *ApiConfigService) getConfigFromFile(provider string) (*model.APIConfig, error) { var apiKey string var hasKey bool @@ -44,7 +44,7 @@ func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig, return nil, ErrApiConfigNotFound } - return &model.ApiConfig{ + return &model.APIConfig{ Provider: provider, APIKey: apiKey, Enabled: true, diff --git a/internal/service/api_config_svc.go b/internal/service/api_config_svc.go index 893217a..e86e086 100644 --- a/internal/service/api_config_svc.go +++ b/internal/service/api_config_svc.go @@ -33,7 +33,7 @@ var ( ) // GetByProvider 获取指定提供者的 API 配置。 -func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) { +func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (*model.APIConfig, error) { cfg, err := s.repo.ApiConfig.FindByProvider(ctx, provider) if err != nil { return nil, err @@ -49,7 +49,7 @@ func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) ( } // List 返回所有 API 配置。 -func (s *ApiConfigService) List(ctx context.Context) ([]model.ApiConfig, error) { +func (s *ApiConfigService) List(ctx context.Context) ([]model.APIConfig, error) { configs, err := s.repo.ApiConfig.List(ctx) if err != nil { return nil, err @@ -69,7 +69,7 @@ func (s *ApiConfigService) GetProviders() []model.ApiProvider { } // Upsert 创建或更新 API 配置,自动加密敏感字段。 -func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) (*model.ApiConfig, error) { +func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) (*model.APIConfig, error) { // 验证提供者是否有效 if !s.isValidProvider(provider) { return nil, ErrInvalidProvider @@ -81,7 +81,7 @@ func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, encryptedKey = s.crypto.Encrypt(apiKey) } - cfg := &model.ApiConfig{ + cfg := &model.APIConfig{ Provider: provider, APIKey: encryptedKey, BaseURL: baseURL, @@ -112,7 +112,7 @@ func (s *ApiConfigService) Update(ctx context.Context, provider string, apiKey, encryptedKey = s.crypto.Encrypt(apiKey) } - cfg := &model.ApiConfig{ + cfg := &model.APIConfig{ Provider: provider, APIKey: encryptedKey, BaseURL: baseURL, diff --git a/internal/service/cloud/clouddrive2.go b/internal/service/cloud/clouddrive2.go index f92d73c..62afb61 100644 --- a/internal/service/cloud/clouddrive2.go +++ b/internal/service/cloud/clouddrive2.go @@ -8,6 +8,8 @@ import ( "net/url" "path" "strings" + "sync" + "time" ) // cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint. @@ -22,11 +24,18 @@ type cloudDrive2Provider struct { base *url.URL username string password string - token string + token string // 配置的静态令牌(构造后只读) ua string apiBase *url.URL client *http.Client proxy bool + + // tokenMu / loginToken / loginTokenSeen 保护 OpenList 用户名密码登录的 + // token 缓存:多 worker 并发时单飞登录,缓存有效期内直接复用, + // 401 时清缓存重登(见 clouddrive2_openlist.go 的 openListAPIToken)。 + tokenMu sync.Mutex + loginToken string + loginTokenSeen time.Time } func newCloudDrive2(cfg map[string]any, client *http.Client) *cloudDrive2Provider { diff --git a/internal/service/cloud/clouddrive2_dav.go b/internal/service/cloud/clouddrive2_dav.go index 042946c..3ce16c2 100644 --- a/internal/service/cloud/clouddrive2_dav.go +++ b/internal/service/cloud/clouddrive2_dav.go @@ -36,9 +36,10 @@ func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, p.decorateDAVStatusError(resp, target) } - body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20)) + // 流式解码:超大目录(如上万条目的网盘目录)响应可能远超旧 4MB 截断上限, + // 直接 xml.Unmarshal 会截断报错;这里用 LimitReader(64MB) + Decoder 边读边解 var multi cloudDAVMultiStatus - if err := xml.Unmarshal(body, &multi); err != nil { + if err := xml.NewDecoder(io.LimitReader(resp.Body, 64<<20)).Decode(&multi); err != nil { return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err) } basePath := strings.TrimRight(p.base.EscapedPath(), "/") diff --git a/internal/service/cloud/clouddrive2_mutation.go b/internal/service/cloud/clouddrive2_mutation.go index 31ae3d1..b5eb77c 100644 --- a/internal/service/cloud/clouddrive2_mutation.go +++ b/internal/service/cloud/clouddrive2_mutation.go @@ -131,10 +131,13 @@ func (p *cloudDrive2Provider) openListAPIMove(ctx context.Context, source, targe } func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath string, payload any, action string) error { - token, err := p.openListAPIToken(ctx) - if err != nil { - return err - } + _, err := doWithOpenListAPIToken(ctx, p, func(token string) (struct{}, error) { + return struct{}{}, p.openListAPIPostWithToken(ctx, apiPath, payload, action, token) + }) + return err +} + +func (p *cloudDrive2Provider) openListAPIPostWithToken(ctx context.Context, apiPath string, payload any, action, token string) error { body, _ := json.Marshal(payload) req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body)) if err != nil { @@ -151,6 +154,9 @@ func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath strin return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err) } defer resp.Body.Close() + if resp.StatusCode == http.StatusUnauthorized { + return errOpenListAPITokenExpired + } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return fmt.Errorf("%s: api %s returned http %d", p.name, action, resp.StatusCode) } @@ -200,27 +206,43 @@ func (p *cloudDrive2Provider) PutFile(ctx context.Context, remotePath string, r } // openListAPIPutFile 通过 OpenList /api/fs/form 上传(QMediaSync 同款契约: -// PUT + multipart + File-Path 头)。 +// PUT + multipart + File-Path 头)。使用 io.Pipe + multipart.Writer 边写边发, +// 避免把整个文件读进内存。 func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath string, r io.Reader) error { token, err := p.openListAPIToken(ctx) if err != nil { return err } encodedPath := openListPathEscape(remotePath) - body := &bytes.Buffer{} - writer := multipart.NewWriter(body) - formFile, err := writer.CreateFormFile("file", path.Base(remotePath)) - if err != nil { - return err - } - if _, err := io.Copy(formFile, r); err != nil { - return err - } - if err := writer.Close(); err != nil { - return err - } - req, err := http.NewRequestWithContext(ctx, http.MethodPut, p.openListAPIURL("/api/fs/form"), body) + + pr, pw := io.Pipe() + writer := multipart.NewWriter(pw) + go func() { + var writeErr error + defer func() { + // 读源失败必须传给 pipe 写端,让 HTTP 请求以失败收场而不是静默截断 + if writeErr != nil { + _ = pw.CloseWithError(writeErr) + return + } + _ = pw.Close() + }() + formFile, err := writer.CreateFormFile("file", path.Base(remotePath)) + if err != nil { + writeErr = err + return + } + if _, err := io.Copy(formFile, r); err != nil { + writeErr = err + return + } + writeErr = writer.Close() + }() + + req, err := http.NewRequestWithContext(ctx, http.MethodPut, p.openListAPIURL("/api/fs/form"), pr) if err != nil { + // 关闭读端以释放仍在等待写入的后台 goroutine(其 Write 会立即失败返回) + _ = pr.Close() return err } req.Header.Set("Authorization", token) @@ -230,9 +252,15 @@ func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath req.Header.Set("Overwrite", "true") resp, err := p.client.Do(req) if err != nil { + // 传输层失败(含提前断开)时 net/http 会关闭请求 body,解除后台 goroutine 阻塞 return decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/form"), err) } defer resp.Body.Close() + if resp.StatusCode == http.StatusUnauthorized { + // 流式 body 无法重放,不能自动重试:清除登录 token 缓存让下次上传重新登录, + // 本次返回明确错误交由调用方重试 + p.invalidateOpenListAPIToken() + } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return p.openListAPIStatusError("upload", remotePath, resp.StatusCode) } diff --git a/internal/service/cloud/clouddrive2_openlist.go b/internal/service/cloud/clouddrive2_openlist.go index 6c4e2f4..ed38075 100644 --- a/internal/service/cloud/clouddrive2_openlist.go +++ b/internal/service/cloud/clouddrive2_openlist.go @@ -4,18 +4,50 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "net/http" "net/url" "strings" + "time" ) -func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) { +// errOpenListAPITokenExpired 标记 OpenList 返回 401(登录 token 已失效): +// 调用方收到后应清缓存重登一次再重试原请求。 +var errOpenListAPITokenExpired = errors.New("openlist api token expired") + +// openListAPITokenCacheTTL 登录 token 缓存有效期(OpenList 默认签发 48h JWT, +// 这里保守取 30 分钟,过期自动重新登录)。 +const openListAPITokenCacheTTL = 30 * time.Minute + +// doWithOpenListAPIToken 获取 OpenList API token 后执行 fn;若请求命中 401 +// (登录 token 失效)则清缓存重登一次并重试,避免一次 token 轮换导致整批请求失败。 +func doWithOpenListAPIToken[T any](ctx context.Context, p *cloudDrive2Provider, fn func(token string) (T, error)) (T, error) { + var zero T token, err := p.openListAPIToken(ctx) if err != nil { - return nil, err + return zero, err } + result, err := fn(token) + if err == nil || !errors.Is(err, errOpenListAPITokenExpired) { + return result, err + } + p.invalidateOpenListAPIToken() + token, err = p.openListAPIToken(ctx) + if err != nil { + return zero, err + } + return fn(token) +} + +func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) { + return doWithOpenListAPIToken(ctx, p, func(token string) ([]FileEntry, error) { + return p.listOpenListAPIWithToken(ctx, dir, token) + }) +} + +func (p *cloudDrive2Provider) listOpenListAPIWithToken(ctx context.Context, dir, token string) ([]FileEntry, error) { const pageSize = 500 target := normalizeCloudDAVPath(dir) out := make([]FileEntry, 0, pageSize) @@ -45,6 +77,9 @@ func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ( var decoded openListListResponse decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded) resp.Body.Close() + if resp.StatusCode == http.StatusUnauthorized { + return nil, errOpenListAPITokenExpired + } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, p.openListAPIStatusError("list", target, resp.StatusCode) } @@ -85,10 +120,12 @@ func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ( } func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) { - token, err := p.openListAPIToken(ctx) - if err != nil { - return nil, err - } + return doWithOpenListAPIToken(ctx, p, func(token string) (*DirectLink, error) { + return p.resolveOpenListAPIDirectWithToken(ctx, fileRef, token) + }) +} + +func (p *cloudDrive2Provider) resolveOpenListAPIDirectWithToken(ctx context.Context, fileRef, token string) (*DirectLink, error) { 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 { @@ -105,6 +142,9 @@ func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, file return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err) } defer resp.Body.Close() + if resp.StatusCode == http.StatusUnauthorized { + return nil, errOpenListAPITokenExpired + } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, p.openListAPIStatusError("get", fileRef, resp.StatusCode) } @@ -163,6 +203,10 @@ func (p *cloudDrive2Provider) hasOpenListAPICredentials() bool { return strings.TrimSpace(p.token) != "" || (strings.TrimSpace(p.username) != "" && p.password != "") } +// openListAPIToken 返回 OpenList API 访问令牌: +// - 配置了静态 token 时直接使用(构造后只读,无并发问题); +// - 否则用用户名密码登录,并在缓存有效期内单飞复用——8 个同步 worker 并发时 +// 只会有一个 goroutine 真正执行登录,避免登录风暴;登录 token 的写入受 tokenMu 保护。 func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) { if token := strings.TrimSpace(p.token); token != "" { return token, nil @@ -170,6 +214,30 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err if strings.TrimSpace(p.username) == "" || p.password == "" { return "", nil } + p.tokenMu.Lock() + defer p.tokenMu.Unlock() + if p.loginToken != "" && time.Since(p.loginTokenSeen) < openListAPITokenCacheTTL { + return p.loginToken, nil + } + token, err := p.openListAPILogin(ctx) + if err != nil { + return "", err + } + p.loginToken = token + p.loginTokenSeen = time.Now() + return token, nil +} + +// invalidateOpenListAPIToken 清除登录 token 缓存(收到 401 时调用,下次请求重新登录)。 +func (p *cloudDrive2Provider) invalidateOpenListAPIToken() { + p.tokenMu.Lock() + p.loginToken = "" + p.loginTokenSeen = time.Time{} + p.tokenMu.Unlock() +} + +// openListAPILogin 调用 OpenList /api/auth/login 换取登录 token。 +func (p *cloudDrive2Provider) openListAPILogin(ctx context.Context) (string, error) { payload, _ := json.Marshal(map[string]string{ "username": p.username, "password": p.password, @@ -204,7 +272,6 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err if token == "" { return "", fmt.Errorf("%s: api login returned empty token", p.name) } - p.token = token return token, nil } diff --git a/internal/service/cloud/pan115_openapi.go b/internal/service/cloud/pan115_openapi.go index e81d169..a404b34 100644 --- a/internal/service/cloud/pan115_openapi.go +++ b/internal/service/cloud/pan115_openapi.go @@ -48,7 +48,7 @@ func (p *openAPI115Provider) Ping(ctx context.Context) error { if strings.TrimSpace(p.c.AppID) == "" { return fmt.Errorf("115: 缺少开放平台应用 ID,请重新授权") } - if strings.TrimSpace(p.c.AccessToken) == "" { + if strings.TrimSpace(p.c.CurrentAccessToken()) == "" { return fmt.Errorf("115: 缺少访问令牌,请重新授权") } _, _, err := p.c.GetFsList(ctx, "0", 0, 1) @@ -70,7 +70,7 @@ func (p *openAPI115Provider) List(ctx context.Context, dirID string) ([]FileEntr Name: f.FileName, IsDir: f.Category == cloud115.TypeDir, Size: f.FileSize, - MTime: f.Utime, + MTime: f.ModifiedAt(), PickCode: f.PickCode, }) } @@ -127,12 +127,15 @@ func (p *openAPI115Provider) PutFileNamed(ctx context.Context, parentCID, fileNa if err := tmp.Close(); err != nil { return fmt.Errorf("115: 关闭临时文件失败:%w", err) } - // 重命名为目标文件名,保证上传到 115 后保留原始文件名 + // 重命名为目标文件名,保证上传到 115 后保留原始文件名。 + // 重命名失败必须 fail fast:静默用随机临时名上传会导致 115 上的文件名 + // 变成 mebox-upload-xxx,破坏元数据文件名契约。 if fileName != "" && fileName != filepath.Base(tmpPath) { namedPath := filepath.Join(filepath.Dir(tmpPath), fileName) - if err := os.Rename(tmpPath, namedPath); err == nil { - tmpPath = namedPath + if err := os.Rename(tmpPath, namedPath); err != nil { + return fmt.Errorf("115: 重命名临时文件为 %s 失败:%w", fileName, err) } + tmpPath = namedPath } _, err = p.c.Upload(ctx, tmpPath, parentCID, "", "") if err != nil { diff --git a/internal/service/cloud115/client.go b/internal/service/cloud115/client.go index a3d915d..7abd788 100644 --- a/internal/service/cloud115/client.go +++ b/internal/service/cloud115/client.go @@ -23,9 +23,14 @@ type OpenClient struct { RefreshTokenStr string executor *QueueExecutor - // tokenMu 保护令牌刷新:业务请求中途 access_token 失效时自动刷新重试, - // 多 goroutine(同步列表 + 下载队列)并发下只允许一次刷新进行。 - tokenMu sync.Mutex + // OnTokenRefreshed 在 access_token 刷新成功后回调(参数为新令牌对), + // 供上层持久化新令牌使用;nil 安全,且在 tokenMu 释放后调用以避免死锁。 + OnTokenRefreshed func(accessToken, refreshToken string) + + // tokenMu 保护 AccessToken / RefreshTokenStr 的并发读写:业务请求中途 + // access_token 失效时自动刷新重试,多 goroutine(同步列表 + 下载队列) + // 并发下只允许一次刷新进行。 + tokenMu sync.RWMutex } // default115HTTPClient 创建带有防 405 重定向保护的 http.Client。 @@ -57,12 +62,40 @@ func NewOpenClient(appID, accessToken, refreshToken string) *OpenClient { } } -// SetAuthToken 更新认证令牌。 +// SetAuthToken 更新认证令牌(并发安全)。 func (c *OpenClient) SetAuthToken(accessToken, refreshToken string) { + c.tokenMu.Lock() + c.setAuthTokenLocked(accessToken, refreshToken) + c.tokenMu.Unlock() +} + +// setAuthTokenLocked 无锁更新令牌,调用方必须已持有 tokenMu 写锁 +// (tryRefreshTokenLocked 等已持锁流程内部使用,避免重入死锁)。 +func (c *OpenClient) setAuthTokenLocked(accessToken, refreshToken string) { c.AccessToken = accessToken c.RefreshTokenStr = refreshToken } +// currentAccessToken 返回当前 access_token(并发安全)。 +func (c *OpenClient) currentAccessToken() string { + c.tokenMu.RLock() + defer c.tokenMu.RUnlock() + return c.AccessToken +} + +// currentRefreshToken 返回当前 refresh_token(并发安全)。 +func (c *OpenClient) currentRefreshToken() string { + c.tokenMu.RLock() + defer c.tokenMu.RUnlock() + return c.RefreshTokenStr +} + +// CurrentAccessToken 返回当前 access_token 快照(并发安全), +// 供上层在无锁环境下安全读取(如 Ping 时探测令牌是否存在)。 +func (c *OpenClient) CurrentAccessToken() string { + return c.currentAccessToken() +} + // RespState 兼容 115 不同端点返回的 state 类型(proapi 返回布尔、passport 返回数字)。 type RespState bool @@ -191,6 +224,14 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map if access { // 刷新失败(或已刷新仍失败)时返回明确错误 lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115: access_token 校验失败且刷新未成功") + } else { + // 未携带令牌的请求(登录/刷新流程)命中 token 类错误码: + // 必须返回显式 error,避免调用方把 (resp, nil) 当作成功处理 + msg := base.Message + if msg == "" { + msg = base.Error + } + lastErr = fmt.Errorf("115: 认证失败(code=%d): %s", base.Code, msg) } return &base, lastErr } @@ -242,8 +283,11 @@ func (c *OpenClient) buildRequestWithUA(ctx context.Context, method, rawURL stri if method == http.MethodPost && len(form) > 0 { req.Header.Set("Content-Type", "application/x-www-form-urlencoded") } - if access && c.AccessToken != "" { - req.Header.Set("Authorization", "Bearer "+c.AccessToken) + if access { + // RLock 读取令牌,避免与刷新流程的写入产生数据竞争 + if accessToken := c.currentAccessToken(); accessToken != "" { + req.Header.Set("Authorization", "Bearer "+accessToken) + } } return req, nil } @@ -261,33 +305,57 @@ func (c *OpenClient) doAuthJSONWithUA(ctx context.Context, method, rawURL string // tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方 // 应使用内存中的新 token 重试原请求)。 // +// 拿到写锁后在锁内读取 oldAccess,与持锁期间的当前值对比:若已被其他 +// goroutine 刷新过则直接复用新 token,避免并发请求连环轮转消耗 115 的 +// 一次性 refresh_token。全程持写锁读写 token 字段,无 TOCTOU 窗口。 +// // 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120) // 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。 // 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃, // 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。 func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool { c.tokenMu.Lock() - defer c.tokenMu.Unlock() + // 在已持有写锁内读取当前 token 作为"刷新前快照",消除双重加锁窗口: + // 若在拿锁期间已有其他 goroutine 完成刷新,refreshTokenWhileLocked + // 内的 c.AccessToken != oldAccess 判断会立即命中并返回复用。 + oldAccess := c.AccessToken + newToken, ok := c.refreshTokenWhileLocked(ctx, oldAccess) + c.tokenMu.Unlock() + // 回调必须在 tokenMu 释放后调用,避免上层在回调内访问客户端时死锁 + if ok && newToken != nil && c.OnTokenRefreshed != nil { + c.OnTokenRefreshed(newToken.AccessToken, newToken.RefreshToken) + } + return ok +} + +// refreshTokenWhileLocked 在已持有 tokenMu 写锁的前提下执行刷新。 +// 返回 (新令牌, 是否成功);命中"他人已刷新"捷径时新令牌为 nil。 +func (c *OpenClient) refreshTokenWhileLocked(ctx context.Context, oldAccess string) (*TokenData, bool) { + if c.AccessToken != oldAccess { + // 其他 goroutine 刚刷新过:直接复用内存中的新 token 重试原请求 + return nil, true + } + refreshToken := c.RefreshTokenStr for attempt := 0; attempt < refreshAttempts; attempt++ { - token, err := c.RefreshToken(c.RefreshTokenStr) + token, err := c.doRefreshToken(refreshToken) if err == nil { - c.SetAuthToken(token.AccessToken, token.RefreshToken) - return true + c.setAuthTokenLocked(token.AccessToken, token.RefreshToken) + return token, true } if IsRefreshTokenDead(err) { - c.SetAuthToken("", "") - return false + c.setAuthTokenLocked("", "") + return nil, false } // 可恢复失败:退避后重试。ctx 取消时立即放弃。 if attempt < refreshAttempts-1 { select { case <-ctx.Done(): - return false + return nil, false case <-time.After(refreshBackoff(attempt)): } } } - return false + return nil, false } // refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。 @@ -312,7 +380,14 @@ func isTokenCode(code int) bool { } // openList 解析 data 为对象或数组(StructOrArray 语义)。 +// 115 部分接口在鉴权/业务异常时会返回 data:null 或 data:{},此时若直接 +// 反序列化会得到零值元素 + nil error,调用方会把空数据当成功处理; +// 这里对 null/空对象显式报错。 func openList[T any](raw json.RawMessage) ([]T, error) { + trimmed := bytes.TrimSpace(raw) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) || bytes.Equal(trimmed, []byte("{}")) { + return nil, fmt.Errorf("115: data 为空(%s)", string(trimmed)) + } var single T if err := json.Unmarshal(raw, &single); err == nil { return []T{single}, nil @@ -324,12 +399,16 @@ func openList[T any](raw json.RawMessage) ([]T, error) { return nil, fmt.Errorf("115: data 既不是对象也不是数组") } -// openFirstList 取 data 的第一个元素。 +// openFirstList 取 data 的第一个元素;data 为空(null/空数组)时返回显式错误, +// 避免调用方拿到 (nil, nil) 后解引用空指针。 func openFirstList[T any](raw json.RawMessage) (*T, error) { items, err := openList[T](raw) - if err != nil || len(items) == 0 { + if err != nil { return nil, err } + if len(items) == 0 { + return nil, fmt.Errorf("115: data 为空数组") + } return &items[0], nil } diff --git a/internal/service/cloud115/oauth.go b/internal/service/cloud115/oauth.go index 3ec9743..39e5f13 100644 --- a/internal/service/cloud115/oauth.go +++ b/internal/service/cloud115/oauth.go @@ -317,13 +317,18 @@ func appendCallbackParams(rawURL string, params url.Values) (string, error) { return callbackURL.String(), nil } +// oauthHTTPClient 是 OAuth 授权服务专用 HTTP 客户端。http.DefaultClient 无超时, +// 授权服务无响应时会永久阻塞授权/轮询协程,这里统一 30s 超时(ctx 仍经 +// NewRequestWithContext 传导,可提前取消)。 +var oauthHTTPClient = &http.Client{Timeout: 30 * time.Second} + func httpGetJSON(ctx context.Context, endpoint string) (map[string]any, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) if err != nil { return nil, err } req.Header.Set("User-Agent", DefaultUA) - resp, err := http.DefaultClient.Do(req) + resp, err := oauthHTTPClient.Do(req) if err != nil { return nil, err } diff --git a/internal/service/cloud115/open.go b/internal/service/cloud115/open.go index 9a51c1d..27ff4ee 100644 --- a/internal/service/cloud115/open.go +++ b/internal/service/cloud115/open.go @@ -300,6 +300,11 @@ func (c *OpenClient) GetQrCode() (*QrCodeDataReturn, error) { if err != nil { return nil, err } + // 关键字段缺失时显式报错:空 uid/sign 会导致后续扫码轮询必然失败, + // 不能把残缺响应当成功返回给界面。 + if code.Uid == "" || code.Sign == "" { + return nil, fmt.Errorf("115: 设备码响应缺少 uid/sign,无法发起扫码授权") + } return &QrCodeDataReturn{QrCodeData: *code, CodeVerifier: codeVerifier}, nil } @@ -352,6 +357,10 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) { if err != nil { return nil, err } + // 空凭证绝不能 SetAuthToken 后当成功返回:界面会显示"授权成功"但账号不可用 + if token.AccessToken == "" || token.RefreshToken == "" { + return nil, fmt.Errorf("115: 设备码换 token 返回空凭证(access_token/refresh_token 缺失)") + } c.SetAuthToken(token.AccessToken, token.RefreshToken) return token, nil } @@ -359,11 +368,30 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) { // RefreshToken 刷新访问令牌。 func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) { if refreshToken == "" { - refreshToken = c.RefreshTokenStr + refreshToken = c.currentRefreshToken() } if refreshToken == "" { return nil, fmt.Errorf("没有可用的 refresh_token") } + token, err := c.doRefreshToken(refreshToken) + if err != nil { + // refresh_token 已失效时清空内存令牌(提示需重新授权) + if IsRefreshTokenDead(err) { + c.SetAuthToken("", "") + } + return nil, err + } + if token.AccessToken == "" || token.RefreshToken == "" { + return nil, fmt.Errorf("115: 刷新返回空凭证(access_token/refresh_token 缺失)") + } + c.SetAuthToken(token.AccessToken, token.RefreshToken) + return token, nil +} + +// doRefreshToken 调用 115 刷新接口换取新令牌,不修改客户端内存状态; +// 拆出无状态方法供 tryRefreshTokenLocked(已持 tokenMu 写锁)复用, +// 避免在持锁期间重入 SetAuthToken 造成死锁。 +func (c *OpenClient) doRefreshToken(refreshToken string) (*TokenData, error) { params := map[string]string{"refresh_token": refreshToken} resp, err := c.doJSON(context.Background(), "POST", PassportAPIBase+"/open/refreshToken", params, false, 0) if err != nil && resp == nil { @@ -373,18 +401,9 @@ func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) { return nil, err } if !resp.State { - apiErr := NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败") - if IsRefreshTokenDead(apiErr) { - c.SetAuthToken("", "") - } - return nil, apiErr + return nil, NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败") } - token, err := openFirstList[TokenData](resp.Data) - if err != nil { - return nil, err - } - c.SetAuthToken(token.AccessToken, token.RefreshToken) - return token, nil + return openFirstList[TokenData](resp.Data) } // ─── 用户信息 ────────────────────────────────────────────────────────────────── diff --git a/internal/service/cloud115/oss_multipart.go b/internal/service/cloud115/oss_multipart.go index a3bc436..d5b7273 100644 --- a/internal/service/cloud115/oss_multipart.go +++ b/internal/service/cloud115/oss_multipart.go @@ -10,6 +10,7 @@ import ( "errors" "fmt" "io" + "log" "os" "sort" @@ -106,14 +107,23 @@ func (u *OSSMultipartUploader) UploadFile(ctx context.Context, input OSSMultipar return result.CallbackResult, nil } +// UploadedPart 是 OSS 已上传分片的定位信息(断点续传时复用 ETag 用)。 +type UploadedPart struct { + PartNumber int32 + Size int64 + ETag string +} + // UploadFileWithResult 上传文件并返回 multipart 结果。 -func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input OSSMultipartUploadInput) (OSSMultipartUploadResult, error) { +// 任一失败路径(分片上传失败 / callback 校验失败 / Complete 失败 / 文件打开失败等) +// 都会经 defer 统一 AbortMultipartUpload 丢弃本次 Initiate 出的 multipart +// (abort 失败仅记日志),避免 OSS 分片永久泄漏;成功路径不 Abort。 +func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input OSSMultipartUploadInput) (result OSSMultipartUploadResult, err error) { if input.PartRetryMax <= 0 { input.PartRetryMax = 3 } partSize := input.PartSize totalParts := 0 - var err error if partSize <= 0 { partSize, totalParts, err = CalculateMultipartPartSize(input.FileSize) if err != nil { @@ -124,28 +134,45 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O } uploadId := input.UploadId - if uploadId == "" { - initResult, err := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{ + // ownUploadId 标记 uploadId 是否为本调用 Initiate 出来的:仅自建的 + // multipart 在失败时由本函数 Abort;调用方显式传入的 uploadId(断点续传) + // 失败后保留现场,由调用方决定重试或清理。 + ownUploadId := uploadId == "" + if ownUploadId { + initResult, initErr := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{ Bucket: oss.Ptr(input.Bucket), Key: oss.Ptr(input.Object), RequestCommon: oss.RequestCommon{ Parameters: map[string]string{"sequential": "1"}, }, }) - if err != nil { - return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", err) + if initErr != nil { + return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", initErr) } if initResult.UploadId == nil || *initResult.UploadId == "" { return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 返回空 upload_id") } uploadId = *initResult.UploadId } + defer func() { + if err == nil || !ownUploadId || uploadId == "" { + return + } + // 失败路径统一 Abort 丢弃已上传分片;ctx 可能已取消,脱离其取消信号尽力清理 + abortCtx := context.WithoutCancel(ctx) + if _, abortErr := u.client.AbortMultipartUpload(abortCtx, &oss.AbortMultipartUploadRequest{ + Bucket: oss.Ptr(input.Bucket), + Key: oss.Ptr(input.Object), + UploadId: oss.Ptr(uploadId), + }); abortErr != nil { + log.Printf("115: 中止 OSS multipart 失败(upload_id=%s,可能残留分片):%v", uploadId, abortErr) + } + }() - existingPartMap := make(map[int32]int64) - existingParts, err := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId) - if err == nil { + existingPartMap := make(map[int32]UploadedPart) + if existingParts, listErr := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId); listErr == nil { for _, part := range existingParts { - existingPartMap[part.PartNumber] = part.Size + existingPartMap[part.PartNumber] = part } } @@ -164,13 +191,20 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O if length < 0 { length = 0 } - if existingSize, ok := existingPartMap[int32(partNumber)]; ok && existingSize == length { + // 断点续传:分片已完整上传(大小一致即代表分片大小未变)时直接复用 + // ListParts 返回的 ETag,跳过重传,也不再重复累加统计 + if existing, ok := existingPartMap[int32(partNumber)]; ok && existing.Size == length && existing.ETag != "" { uploadedBytes += length uploadedParts++ + completeParts = append(completeParts, oss.UploadPart{ + PartNumber: int32(partNumber), + ETag: oss.Ptr(existing.ETag), + }) + continue } - etag, err := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length) - if err != nil { - return OSSMultipartUploadResult{}, err + etag, uploadErr := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length) + if uploadErr != nil { + return OSSMultipartUploadResult{}, uploadErr } uploadedBytes += length uploadedParts++ @@ -224,29 +258,34 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O }, nil } -// ListUploadedParts 查询 OSS 已上传分片。 -func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]struct { - PartNumber int32 - Size int64 -}, error) { - parts := []struct { - PartNumber int32 - Size int64 - }{} - result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{ - Bucket: oss.Ptr(bucket), - Key: oss.Ptr(object), - UploadId: oss.Ptr(uploadId), - MaxParts: 1000, - }) - if err != nil { - return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err) - } - for _, part := range result.Parts { - parts = append(parts, struct { - PartNumber int32 - Size int64 - }{PartNumber: part.PartNumber, Size: part.Size}) +// ListUploadedParts 查询 OSS 已上传分片(MaxParts 上限 1000,超过时按 +// NextPartNumberMarker 自动翻页取全量,否则断点续传只能看到前 1000 片)。 +func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]UploadedPart, error) { + parts := []UploadedPart{} + var marker int32 + for { + result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{ + Bucket: oss.Ptr(bucket), + Key: oss.Ptr(object), + UploadId: oss.Ptr(uploadId), + MaxParts: 1000, + PartNumberMarker: marker, + }) + if err != nil { + return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err) + } + for _, part := range result.Parts { + etag := "" + if part.ETag != nil { + etag = *part.ETag + } + parts = append(parts, UploadedPart{PartNumber: part.PartNumber, Size: part.Size, ETag: etag}) + } + if !result.IsTruncated || result.NextPartNumberMarker <= marker { + // 防御:marker 不前进时终止循环,避免异常响应导致死循环 + break + } + marker = result.NextPartNumberMarker } return parts, nil } diff --git a/internal/service/cloud115/upload.go b/internal/service/cloud115/upload.go index eae8e9d..77e4d99 100644 --- a/internal/service/cloud115/upload.go +++ b/internal/service/cloud115/upload.go @@ -262,7 +262,10 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s } switch status { case UploadInitStatusRapidUploaded: - // 秒传成功 + // 秒传成功:必须带远端文件定位信息,否则视为异常响应 + if initResult.FileId == "" || initResult.PickCode == "" { + return nil, fmt.Errorf("115: 秒传成功但缺少 file_id/pick_code(status=%d)", status) + } return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil case UploadInitStatusSignFailed: return nil, fmt.Errorf("115: 签名验证后失败") @@ -271,7 +274,8 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s case UploadInitStatusNeedUpload: // 真实上传:OSS multipart default: - return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil + // 未知状态不能当成功返回(会静默丢文件),显式报错便于排查 + return nil, fmt.Errorf("115: 未知的 upload/init 状态 %d", status) } if initResult.Bucket == "" || initResult.Object == "" { diff --git a/internal/service/cloud115/utils.go b/internal/service/cloud115/utils.go index fa7f6b7..2e76365 100644 --- a/internal/service/cloud115/utils.go +++ b/internal/service/cloud115/utils.go @@ -1,14 +1,26 @@ package cloud115 -import "math/rand" +import ( + "crypto/rand" + "fmt" + "math/big" +) const randCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" -// RandomString 生成指定长度的随机字符串(PKCE code_verifier 等)。 +// RandomString 生成指定长度的密码学安全随机字符串(PKCE code_verifier、 +// OAuth state 等安全敏感场景)。必须使用 crypto/rand:math/rand 未播种时 +// 序列可预测,会造成 PKCE 防御失效。 func RandomString(length int) string { b := make([]byte, length) + max := big.NewInt(int64(len(randCharset))) for i := range b { - b[i] = randCharset[rand.Intn(len(randCharset))] + n, err := rand.Int(rand.Reader, max) + if err != nil { + // 仅在系统熵源不可用时发生;静默降级为弱随机不可接受,直接暴露 + panic(fmt.Errorf("115: 生成安全随机字符串失败:%w", err)) + } + b[i] = randCharset[n.Int64()] } return string(b) } diff --git a/internal/service/dlna.go b/internal/service/dlna.go index 81c3c5f..fb176ef 100644 --- a/internal/service/dlna.go +++ b/internal/service/dlna.go @@ -36,6 +36,10 @@ type DLNAService struct { cachedAt time.Time } +// dlnaHTTPClient 是 DLNA 专用 HTTP 客户端:SSDP 描述拉取与 SOAP 投递 +// 都应快速失败,不占用全局 DefaultClient,也不无限悬挂。 +var dlnaHTTPClient = &http.Client{Timeout: 15 * time.Second} + // NewDLNAService is the constructor. func NewDLNAService(log *zap.Logger) *DLNAService { return &DLNAService{log: log} @@ -153,7 +157,7 @@ func (d *DLNAService) fetchDescription(ctx context.Context, location string) (*D if err != nil { return nil, err } - resp, err := http.DefaultClient.Do(req) + resp, err := dlnaHTTPClient.Do(req) if err != nil { return nil, err } @@ -267,7 +271,7 @@ func (d *DLNAService) soap(ctx context.Context, controlURL, action, envelope str req.Header.Set("Content-Type", `text/xml; charset="utf-8"`) req.Header.Set("SOAPAction", fmt.Sprintf(`"urn:schemas-upnp-org:service:AVTransport:1#%s"`, action)) - resp, err := http.DefaultClient.Do(req) + resp, err := dlnaHTTPClient.Do(req) if err != nil { return err } diff --git a/internal/service/emby_compat.go b/internal/service/emby_compat.go index edc7d8a..389469c 100644 --- a/internal/service/emby_compat.go +++ b/internal/service/emby_compat.go @@ -18,6 +18,7 @@ import ( "time" "github.com/truewhile/MeBox/internal/config" + "github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/repository" "go.uber.org/zap" ) @@ -267,12 +268,19 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[ if err != nil { return nil, err } - type remoteResult struct { - items []any + type remoteReply struct { + acct *model.StrmAccount + envelope map[string]any } mounts, aerr := e.remote.ListMounts(ctx) - results := make([]remoteResult, 0, len(mounts)) + replies := make([]*remoteReply, 0, len(mounts)) if aerr == nil { + type mountSearchJob struct { + idx int + mount *model.EmbyMount + acct *model.StrmAccount + } + jobs := make([]*mountSearchJob, 0, len(mounts)) for i := range mounts { m := mounts[i] if !m.Enabled { @@ -285,40 +293,76 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[ if acct == nil { continue } - // 按挂载逐个搜索:搜索结果归属明确(伪装 ID 正确),也天然只搜已 - // 挂载的媒体库。 - searchParams := p - searchParams.ParentID = "" // RemoteSearchMount 内部设 ParentId - remote, rerr := e.remote.RemoteSearchMount(ctx, &m, acct, p) - if rerr != nil { - if e.log != nil { - e.log.Warn("remote emby search failed", - zap.String("account", acct.Name), zap.Error(rerr)) + // idx 使用 jobs 内的序号(而非 mounts 下标):fetched 按 + // len(jobs) 分配,必须与 jobs 下标对齐,否则越界 panic。 + jobs = append(jobs, &mountSearchJob{idx: len(jobs), mount: &mounts[i], acct: acct}) + } + // 并发搜索各挂载(限并发 + 单挂载超时):串行时每挂载最多 + // 15s×线路数,多挂载下首屏延迟被成倍放大。结果按挂载顺序合并。 + sem := make(chan struct{}, 4) + var wg sync.WaitGroup + fetched := make([]*remoteReply, len(jobs)) + for _, job := range jobs { + wg.Add(1) + go func(job *mountSearchJob) { + defer wg.Done() + sem <- struct{}{} + defer func() { <-sem }() + sctx, cancel := context.WithTimeout(ctx, 8*time.Second) + defer cancel() + sp := p + // 远程只取首页:此前每个远程各自按 StartIndex 分页,拼接后 + // 又被 sliceSearchItems 再切一次——分页被二次偏移,远程结果 + // 首屏不可见、翻页错位。合并后由 sliceSearchItems 单点分页。 + sp.StartIndex = 0 + sp.ParentID = "" // RemoteSearchMount 内部设 ParentId + remote, rerr := e.remote.RemoteSearchMount(sctx, job.mount, job.acct, sp) + if rerr != nil { + if e.log != nil { + e.log.Warn("remote emby search failed", + zap.String("account", job.acct.Name), zap.Error(rerr)) + } + return } - continue - } - if err := e.mergeRemoteUserData(ctx, p.UserID, remote); err != nil { - return nil, err - } - if raw, ok := remote["Items"].([]any); ok { - results = append(results, remoteResult{items: raw}) - } else if rawMap, ok := remote["Items"].([]map[string]any); ok { - converted := make([]any, 0, len(rawMap)) - for _, m := range rawMap { - converted = append(converted, any(m)) - } - results = append(results, remoteResult{items: converted}) + fetched[job.idx] = &remoteReply{acct: job.acct, envelope: remote} + }(job) + } + wg.Wait() + for _, r := range fetched { + if r != nil { + replies = append(replies, r) } } } - items := make([]any, 0, len(localItemsAsAny(local))+len(results)*p.Limit) + items := make([]any, 0, len(localItemsAsAny(local))+len(replies)*p.Limit) items = append(items, localItemsAsAny(local)...) - for _, res := range results { - items = append(items, res.items...) + for _, reply := range replies { + if err := e.mergeRemoteUserData(ctx, p.UserID, reply.envelope); err != nil { + return nil, err + } + items = append(items, remoteItemsAsAny(reply.envelope)...) } return sliceSearchItems(items, p), nil } +// remoteItemsAsAny 提取远程载荷的 Items 列表(兼容 []any 与 []map 形态)。 +func remoteItemsAsAny(envelope map[string]any) []any { + if envelope == nil { + return nil + } + if raw, ok := envelope["Items"].([]any); ok { + return raw + } + if rawMap, ok := envelope["Items"].([]map[string]any); ok { + converted := make([]any, 0, len(rawMap)) + for _, m := range rawMap { + converted = append(converted, any(m)) + } + return converted + } + return nil +} + func localItemsAsAny(envelope map[string]any) []any { if envelope == nil { return nil diff --git a/internal/service/emby_items_detail.go b/internal/service/emby_items_detail.go index 0ef7f75..b8dbdd6 100644 --- a/internal/service/emby_items_detail.go +++ b/internal/service/emby_items_detail.go @@ -338,10 +338,12 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st return map[string]any{"Items": []any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil } + // 历史记录限行:此前无上限全量加载,远程条目多时既拖慢 SQL 也放大 + // 下面的远程详情请求量。 var hist []model.PlaybackHistory if err := e.repo.DB.WithContext(ctx). Where("user_id = ? AND completed = ? AND position_ms > 0", p.UserID, false). - Order("watched_at desc").Find(&hist).Error; err != nil { + Order("watched_at desc").Limit(200).Find(&hist).Error; err != nil { return nil, err } if len(hist) == 0 { @@ -367,18 +369,31 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st } } - items := make([]map[string]any, 0, len(hist)) + // 分页前置:凑满 StartIndex+Limit 条即停,不再为「总数」逐条发远程 + // 详情 GET(此前每条远程记录一次串行 GET,远程慢时请求挂起数分钟)。 + // 总数用候选行数(本地过滤后 + 远程候选),对继续观看行的翻页语义 + // 足够准确。 + needed := p.StartIndex + p.Limit + items := make([]map[string]any, 0, p.Limit) + localTotal, remoteTotal := 0, 0 for _, h := range hist { if m, ok := byID[h.MediaID]; ok { if p.ParentID != "" && m.LibraryID != p.ParentID && m.SeriesID != p.ParentID { continue } - items = append(items, e.itemPayload(ctx, m, false, h.PositionMs)) + localTotal++ + if produced := len(items); produced < needed { + items = append(items, e.itemPayload(ctx, m, false, h.PositionMs)) + } continue } if e.remote == nil || !IsEmbyRemoteID(h.MediaID) { continue } + remoteTotal++ + if len(items) >= needed { + continue + } mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID) mount, acct, err := e.remote.ResolveMount(ctx, mountID) if err != nil || mount == nil || acct == nil { @@ -399,7 +414,7 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st items = append(items, item) } - total := int64(len(items)) + total := int64(localTotal + remoteTotal) if p.StartIndex >= len(items) { return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil } diff --git a/internal/service/emby_remote.go b/internal/service/emby_remote.go index afd5aa2..4ef77fe 100644 --- a/internal/service/emby_remote.go +++ b/internal/service/emby_remote.go @@ -26,6 +26,7 @@ import ( "regexp" "strconv" "strings" + "sync" "time" "go.uber.org/zap" @@ -73,6 +74,7 @@ type EmbyRemoteService struct { repo *repository.Container crypto *CryptoService http *http.Client + stream *http.Client // 流式代理专用(视频/字幕),无整体 Timeout cache *RuntimeCacheService } @@ -87,6 +89,12 @@ func NewEmbyRemoteService(cfg *config.Config, log *zap.Logger, repo *repository. Timeout: embyRemoteHTTPTimeout, Transport: &embyRemoteTransport{base: http.DefaultTransport}, }, + // 流式代理必须用无整体 Timeout 的 client:http.Client.Timeout + // 覆盖整个响应体读取过程,15s 的常规超时会让代理播放播到 + // 15 秒整被掐断。生命周期由请求 ctx 控制。 + stream: &http.Client{ + Transport: &embyRemoteTransport{base: http.DefaultTransport}, + }, } } @@ -352,15 +360,9 @@ func (r *EmbyRemoteService) resolveRemoteUserID(ctx context.Context, acct *model return } cfg.RemoteUserID = uid - raw := map[string]string{} - _ = json.Unmarshal([]byte(acct.Config), &raw) - raw["remote_user_id"] = uid - data, err := json.Marshal(raw) - if err != nil { - return - } - acct.Config = string(data) - _ = r.repo.StrmAccount.Update(ctx, acct) + _ = r.updateAccountConfig(ctx, acct, func(raw map[string]string) { + raw["remote_user_id"] = uid + }) } // CleanupOrphanMounts 清理账号已删除的残留挂载(老版本删除账号未级联), @@ -495,19 +497,28 @@ func (r *EmbyRemoteService) ensureTokenOnLine(ctx context.Context, acct *model.S return nil } -// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。 -func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error { - if acct == nil { +// acctCfgMu 序列化对账号 Config 的读-改-写。并发请求若各自基于请求开始 +// 时的快照做整包覆盖,会互相丢失更新(刚持久化的 token / active_line 被 +// 旧快照覆盖回去)。 +var acctCfgMu sync.Mutex + +// updateAccountConfig 在互斥下重读账号最新 Config,应用 mutate 后写回, +// 并同步调用方持有的 acct 快照。 +func (r *EmbyRemoteService) updateAccountConfig(ctx context.Context, acct *model.StrmAccount, mutate func(raw map[string]string)) error { + if acct == nil || r.repo == nil { return nil } + acctCfgMu.Lock() + defer acctCfgMu.Unlock() raw := map[string]string{} + if fresh, err := r.repo.StrmAccount.FindByID(ctx, acct.ID); err == nil && fresh != nil { + acct.Config = fresh.Config // 以 DB 最新值为基线,避免覆盖并发写入 + } if strings.TrimSpace(acct.Config) != "" { _ = json.Unmarshal([]byte(acct.Config), &raw) } - raw["api_key"] = r.crypto.Encrypt(cfg.Token) - raw["remote_user_id"] = cfg.RemoteUserID - if strings.TrimSpace(raw["username"]) == "" { - raw["username"] = cfg.Username + if mutate != nil { + mutate(raw) } data, err := json.Marshal(raw) if err != nil { @@ -517,6 +528,20 @@ func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAc return r.repo.StrmAccount.Update(ctx, acct) } +// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。 +func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error { + if acct == nil { + return nil + } + return r.updateAccountConfig(ctx, acct, func(raw map[string]string) { + raw["api_key"] = r.crypto.Encrypt(cfg.Token) + raw["remote_user_id"] = cfg.RemoteUserID + if strings.TrimSpace(raw["username"]) == "" { + raw["username"] = cfg.Username + } + }) +} + // doGet 向远程 Emby 发起带 api_key 的 GET,把响应 JSON 解码到 out。 // 401 时自动重认证一次再重试(凭据过期场景)。连接失败时按线路优先级自动切换。 func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, path string, q url.Values, out any) error { @@ -567,22 +592,28 @@ func (r *EmbyRemoteService) doGetOnLine(ctx context.Context, acct *model.StrmAcc if err != nil { return fmt.Errorf("请求远程 Emby 失败: %w", err) } - data, readErr := io.ReadAll(io.LimitReader(resp.Body, 8<<20)) + // 读 8MB+1 以区分"刚好 8MB"与"被截断":截断的 JSON 会让 + // Unmarshal 报 unexpected end,难以定位;这里显式报错。 + data, readErr := io.ReadAll(io.LimitReader(resp.Body, (8<<20)+1)) resp.Body.Close() if readErr != nil { return readErr } + if len(data) > 8<<20 { + return fmt.Errorf("远程 Emby 响应超过 8MB 上限(路径 %s):请减小分页或 Fields 字段", path) + } if resp.StatusCode == http.StatusUnauthorized && attempt == 0 { + // 401:只清当前线路的内存 token 并立即重认证;不在此时删除 + // DB 里的 api_key——①外层还会按线路故障转移(其他线路可能 + // 存有自己的 token);②纯 api_key 账号删除后无法再认证,一次 + // 线路误报就会把账号“砖化”。重认证成功后 persistToken 会用 + // 新 token 覆盖 api_key。 cfg.Token = "" - master.Token = "" - if acct != nil { - raw := map[string]string{} - _ = json.Unmarshal([]byte(acct.Config), &raw) - delete(raw, "api_key") - enc, _ := json.Marshal(raw) - acct.Config = string(enc) - _ = r.repo.StrmAccount.Update(ctx, acct) + if err := r.ensureTokenOnLine(ctx, acct, cfg); err != nil { + return fmt.Errorf("认证重试失败: %w", err) } + master.Token = cfg.Token + master.RemoteUserID = cfg.RemoteUserID continue } if resp.StatusCode >= 300 { @@ -957,7 +988,7 @@ func (r *EmbyRemoteService) proxyVideoStreamOnLine(ctx context.Context, w http.R if rangeHeader := req.Header.Get("Range"); rangeHeader != "" { upstream.Header.Set("Range", rangeHeader) } - resp, err := r.http.Do(upstream) + resp, err := r.stream.Do(upstream) if err != nil { return fmt.Errorf("连接远程 Emby 视频流失败: %w", err) } @@ -1028,7 +1059,7 @@ func (r *EmbyRemoteService) proxySubtitleOnLine(ctx context.Context, w http.Resp return err } upstream.Header.Set("X-Emby-Token", cfg.Token) - resp, err := r.http.Do(upstream) + resp, err := r.stream.Do(upstream) if err != nil { return fmt.Errorf("连接远程 Emby 字幕流失败: %w", err) } diff --git a/internal/service/emby_remote_lines.go b/internal/service/emby_remote_lines.go index e46e13b..a3b1785 100644 --- a/internal/service/emby_remote_lines.go +++ b/internal/service/emby_remote_lines.go @@ -124,10 +124,11 @@ func isEmbyLineFailoverError(err error) bool { return false } msg := strings.ToLower(err.Error()) + // 注意:认证类错误(重认证失败 / 缺少凭据)不在此排除——401 后清空 + // 内存 token 重认证失败时应继续按线路故障转移,其他线路可能存有 + // 自己的 token。仅“登录失败”(密码错误)是账号级问题,无需换线。 if strings.Contains(msg, "登录失败") || - strings.Contains(msg, "未返回 accesstoken") || - strings.Contains(msg, "缺少 emby 凭据") || - strings.Contains(msg, "认证重试失败") { + strings.Contains(msg, "未返回 accesstoken") { return false } var urlErr *url.Error @@ -147,20 +148,13 @@ func (r *EmbyRemoteService) persistActiveLine(ctx context.Context, acct *model.S if acct == nil || cfg == nil || lineIndex < 0 || lineIndex >= len(cfg.Lines) { return nil } - raw := map[string]string{} - if strings.TrimSpace(acct.Config) != "" { - _ = json.Unmarshal([]byte(acct.Config), &raw) - } - raw["active_line"] = strconv.Itoa(lineIndex) - raw["url"] = cfg.Lines[lineIndex].URL - data, err := json.Marshal(raw) - if err != nil { - return err - } - acct.Config = string(data) + err := r.updateAccountConfig(ctx, acct, func(raw map[string]string) { + raw["active_line"] = strconv.Itoa(lineIndex) + raw["url"] = cfg.Lines[lineIndex].URL + }) cfg.ActiveLine = lineIndex cfg.BaseURL = normalizeEmbyRemoteURL(cfg.Lines[lineIndex].URL) - return r.repo.StrmAccount.Update(ctx, acct) + return err } func (r *EmbyRemoteService) adoptWorkingLine(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, lineIndex int) { diff --git a/internal/service/emby_remote_web.go b/internal/service/emby_remote_web.go index cd60967..6dd7a09 100644 --- a/internal/service/emby_remote_web.go +++ b/internal/service/emby_remote_web.go @@ -371,37 +371,52 @@ func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model // RemoteMediaDetail 拉远程单条目映射为 Media(网页详情页)。 func (r *EmbyRemoteService) RemoteMediaDetail(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, error) { + m, _, err := r.remoteMediaDetailRaw(ctx, mount, acct, remoteID) + return m, err +} + +// remoteMediaDetailRaw 拉取远程条目详情,同时返回原始载荷(ID 已伪装), +// 供调用方免二次请求读取 Type / SeriesId 等字段。 +func (r *EmbyRemoteService) remoteMediaDetailRaw(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, map[string]any, error) { cfg, err := r.remoteConfigWithToken(ctx, acct) if err != nil { - return nil, err + return nil, nil, err } path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID) path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating" var out map[string]any if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil { - return nil, err + return nil, nil, err } RewriteEmbyRemoteIDs(out, mount.ID) m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, out) - return &m, nil + return &m, out, nil } // RemoteEpisodes 拉远程条目下的集列表(Series/Season/Folder→子集;Episode→同系列; // Movie→自身单条),按季/集排序,与本地 ListMediaEpisodes 行为一致。 func (r *EmbyRemoteService) RemoteEpisodes(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) ([]model.Media, error) { - detail, err := r.RemoteMediaDetail(ctx, mount, acct, remoteID) + detail, rawDetail, err := r.remoteMediaDetailRaw(ctx, mount, acct, remoteID) if err != nil { return nil, err } // 用远程详情载荷精判类型(Episode→同系列;Series/Season/Folder→子集;Movie→单条)。 - itemType := r.remoteItemType(ctx, acct, remoteID) + // Type/SeriesId 都在详情载荷里现成可用,不再为判定类型/系列额外发起 + // 两次重复的远程全量 GET(远程慢时页面延迟直接×3)。 + itemType := remoteItemString(rawDetail, "Type") if itemType == "" { itemType = remoteItemTypeOf(detail) } var parentID string switch itemType { case "Episode": - parentID = r.remoteItemSeriesID(ctx, acct, remoteID) + parentID = remoteItemString(rawDetail, "SeriesId") + if _, rid, ok := DecodeEmbyRemoteID(parentID); ok { + parentID = rid // 载荷 ID 已伪装,远程查询需要原始 ID + } + if parentID == "" { + parentID = r.remoteItemSeriesID(ctx, acct, remoteID) + } if parentID == "" { parentID = remoteID } @@ -435,23 +450,36 @@ func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.E q.Set("ParentId", parentID) q.Set("IncludeItemTypes", "Episode") q.Set("Recursive", "true") - q.Set("StartIndex", "0") - q.Set("Limit", "500") q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating") - var body struct { - Items []map[string]any `json:"Items"` - TotalRecordCount int64 `json:"TotalRecordCount"` + items := make([]model.Media, 0, 64) + total := int64(0) + // 每页 200 循环拉全:MediaStreams/MediaSources 重字段下单页 500 条 + // 已贴近 8MB 截断上限;单次大页超限会静默解析失败。 + const episodePageSize = 200 + for startIndex := 0; ; startIndex += episodePageSize { + q.Set("StartIndex", strconv.Itoa(startIndex)) + q.Set("Limit", strconv.Itoa(episodePageSize)) + var body struct { + Items []map[string]any `json:"Items"` + TotalRecordCount int64 `json:"TotalRecordCount"` + } + if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil { + return nil, 0, err + } + total = body.TotalRecordCount + if len(body.Items) == 0 { + break + } + for _, it := range body.Items { + RewriteEmbyRemoteIDs(it, mount.ID) + m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it) + items = append(items, m) + } + if len(body.Items) < episodePageSize { + break + } } - if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil { - return nil, 0, err - } - items := make([]model.Media, 0, len(body.Items)) - for _, it := range body.Items { - RewriteEmbyRemoteIDs(it, mount.ID) - m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it) - items = append(items, m) - } - return items, body.TotalRecordCount, nil + return items, total, nil } // RemoteSeriesCards 远程剧集库的系列卡片(ChildCount 作为集数)。 @@ -477,14 +505,16 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model. q.Set("Recursive", "false") q.Set("SortBy", "DateLastContentAdded") q.Set("SortOrder", "Descending") - q.Set("Limit", "1000") + // 每页 200:Fields 带全量重字段(Overview/MediaStreams 等)时单页 1000 + // 条的载荷会超过 doGet 的 8MB 截断上限,JSON 被静默截断直接解析失败。 + q.Set("Limit", "200") q.Set("Fields", "Overview,Genres,ProviderIds,Path,RecursiveItemCount,SeriesPrimaryImage,DateCreated,DateLastMediaAdded,PremiereDate,ProductionYear,CommunityRating,CriticRating") var body struct { Items []map[string]any `json:"Items"` TotalRecordCount int64 `json:"TotalRecordCount"` } cards := make([]SeriesCard, 0) - for startIndex := 0; ; startIndex += 1000 { + for startIndex := 0; ; startIndex += 200 { q.Set("StartIndex", strconv.Itoa(startIndex)) body.Items = nil if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil { diff --git a/internal/service/emby_user_data.go b/internal/service/emby_user_data.go index eba8026..8a3a2fb 100644 --- a/internal/service/emby_user_data.go +++ b/internal/service/emby_user_data.go @@ -5,6 +5,7 @@ import ( "errors" "strconv" "strings" + "sync" "time" "github.com/truewhile/MeBox/internal/model" @@ -218,7 +219,22 @@ func mergedRemoteUserData(raw any, history *model.PlaybackHistory) map[string]an return userData } +// embyInvalMu 节流全量缓存失效:播放期间客户端每 5-10s 上报一次进度, +// 每次都 SCAN+DEL 全部 media:emby:* 缓存会把缓存命中率持续打穿(其他 +// 客户端每次翻页都回源 SQL)。条目载荷的 UserData 在请求时动态合并, +// 进度类变更做 30s 节流即可,不影响正确性观感。 +var ( + embyInvalMu sync.Mutex + embyInvalLast time.Time +) + func (e *EmbyService) invalidateEmbyItemsCache(ctx context.Context) { + embyInvalMu.Lock() + defer embyInvalMu.Unlock() + if !embyInvalLast.IsZero() && time.Since(embyInvalLast) < 30*time.Second { + return + } + embyInvalLast = time.Now() if e.cache != nil { e.cache.DeletePrefix(ctx, "media:emby:") } diff --git a/internal/service/image_proxy.go b/internal/service/image_proxy.go index 67dd99d..c3e32cd 100644 --- a/internal/service/image_proxy.go +++ b/internal/service/image_proxy.go @@ -13,9 +13,12 @@ package service import ( + "errors" + "net" "net/http" "path/filepath" "sync" + "syscall" "time" "go.uber.org/zap" @@ -52,6 +55,33 @@ func NewImageProxy(cfg *config.Config, log *zap.Logger) *ImageProxy { // from image.tmdb.org via their HTTP proxy without extra config. On // Windows we also honor the current user's system proxy settings. transport := NewExternalTransport() + if proxyConfiguredForImageFetch() { + // 走本地代理(如 127.0.0.1:7890)时,拨号目标是代理本身, + // 连接层 SSRF 校验会误杀本地回环代理;此时沿用 URL 级校验。 + log.Info("image proxy: outbound proxy detected, connection-level SSRF guard disabled") + } else { + // 仅 URL 解析层的 isPrivateHost 可被十进制/十六进制 IP、解析到 + // 私网的域名与 DNS rebinding 绕过;在拨号层对最终连接 IP 做二次 + // 校验(含重定向后的每条连接)堵住该旁路。 + dialer := &net.Dialer{ + Timeout: 15 * time.Second, + Control: func(_, address string, _ syscall.RawConn) error { + host, _, err := net.SplitHostPort(address) + if err != nil { + return err + } + ip := net.ParseIP(host) + if ip == nil { + return errors.New("image proxy: refusing non-IP dial target") + } + if isPrivateIP(ip) { + return errors.New("image proxy: requests to private/internal hosts are not allowed") + } + return nil + }, + } + transport.DialContext = dialer.DialContext + } return &ImageProxy{ cfg: cfg, log: log, @@ -60,6 +90,16 @@ func NewImageProxy(cfg *config.Config, log *zap.Logger) *ImageProxy { } } +// proxyConfiguredForImageFetch 探测环境变量或系统代理是否会影响图片抓取。 +func proxyConfiguredForImageFetch() bool { + req, err := http.NewRequest(http.MethodGet, "https://image.tmdb.org/", nil) + if err != nil { + return false + } + proxy, err := ProxyFromEnvironmentOrSystem(req) + return err == nil && proxy != nil +} + // SetLibraryRootsProvider injects a callback that returns the current set of // media library root directories. Sidecar posters live under these roots // (which are arbitrary, user-defined, and not necessarily under the diff --git a/internal/service/image_proxy_paths.go b/internal/service/image_proxy_paths.go index e73280a..adca575 100644 --- a/internal/service/image_proxy_paths.go +++ b/internal/service/image_proxy_paths.go @@ -35,13 +35,18 @@ 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() + if ip := net.ParseIP(host); ip != nil { + return isPrivateIP(ip) } return false } +// isPrivateIP 判定单个 IP 是否属于回环/私网/链路本地/未指定地址。 +func isPrivateIP(ip net.IP) bool { + return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || + ip.IsLinkLocalMulticast() || ip.IsUnspecified() +} + // 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} diff --git a/internal/service/organizer_directory_versions.go b/internal/service/organizer_directory_versions.go index e85d25d..1be8b62 100644 --- a/internal/service/organizer_directory_versions.go +++ b/internal/service/organizer_directory_versions.go @@ -2,6 +2,7 @@ package service import ( "context" + "fmt" "os" "path/filepath" "strconv" @@ -219,24 +220,50 @@ func (o *OrganizerService) replaceVersions(ctx context.Context, src string, exis o.log.Warn("organize replace sidecar artwork failed", zap.String("from", src), zap.String("to", dst), zap.Error(err)) } - // New file is safely staged; the transfer succeeded so it is now safe to - // supersede the existing lower-res versions. - for _, e := range existing { - if nfo := nfoPath(e); nfo != "" { - _ = os.Remove(nfo) + // New file is safely staged. 先把现有 dst 改名为备份、rename stage→dst + // 成功后,才删除旧版本:此前顺序是先删旧版本再 rename,一旦 rename + // 失败(Windows 下 dst 被播放器/杀软占用很常见),cleanup 会删掉 + // stage——旧版本已删、move 模式下源已不在、新文件也删,数据彻底丢失。 + var backup string + if _, err := os.Stat(dst); err == nil { + backup = dst + ".replacing-" + randomSuffix() + if err := os.Rename(dst, backup); err != nil { + cleanup() + return fmt.Errorf("备份现有文件失败(可能被其他程序占用):%w", err) } - if err := os.Remove(e); err != nil && !os.IsNotExist(err) { - o.log.Warn("organize replace remove existing failed", - zap.String("path", e), zap.Error(err)) + } + if err := os.Rename(stage, dst); err != nil { + if backup != "" { + if rbErr := os.Rename(backup, dst); rbErr != nil { + o.log.Error("organize replace restore backup failed", + zap.String("backup", backup), zap.Error(rbErr)) + } + } + cleanup() + return err + } + // 新文件已就位,现在才删除被取代的旧版本。dst 路径此时已是新文件, + // 文件级删除必须跳过(DB 行仍按原语义清理)。 + for _, e := range existing { + if e != dst { + if nfo := nfoPath(e); nfo != "" { + _ = os.Remove(nfo) + } + if err := os.Remove(e); err != nil && !os.IsNotExist(err) { + o.log.Warn("organize replace remove existing failed", + zap.String("path", e), zap.Error(err)) + } } if o.repo != nil && o.repo.DB != nil { _ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", e).Delete(&model.Media{}).Error } } - // Move staged file + sidecars into the final path. - if err := os.Rename(stage, dst); err != nil { - cleanup() - return err + if backup != "" { + // 备份文件即被取代的旧 dst 内容,新文件已成功落位后移除。 + if err := os.Remove(backup); err != nil && !os.IsNotExist(err) { + o.log.Warn("organize replace remove backup failed", + zap.String("path", backup), zap.Error(err)) + } } moveSidecarRename(nfoPath(stage), nfoPath(dst)) moveStagedArtwork(stage, dst) diff --git a/internal/service/playback.go b/internal/service/playback.go index cede9ba..270ed34 100644 --- a/internal/service/playback.go +++ b/internal/service/playback.go @@ -291,10 +291,29 @@ func (p *PlaybackService) GetPlaylist(ctx context.Context, playlistID string) (* return &PlaylistDetail{Playlist: pl, Items: ordered}, nil } +// ErrPlaylistForbidden 表示当前用户无权操作目标播放列表。 +var ErrPlaylistForbidden = errors.New("forbidden") + +// EnsurePlaylistOwner 校验播放列表属主;admin 可操作任意列表。 +// 非存在的列表返回 gorm.ErrRecordNotFound。 +func (p *PlaybackService) EnsurePlaylistOwner(ctx context.Context, playlistID, userID string, isAdmin bool) error { + var pl model.Playlist + if err := p.repo.DB.WithContext(ctx).Select("user_id").Where("id = ?", playlistID).First(&pl).Error; err != nil { + return err + } + if pl.UserID != userID && !isAdmin { + return ErrPlaylistForbidden + } + return nil +} + // AddToPlaylist appends a media item to the end of a playlist. -func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID string) error { +func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, userID, mediaID string, isAdmin bool) error { + if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil { + return err + } var count int64 - if err := p.repo.DB.Model(&model.PlaylistItem{}). + if err := p.repo.DB.WithContext(ctx).Model(&model.PlaylistItem{}). Where("playlist_id = ?", playlistID).Count(&count).Error; err != nil { return err } @@ -307,14 +326,20 @@ func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID } // RemoveFromPlaylist 物理删除播放列表项(幂等)。 -func (p *PlaybackService) RemoveFromPlaylist(ctx context.Context, playlistID, mediaID string) error { +func (p *PlaybackService) RemoveFromPlaylist(ctx context.Context, playlistID, userID, mediaID string, isAdmin bool) error { + if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil { + return err + } return p.repo.DB.WithContext(ctx).Unscoped(). Where("playlist_id = ? AND media_id = ?", playlistID, mediaID). Delete(&model.PlaylistItem{}).Error } // DeletePlaylist 物理删除播放列表及其全部条目。 -func (p *PlaybackService) DeletePlaylist(ctx context.Context, playlistID string) error { +func (p *PlaybackService) DeletePlaylist(ctx context.Context, playlistID, userID string, isAdmin bool) error { + if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil { + return err + } if err := p.repo.DB.WithContext(ctx).Unscoped().Where("playlist_id = ?", playlistID). Delete(&model.PlaylistItem{}).Error; err != nil { return err diff --git a/internal/service/runtime_settings.go b/internal/service/runtime_settings.go index 95e0a1c..f436695 100644 --- a/internal/service/runtime_settings.go +++ b/internal/service/runtime_settings.go @@ -28,6 +28,8 @@ func ApplyRuntimeSettings(ctx context.Context, cfg *config.Config, repos *reposi } func ApplyRuntimeSetting(cfg *config.Config, key, value string) { + config.RuntimeMu.Lock() + defer config.RuntimeMu.Unlock() if cfg == nil { return } diff --git a/internal/service/scanner_prune.go b/internal/service/scanner_prune.go index 12a48af..5513bb0 100644 --- a/internal/service/scanner_prune.go +++ b/internal/service/scanner_prune.go @@ -16,8 +16,33 @@ func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, er if _, err := os.Stat(path); err == nil { return 0, nil // still exists; nothing to remove } + // 目录整体消失(删除/改名离开):连同其子树下的媒体行一并移除。 + // 此前只删 path 精确匹配的行——目录本身通常没有 media 行,导致目录 + // 改名后旧子树记录全部失联,只有全量扫描才能修复。 + prefix := filepath.Clean(path) + string(filepath.Separator) + var rows []struct { + ID string + Path string + } + if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}). + Select("id, path"). + Where("path = ? OR path LIKE ?", path, prefix+"%"). + Find(&rows).Error; err != nil { + return 0, err + } + ids := make([]string, 0, len(rows)) + for _, row := range rows { + // LIKE 里的 % _ 是通配符(候选集只会偏大),用 Go 前缀精确过滤, + // 避免对含 % / _ 的路径误删。 + if row.Path == path || strings.HasPrefix(filepath.Clean(row.Path), prefix) { + ids = append(ids, row.ID) + } + } + if len(ids) == 0 { + return 0, nil + } res := s.repo.DB.WithContext(ctx).Unscoped(). - Where("path = ?", path). + Where("id IN ?", ids). Delete(&model.Media{}) if res.Error == nil && res.RowsAffected > 0 { s.invalidateMediaCache(ctx) diff --git a/internal/service/scraper_queue.go b/internal/service/scraper_queue.go index a63946c..a782e90 100644 --- a/internal/service/scraper_queue.go +++ b/internal/service/scraper_queue.go @@ -35,6 +35,13 @@ func (s *ScraperService) Start(ctx context.Context) { if s == nil { return } + // 启动自愈:进程中断遗留的 running 任务重置为 pending,否则永久卡死 + // (ClaimPending 只认 pending,重试按钮也拒绝 running)。 + if n, err := s.repo.ScrapeTask.ResetRunningToPending(ctx); err == nil && n > 0 && s.log != nil { + s.log.Warn("scrape tasks reset from running to pending after restart", zap.Int64("count", n)) + } else if err != nil && s.log != nil { + s.log.Warn("reset running scrape tasks failed", zap.Error(err)) + } go s.queueWorker(ctx) } @@ -69,6 +76,9 @@ func (s *ScraperService) queueWorker(ctx context.Context) { defer wg.Done() select { case <-ctx.Done(): + // 任务已被 Claim 置为 running:停机前回写 pending, + // 避免留下永久卡死的任务。 + s.requeueClaimedScrapeTask(t) return case sem <- struct{}{}: } @@ -84,6 +94,15 @@ func (s *ScraperService) queueWorker(ctx context.Context) { } } +// requeueClaimedScrapeTask 把已认领但未开始执行的任务回写为 pending。 +func (s *ScraperService) requeueClaimedScrapeTask(t *model.ScrapeTask) { + t.Status = model.ScrapeTaskPending + t.StartedAt = nil + if err := s.repo.ScrapeTask.Update(context.Background(), t); err != nil && s.log != nil { + s.log.Warn("requeue claimed scrape task failed", zap.Error(err), zap.String("id", t.ID)) + } +} + func (s *ScraperService) processScrapeTask(ctx context.Context, task *model.ScrapeTask) { media, err := s.repo.Media.FindByID(ctx, task.MediaID) if err != nil || media == nil { @@ -220,8 +239,22 @@ func (s *ScraperService) EnqueueLibrary(ctx context.Context, libraryID string, o return 0, nil } + // 去重:排除已有 pending/running 任务的媒体,防止"先单集入队再点 + // 整库刮削"产生重复任务并被并发双刮(同一 Media 行被并发写两次)。 + mediaIDs := make([]string, 0, len(rows)) + for _, m := range rows { + mediaIDs = append(mediaIDs, m.ID) + } + activeByMedia, err := s.repo.ScrapeTask.FindActiveByMediaIDs(ctx, mediaIDs) + if err != nil { + activeByMedia = nil // 去重查询失败不阻塞入队,仅退化为不去重 + } + tasks := make([]model.ScrapeTask, 0, len(rows)) for _, m := range rows { + if activeByMedia != nil && activeByMedia[m.ID] { + continue + } tasks = append(tasks, model.ScrapeTask{ MediaID: m.ID, LibraryID: lib.ID, @@ -310,7 +343,16 @@ func (s *ScraperService) RetryScrapeTask(ctx context.Context, id string) error { return errors.New("刮削任务不存在") } if task.Status != model.ScrapeTaskFailed && task.Status != model.ScrapeTaskCanceled { - return errors.New("只有失败或已取消的任务可以重试") + // running 任务仅在其长时间无进展(>1h)时允许重试,作为卡死 + // 任务的逃生通道;正常执行中的任务仍拒绝重试以防双跑。 + stale := task.Status == model.ScrapeTaskRunning && + (task.StartedAt == nil || time.Since(*task.StartedAt) > time.Hour) + if !stale { + return errors.New("只有失败或已取消的任务可以重试") + } + if s.log != nil { + s.log.Warn("retrying stuck running scrape task", zap.String("id", id)) + } } task.Status = model.ScrapeTaskPending task.Error = "" diff --git a/internal/service/strm_115_oauth.go b/internal/service/strm_115_oauth.go index 435f9a3..78f0abd 100644 --- a/internal/service/strm_115_oauth.go +++ b/internal/service/strm_115_oauth.go @@ -14,6 +14,7 @@ import ( "context" "errors" "fmt" + "strconv" "strings" "time" @@ -371,6 +372,16 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) { if err != nil || cfg["access_token"] == "" || cfg["refresh_token"] == "" { continue } + // 临期才刷:115 refresh_token 是一次性轮转,无条件周期刷新会与 + // 运行中任务的自动刷新互相作废对方的凭据。24h 内刷新过(含运行 + // 中回调落库)就跳过;access_token 有效期远长于 24h。 + if last := strings.TrimSpace(cfg["token_refreshed_at"]); last != "" { + if ts, perr := strconv.ParseInt(last, 10, 64); perr == nil { + if time.Since(time.Unix(ts, 0)) < 24*time.Hour { + continue + } + } + } client := cloud115.NewOpenClient(cfg["app_id"], cfg["access_token"], cfg["refresh_token"]) token, err := client.RefreshToken(cfg["refresh_token"]) if err != nil { @@ -395,6 +406,7 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) { } cfg["access_token"] = s.crypto.Encrypt(token.AccessToken) cfg["refresh_token"] = s.crypto.Encrypt(token.RefreshToken) + cfg["token_refreshed_at"] = strconv.FormatInt(time.Now().Unix(), 10) enc, err := s.strmAccountConfigJSON(cfg, false) if err != nil { continue @@ -410,6 +422,37 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) { } } +// persist115Tokens 把运行中任务自动刷新得到的新令牌加密写回账号配置, +// 并记录刷新时间供定时刷新线程做临期判断。 +func (s *StrmService) persist115Tokens(accountID, accessToken, refreshToken string) { + ctx := context.Background() + acct, err := s.repo.StrmAccount.FindByID(ctx, accountID) + if err != nil || acct == nil { + return + } + cfg, err := s.strmAccountConfig(acct) + if err != nil { + return + } + cfg["access_token"] = s.crypto.Encrypt(accessToken) + cfg["refresh_token"] = s.crypto.Encrypt(refreshToken) + cfg["token_refreshed_at"] = strconv.FormatInt(time.Now().Unix(), 10) + enc, err := s.strmAccountConfigJSON(cfg, false) + if err != nil { + return + } + acct.Config = enc + now := time.Now() + acct.LastTestAt = &now + acct.LastTestResult = "ok" + acct.LastTestOK = true + if err := s.repo.StrmAccount.Update(ctx, acct); err != nil { + s.log.Warn("persist refreshed 115 token failed", zap.Error(err), zap.String("account", acct.Name)) + return + } + s.log.Info("115 token refreshed and persisted", zap.String("account", acct.Name)) +} + // sync115RelayKey 把设置里的中继密钥同步给 cloud115(启动与设置保存时调用)。 func (s *StrmService) sync115RelayKey(ctx context.Context) { cloud115.RelayEncryptionKey = s.strmSetting(ctx, Strm115RelayKeySetting) diff --git a/internal/service/strm_queue.go b/internal/service/strm_queue.go index 4fd51ec..9d8b30d 100644 --- a/internal/service/strm_queue.go +++ b/internal/service/strm_queue.go @@ -67,9 +67,17 @@ func (s *StrmService) downloadWorker(ctx context.Context) { return } defer s.releaseDownloadSlot(task.Provider) - // 单个任务 panic 不应拖垮整个下载 worker。 + // 单个任务 panic 不应拖垮整个下载 worker,且 panic 时任务 + // 会永远停在 running:兜底走失败重试路径。 + completed := false helper.Run(s.log, "strm.downloadTask", func() { + defer func() { + if !completed { + s.downloadTaskFailWithRetry(task, "任务执行异常中断") + } + }() s.processDownloadTask(ctx, task) + completed = true }) }(i) } @@ -78,9 +86,19 @@ func (s *StrmService) downloadWorker(ctx context.Context) { } // requeueDownloadTask 把已认领但未实际执行的任务退回 pending,避免长期停留在 running。 +// 退回时必须设置 NextTryAt(WAF 冷却剩余时间):claim 只过滤 next_try_at +// 已过期的任务,不设会让同一批任务被立刻再认领,形成 claim/requeue +// 热循环(占用 SQLite 写锁并饿死上传队列)。 func (s *StrmService) requeueDownloadTask(task *model.StrmDownloadTask) { task.Status = model.StrmTaskPending task.StartedAt = nil + task.NextTryAt = nil + if task.Provider == model.StrmProvider115 { + if left := s.wafCooldownLeft(); left > 0 { + next := time.Now().Add(left) + task.NextTryAt = &next + } + } if err := s.repo.StrmDownload.Update(context.Background(), task); err != nil { s.log.Warn("requeue strm download task failed", zap.Error(err), zap.String("id", task.ID)) } @@ -97,8 +115,16 @@ func (s *StrmService) processDownloadTask(ctx context.Context, task *model.StrmD task.Status = status task.Error = message task.FinishedAt = &now - if err := s.repo.StrmDownload.Update(context.Background(), task); err != nil { + // 条件化收尾:用户取消会直接把 running 改为 canceled,无条件 + // Update 会把已取消任务覆盖回 done。 + if ok, err := s.repo.StrmDownload.UpdateIfRunning(context.Background(), task.ID, map[string]any{ + "status": status, + "error": message, + "finished_at": &now, + }); err != nil { s.log.Warn("update strm download task failed", zap.Error(err)) + } else if !ok { + s.log.Info("strm download task already closed elsewhere", zap.String("id", task.ID)) } } acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID) @@ -151,7 +177,19 @@ func (s *StrmService) uploadWorker(ctx context.Context) { continue } for i := range tasks { - s.processUploadTask(ctx, &tasks[i]) + t := &tasks[i] + // 与下载侧一致:单任务 panic 不损失 worker 线程,且兜底走 + // 失败重试路径(否则任务永久 running)。 + completed := false + helper.Run(s.log, "strm.uploadTask", func() { + defer func() { + if !completed { + s.uploadTaskFailWithRetry(t, "任务执行异常中断") + } + }() + s.processUploadTask(ctx, t) + completed = true + }) } } } @@ -162,8 +200,15 @@ func (s *StrmService) processUploadTask(ctx context.Context, task *model.StrmUpl task.Status = status task.Error = message task.FinishedAt = &now - if err := s.repo.StrmUpload.Update(context.Background(), task); err != nil { + // 条件化收尾:与下载侧一致,防止覆盖已取消任务。 + if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{ + "status": status, + "error": message, + "finished_at": &now, + }); err != nil { s.log.Warn("update strm upload task failed", zap.Error(err)) + } else if !ok { + s.log.Info("strm upload task already closed elsewhere", zap.String("id", task.ID)) } } if task.Provider == model.StrmProvider115 { @@ -213,8 +258,15 @@ func (s *StrmService) processUpload115(ctx context.Context, task *model.StrmUplo task.Status = status task.Error = message task.FinishedAt = &now - if err := s.repo.StrmUpload.Update(context.Background(), task); err != nil { + // 条件化收尾:与下载侧一致,防止覆盖已取消任务。 + if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{ + "status": status, + "error": message, + "finished_at": &now, + }); err != nil { s.log.Warn("update strm upload task failed", zap.Error(err)) + } else if !ok { + s.log.Info("strm upload task already closed elsewhere", zap.String("id", task.ID)) } } acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID) @@ -253,7 +305,19 @@ func (s *StrmService) downloadTaskFailWithRetry(task *model.StrmDownloadTask, me if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) { return } - _ = s.repo.StrmDownload.Update(context.Background(), task) + // 条件化写入:任务已被取消(DB 中不再是 running)时不得覆盖回 pending, + // 否则用户刚取消的任务会“复活”并自动重试。 + if ok, err := s.repo.StrmDownload.UpdateIfRunning(context.Background(), task.ID, map[string]any{ + "status": task.Status, + "error": task.Error, + "retry_count": task.RetryCount, + "next_try_at": task.NextTryAt, + "finished_at": task.FinishedAt, + }); err != nil { + s.log.Warn("fail strm download task failed", zap.Error(err), zap.String("id", task.ID)) + } else if !ok { + s.log.Info("strm download task already closed elsewhere, skip retry overwrite", zap.String("id", task.ID)) + } } // uploadTaskFailWithRetry 上传失败任务按退避重试,超过上限标记 failed。 @@ -261,7 +325,17 @@ func (s *StrmService) uploadTaskFailWithRetry(task *model.StrmUploadTask, messag if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) { return } - _ = s.repo.StrmUpload.Update(context.Background(), task) + if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{ + "status": task.Status, + "error": task.Error, + "retry_count": task.RetryCount, + "next_try_at": task.NextTryAt, + "finished_at": task.FinishedAt, + }); err != nil { + s.log.Warn("fail strm upload task failed", zap.Error(err), zap.String("id", task.ID)) + } else if !ok { + s.log.Info("strm upload task already closed elsewhere, skip retry overwrite", zap.String("id", task.ID)) + } } // retryTask 失败状态机:重试次数不足则回 pending 并设置退避时间,否则 failed。 diff --git a/internal/service/strm_service.go b/internal/service/strm_service.go index edac44f..a1b456f 100644 --- a/internal/service/strm_service.go +++ b/internal/service/strm_service.go @@ -28,6 +28,7 @@ import ( "github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/repository" "github.com/truewhile/MeBox/internal/service/cloud" + "github.com/truewhile/MeBox/internal/service/cloud115" ) // strm 全局设置键(存于 Setting 表,strm.* 前缀)。 @@ -169,6 +170,11 @@ func NewStrmService(cfg *config.Config, log *zap.Logger, repos *repository.Conta // Start 启动下载/上传队列 worker、定时同步巡检、115 token 刷新与队列清理。 func (s *StrmService) Start(ctx context.Context) { + // baseCtx 挂到服务生命周期 ctx 上(Start 由启动流程传入 stopCtx): + // 此前硬编码 context.Background(),Stop() 关 stopCh 后 worker 会退出, + // 但进行中的全量同步(可能持续数小时)完全不受停机控制,优雅停机 + // 窗口内仍在批量写库/写盘。 + s.baseCtx = ctx s.sync115RelayKey(ctx) s.recoverInterruptedSyncs(ctx) downloadThreads := s.strmIntSetting(ctx, StrmSettingDownloadThreads, 6) @@ -192,6 +198,18 @@ func (s *StrmService) Start(ctx context.Context) { for i := 0; i < uploadThreads; i++ { helper.Go(s.log, "strm.uploadWorker", func() { s.uploadWorker(ctx) }) } + // 队列任务自愈:进程崩溃/停机遗留的 running 任务重置为 pending, + // 否则永久卡死并会通过 GetActiveLocalPathMap 阻塞该文件的重复下载。 + if n, err := s.repo.StrmDownload.ResetRunningToPending(ctx); err == nil && n > 0 { + s.log.Warn("strm download tasks reset from running to pending after restart", zap.Int64("count", n)) + } else if err != nil { + s.log.Warn("reset running strm download tasks failed", zap.Error(err)) + } + if n, err := s.repo.StrmUpload.ResetRunningToPending(ctx); err == nil && n > 0 { + s.log.Warn("strm upload tasks reset from running to pending after restart", zap.Int64("count", n)) + } else if err != nil { + s.log.Warn("reset running strm upload tasks failed", zap.Error(err)) + } helper.Go(s.log, "strm.cronLoop", func() { s.cronLoop(ctx) }) helper.Go(s.log, "strm.queueCleanupLoop", func() { s.queueCleanupLoop(ctx) }) helper.Go(s.log, "strm.refresh115TokensLoop", func() { s.refresh115TokensLoop(ctx) }) @@ -457,6 +475,15 @@ func (s *StrmService) providerFor(ctx context.Context, acct *model.StrmAccount) if err != nil { return nil, err } + // 115 开放平台:运行中自动刷新得到的新令牌必须落库。否则长任务 + // 里的新 token 只存在于内存,定时刷新线程又用 DB 里的旧 + // refresh_token 再刷(一次性轮转),两者互相作废,最终把有效账号 + // 标成“授权已失效”。 + if oc, ok := provider.(interface{ OpenClient() *cloud115.OpenClient }); ok { + oc.OpenClient().OnTokenRefreshed = func(accessToken, refreshToken string) { + s.persist115Tokens(acct.ID, accessToken, refreshToken) + } + } return provider, nil } diff --git a/internal/service/strm_sync.go b/internal/service/strm_sync.go index e8bcbcf..73ac60a 100644 --- a/internal/service/strm_sync.go +++ b/internal/service/strm_sync.go @@ -6,6 +6,7 @@ import ( "context" "errors" "fmt" + "reflect" "net/url" "os" "path/filepath" @@ -97,7 +98,8 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ... StartedAt: &now, } if err := s.repo.StrmSyncRecord.Create(ctx, rec); err != nil { - s.clearRunning(pathID) + s.clearRunning(pathID, cancel) + cancel() return err } status := model.StrmSyncRecordRunning @@ -106,7 +108,7 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ... p.LastSyncMessage = "同步进行中" _ = s.repo.StrmSyncPath.Update(ctx, p) - helper.Go(s.log, "strm.sync", func() { s.runSync(runCtx, p, rec) }) + helper.Go(s.log, "strm.sync", func() { s.runSync(runCtx, p, rec, cancel) }) return nil } @@ -114,11 +116,10 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ... func (s *StrmService) CancelSync(ctx context.Context, pathID string) error { s.mu.Lock() cancel, exists := s.running[pathID] - if exists { - delete(s.running, pathID) - } s.mu.Unlock() + // 不在这里预删 running 标记:runSync 退出时的 clearRunning 会按 + // cancel 身份校验后删除,避免旧同步收尾误删新同步的标记。 if exists && cancel != nil { cancel() } @@ -142,12 +143,24 @@ func (s *StrmService) IsSyncRunning(pathID string) bool { return exists } -func (s *StrmService) clearRunning(pathID string) { +// clearRunning 清除同步的运行标记;仅当 map 中登记的 cancel 与本次同步 +// 一致时才删除,防止慢收尾的旧同步把随后启动的新同步标记误删掉。 +func (s *StrmService) clearRunning(pathID string, cancel context.CancelFunc) { s.mu.Lock() - delete(s.running, pathID) + if cur, ok := s.running[pathID]; ok { + if cancel == nil || cur == nil || sameCancelFunc(cur, cancel) { + delete(s.running, pathID) + } + } s.mu.Unlock() } +// sameCancelFunc 比较两个 cancel 是否为同一实例(每次 WithCancel 返回 +// 独立闭包,函数指针即身份)。约定 running 表只登记 StartSync 的 cancel。 +func sameCancelFunc(a, b context.CancelFunc) bool { + return reflect.ValueOf(a).Pointer() == reflect.ValueOf(b).Pointer() +} + // ListRemoteDir 列出网盘账号某目录下的条目(供前端目录选择器使用)。 func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string) ([]cloud.FileEntry, error) { acct, err := s.repo.StrmAccount.FindByID(ctx, accountID) @@ -169,8 +182,9 @@ func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string) } // runSync 执行同步主体;结束时更新记录与目录状态。 -func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord) { - defer s.clearRunning(p.ID) +func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord, cancel context.CancelFunc) { + defer s.clearRunning(p.ID, cancel) + defer cancel() cfg, err := s.strmEffectiveConfig(ctx, p) if err != nil { @@ -334,24 +348,34 @@ func (st *strmSyncState) walkRemote() error { ctx, cancel := context.WithCancel(st.ctx) defer cancel() - queue := make(chan dirTask, 512) - var pending atomic.Int64 + // 工作队列用「互斥锁 + 条件变量 + 动态 slice」实现,而不是有界 + // channel:有界缓冲下所有 worker 可能同时阻塞在发送上、无人接收, + // closer 又在等 pending 归零,形成永久死锁。push 永不阻塞即可保证 + // 有进度就一定有推进。 + // pending 计数 = 尚未处理完的任务数(在 work 里或正在被 List)。 + var ( + walkMu sync.Mutex + walkCond = sync.NewCond(&walkMu) + work []dirTask + pending int + ) + push := func(t dirTask) { + walkMu.Lock() + work = append(work, t) + pending++ + walkCond.Signal() + walkMu.Unlock() + } + // ctx 取消时唤醒所有等待中的 worker 让其退出。 + go func() { + <-ctx.Done() + walkMu.Lock() + walkCond.Broadcast() + walkMu.Unlock() + }() // 根目录入队 - pending.Add(1) - queue <- dirTask{id: root, rel: ""} - - // 当队列中所有目录都被消费(pending 归零)或出错时关闭 channel, - // 让 worker 全部退出。 - go func() { - for { - if ctx.Err() != nil || pending.Load() == 0 { - close(queue) - return - } - time.Sleep(10 * time.Millisecond) - } - }() + push(dirTask{id: root, rel: ""}) var ( wg sync.WaitGroup @@ -362,11 +386,27 @@ func (st *strmSyncState) walkRemote() error { wg.Add(1) go func() { defer wg.Done() - // worker 解析远端响应 panic 时取消整个同步,让 closer 与其余 - // worker 正常收尾,避免队列与 pending 计数卡死;正常退出不取消。 + // worker 解析远端响应 panic 时取消整个同步,让其余 worker + // 正常收尾;正常退出不取消。 if err := helper.Recover(st.s.log, "strm.sync.walkRemote", func() error { - for task := range queue { + for { + walkMu.Lock() + for len(work) == 0 { + if ctx.Err() != nil || pending == 0 { + walkMu.Unlock() + return nil + } + walkCond.Wait() + } + task := work[0] + work = work[1:] + walkMu.Unlock() + if ctx.Err() != nil { + walkMu.Lock() + pending-- + walkCond.Broadcast() + walkMu.Unlock() return nil } entries, err := st.provider.List(ctx, task.id) @@ -376,6 +416,10 @@ func (st *strmSyncState) walkRemote() error { firstErr = fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err) } errMu.Unlock() + walkMu.Lock() + pending-- + walkCond.Broadcast() + walkMu.Unlock() cancel() return nil } @@ -386,19 +430,18 @@ func (st *strmSyncState) walkRemote() error { rel = task.rel + "/" + cleanName } if entry.IsDir { - pending.Add(1) - select { - case queue <- dirTask{id: entry.ID, rel: rel}: - case <-ctx.Done(): - pending.Add(-1) - } + push(dirTask{id: entry.ID, rel: rel}) } else { st.processRemoteFile(entry, rel) } } - pending.Add(-1) + walkMu.Lock() + pending-- + if pending == 0 { + walkCond.Broadcast() + } + walkMu.Unlock() } - return nil }); err != nil { cancel() } @@ -1329,6 +1372,10 @@ func (st *strmSyncState) updateSyncMessage(msg string) { func (s *StrmService) cronLoop(ctx context.Context) { ticker := time.NewTicker(60 * time.Second) defer ticker.Stop() + // 记录上次检查到的分钟:一轮循环若被慢操作拖过 60s(远端 List 慢、 + // 串行 StartSync、DB 忙),ticker 会丢掉中间的 tick,命中排程的分钟 + // 若只按"当前分钟相等"判定就会被静默跳过。逐分钟回放补触。 + last := time.Now().Truncate(time.Minute) for { select { case <-ctx.Done(): @@ -1336,8 +1383,18 @@ func (s *StrmService) cronLoop(ctx context.Context) { case <-s.stopCh: return case now := <-ticker.C: + now = now.Truncate(time.Minute) paths, err := s.repo.StrmSyncPath.List(ctx) if err != nil { + last = now + continue + } + due := make([]time.Time, 0, 2) + for m := last.Add(time.Minute); !m.After(now); m = m.Add(time.Minute) { + due = append(due, m) + } + last = now + if len(due) == 0 { continue } for i := range paths { @@ -1345,7 +1402,14 @@ func (s *StrmService) cronLoop(ctx context.Context) { if !p.Enabled || !p.EnableCron || strings.TrimSpace(p.Cron) == "" { continue } - if !cronMatches(p.Cron, now) { + matched := false + for _, m := range due { + if cronMatches(p.Cron, m) { + matched = true + break + } + } + if !matched { continue } s.mu.Lock() diff --git a/internal/service/watcher.go b/internal/service/watcher.go index 77958cf..6556427 100644 --- a/internal/service/watcher.go +++ b/internal/service/watcher.go @@ -279,12 +279,29 @@ func (w *WatcherService) process(ctx context.Context, d duePath) { if removed, derr := w.scanner.RemovePath(ctx, d.path); derr != nil { w.log.Warn("watcher remove failed", zap.String("path", d.path), zap.Error(derr)) } else if removed > 0 { - w.log.Info("watcher removed media", zap.String("path", d.path)) + w.log.Info("watcher removed media", zap.String("path", d.path), zap.Int64("count", removed)) } return } if fi.IsDir() { - return // directory events only matter for registering new watches + // 目录事件(新建 / 重命名进入):注册递归监听之外,还要对子树 + // 做一次增量 ingest——否则目录改名后新路径下的文件永远不会入库, + // 旧路径记录已由消失侧的 RemovePath 子树删除。 + w.mu.Lock() + w.watchDirRecursive(d.path, d.libraryID) + w.mu.Unlock() + _ = filepath.WalkDir(d.path, func(p string, entry os.DirEntry, werr error) error { + if werr != nil || entry.IsDir() { + return nil + } + if added, ierr := w.scanner.IngestPath(ctx, d.libraryID, p); ierr != nil { + w.log.Warn("watcher dir ingest failed", zap.String("path", p), zap.Error(ierr)) + } else if added { + w.log.Info("watcher ingested media", zap.String("path", p)) + } + return nil + }) + return } if added, ierr := w.scanner.IngestPath(ctx, d.libraryID, d.path); ierr != nil { w.log.Warn("watcher ingest failed", zap.String("path", d.path), zap.Error(ierr)) diff --git a/internal/service/ws_hub.go b/internal/service/ws_hub.go index 39b49b9..3ad8c57 100644 --- a/internal/service/ws_hub.go +++ b/internal/service/ws_hub.go @@ -92,6 +92,13 @@ func (h *Hub) Subscribe(id string, topics []string) *Subscriber { sub.topics[t] = struct{}{} } h.mu.Lock() + // Stop() 会把 subs 置 nil:停机窗口内仍在握手的 WS 连接若在此写入 + // nil map 会直接 panic。已关闭的 hub 返回一个立刻关闭的空订阅者。 + if h.subs == nil { + h.mu.Unlock() + close(sub.Out) + return sub + } h.subs[id] = sub h.mu.Unlock() return sub diff --git a/web/src/components/LayoutHeaderSections.tsx b/web/src/components/LayoutHeaderSections.tsx index a245145..df82979 100644 --- a/web/src/components/LayoutHeaderSections.tsx +++ b/web/src/components/LayoutHeaderSections.tsx @@ -114,26 +114,31 @@ function LayoutHeaderSearch() { const [loading, setLoading] = useState(false) const [results, setResults] = useState([]) const containerRef = useRef(null) + // 递增序号守卫:快速连续输入时丢弃过期响应 + const searchSeqRef = useRef(0) const navigate = useNavigate() useEffect(() => { const trimmed = query.trim() if (!trimmed) { + searchSeqRef.current += 1 setResults([]) setLoading(false) return } + const seq = ++searchSeqRef.current setLoading(true) const timer = setTimeout(async () => { try { const res = await mediaAPI.search(trimmed, 8) + if (seq !== searchSeqRef.current) return setResults(res.items || []) setIsOpen(true) } catch { - setResults([]) + // 请求失败时保留旧结果,避免网络抖动清空下拉 } finally { - setLoading(false) + if (seq === searchSeqRef.current) setLoading(false) } }, 250) diff --git a/web/src/hooks/useSSE.ts b/web/src/hooks/useSSE.ts deleted file mode 100644 index 16cfb22..0000000 --- a/web/src/hooks/useSSE.ts +++ /dev/null @@ -1,123 +0,0 @@ -import { useEffect, useRef, useCallback } from 'react' - -import { useAuthStore } from '../stores/auth' -import type { SSEEvent } from '../types' - -type SSEEventHandler = (event: SSEEvent) => void - -/** - * useSSE hook - 管理 Server-Sent Events 连接 - * - * @param onEvent - 事件处理函数 - * @param options - 配置选项 - * - * @example - * ```tsx - * function MyComponent() { - * const { connect, disconnect } = useSSE((event) => { - * if (event.type === 'scan') { - * updateScanProgress(event.payload) - * } - * }) - * - * useEffect(() => { - * connect() - * return () => disconnect() - * }, []) - * - * return
SSE Demo
- * } - * ``` - */ -export function useSSE( - onEvent: SSEEventHandler, - options: { autoConnect?: boolean } = {} -) { - const { autoConnect = true } = options - const onEventRef = useRef(onEvent) - onEventRef.current = onEvent - - const eventSourceRef = useRef(null) - const reconnectTimeoutRef = useRef | null>(null) - const isConnectedRef = useRef(false) - const reconnectAttemptsRef = useRef(0) - const maxReconnectAttempts = 5 - - const connect = useCallback(() => { - // 如果已有连接,先断开 - if (eventSourceRef.current) { - eventSourceRef.current.close() - } - - const token = useAuthStore.getState().token - if (!token) { - console.warn('Cannot connect to SSE: No auth token') - return - } - - const url = `/api/events?token=${encodeURIComponent(token)}` - const eventSource = new EventSource(url) - eventSourceRef.current = eventSource - - eventSource.onopen = () => { - isConnectedRef.current = true - reconnectAttemptsRef.current = 0 - } - - eventSource.onmessage = (event) => { - try { - const data = JSON.parse(event.data) as SSEEvent - onEventRef.current(data) - } catch (err) { - console.error('Failed to parse SSE event:', err) - } - } - - eventSource.onerror = () => { - isConnectedRef.current = false - eventSource.close() - - // 尝试重连 - if (reconnectAttemptsRef.current < maxReconnectAttempts) { - const delay = Math.min(1000 * Math.pow(2, reconnectAttemptsRef.current), 30000) - reconnectAttemptsRef.current++ - reconnectTimeoutRef.current = setTimeout(connect, delay) - } else { - console.error('SSE connection failed after max attempts') - } - } - - }, []) - - const disconnect = useCallback(() => { - if (reconnectTimeoutRef.current) { - clearTimeout(reconnectTimeoutRef.current) - reconnectTimeoutRef.current = null - } - if (eventSourceRef.current) { - eventSourceRef.current.close() - eventSourceRef.current = null - } - isConnectedRef.current = false - reconnectAttemptsRef.current = 0 - }, []) - - const isConnected = useCallback(() => { - return isConnectedRef.current - }, []) - - useEffect(() => { - if (autoConnect) { - connect() - } - return () => { - disconnect() - } - }, [autoConnect, connect, disconnect]) - - return { - connect, - disconnect, - isConnected, - } -} diff --git a/web/src/hooks/useWebSocket.ts b/web/src/hooks/useWebSocket.ts index 5401062..ba8edca 100644 --- a/web/src/hooks/useWebSocket.ts +++ b/web/src/hooks/useWebSocket.ts @@ -2,12 +2,15 @@ import { useEffect, useRef } from 'react' import { useAuthStore } from '../stores/auth' -const MAX_RECONNECT_ATTEMPTS = 5 +// 前 5 次沿用原快速退避间隔;之后进入 60s 慢速重试并不再放弃, +// 服务重启或网络恢复后仍能自动重连(清理函数可随时取消定时器)。 +const FAST_RECONNECT_ATTEMPTS = 5 +const SLOW_RECONNECT_INTERVAL = 60_000 // useWebSocket opens a single connection to /api/ws and dispatches every // message to the supplied handler. Auto-reconnects with back-off while the -// auth token is present, but stops after repeated failures so an expired token -// cannot create an endless /api/ws 401 loop. +// auth token is present; after the fast retries are exhausted it keeps a +// slow 60s retry loop instead of giving up permanently. export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) { const ref = useRef(null) const token = useAuthStore((s) => s.token) @@ -22,7 +25,6 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) const open = () => { if (closed) return - if (reconnectAttempts >= MAX_RECONNECT_ATTEMPTS) return const proto = window.location.protocol === 'https:' ? 'wss:' : 'ws:' const url = `${proto}//${window.location.host}/api/ws?token=${encodeURIComponent(token)}` const ws = new WebSocket(url) @@ -43,7 +45,12 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) ws.onclose = () => { if (closed) return reconnectAttempts += 1 - const delay = Math.min(3_000 * reconnectAttempts, 30_000) + // 快速阶段保持原有线性退避,之后固定 60s 慢速重试; + // timer 始终只有一个在途,cleanup 时统一清除,不会堆积。 + const delay = + reconnectAttempts <= FAST_RECONNECT_ATTEMPTS + ? Math.min(3_000 * reconnectAttempts, 30_000) + : SLOW_RECONNECT_INTERVAL timer = window.setTimeout(open, delay) } } diff --git a/web/src/pages/AdminLibraryPanel.tsx b/web/src/pages/AdminLibraryPanel.tsx index d7017da..d5f0bf6 100644 --- a/web/src/pages/AdminLibraryPanel.tsx +++ b/web/src/pages/AdminLibraryPanel.tsx @@ -14,6 +14,7 @@ export function AdminLibraryPanel() { coverURL={createForm.coverURL} roots={createForm.roots} createPerSubfolder={createForm.createPerSubfolder} + creating={createForm.creating} onNameChange={createForm.setName} onTypeChange={createForm.setType} onCoverURLChange={createForm.setCoverURL} diff --git a/web/src/pages/AdminLibraryPanelSections.tsx b/web/src/pages/AdminLibraryPanelSections.tsx index b9b939d..89915f1 100644 --- a/web/src/pages/AdminLibraryPanelSections.tsx +++ b/web/src/pages/AdminLibraryPanelSections.tsx @@ -1,5 +1,5 @@ import { FormEvent, useState } from 'react' -import { Folder, Plus, Trash2 } from 'lucide-react' +import { Folder, Loader2, Plus, Trash2 } from 'lucide-react' import { LocalDirBrowserDialog } from '../components/LocalDirBrowserDialog' import type { RootDraft } from './adminLibraryPanelModel' @@ -10,6 +10,7 @@ type CreateFormProps = { coverURL: string roots: RootDraft[] createPerSubfolder: boolean + creating: boolean onNameChange: (value: string) => void onTypeChange: (value: string) => void onCoverURLChange: (value: string) => void @@ -26,6 +27,7 @@ export function AdminLibraryCreateForm({ coverURL, roots, createPerSubfolder, + creating, onNameChange, onTypeChange, onCoverURLChange, @@ -108,8 +110,9 @@ export function AdminLibraryCreateForm({ 批处理模式:仅取上方第一个路径作为父级目录,会为其中每个子文件夹分别创建媒体库,可自选类型用于整体推断。

)} - diff --git a/web/src/pages/AdminUsersForm.tsx b/web/src/pages/AdminUsersForm.tsx index 6c1e090..fe5b9f5 100644 --- a/web/src/pages/AdminUsersForm.tsx +++ b/web/src/pages/AdminUsersForm.tsx @@ -1,11 +1,12 @@ import { FormEvent } from 'react' -import { Plus, Save } from 'lucide-react' +import { Loader2, Plus, Save } from 'lucide-react' type AdminUsersFormProps = { usersCount: number maxUsers: number maxUsersDraft: string savingLimit: boolean + creating: boolean username: string password: string userLimitReached: boolean @@ -21,6 +22,7 @@ export function AdminUsersForm({ maxUsers, maxUsersDraft, savingLimit, + creating, username, password, userLimitReached, @@ -91,9 +93,13 @@ export function AdminUsersForm({ onChange={(e) => onPasswordChange(e.target.value)} disabled={userLimitReached} /> - ) diff --git a/web/src/pages/AdminUsersPanel.tsx b/web/src/pages/AdminUsersPanel.tsx index d4c6d95..a07fc87 100644 --- a/web/src/pages/AdminUsersPanel.tsx +++ b/web/src/pages/AdminUsersPanel.tsx @@ -16,6 +16,7 @@ export function AdminUsersPanel() { const [maxUsers, setMaxUsers] = useState(DEFAULT_MAX_USERS) const [maxUsersDraft, setMaxUsersDraft] = useState(String(DEFAULT_MAX_USERS)) const [savingLimit, setSavingLimit] = useState(false) + const [creating, setCreating] = useState(false) const [username, setUsername] = useState('') const [password, setPassword] = useState('') const [editingID, setEditingID] = useState(null) @@ -69,18 +70,22 @@ export function AdminUsersPanel() { const handleCreate = async (e: FormEvent) => { e.preventDefault() + if (creating) return + setCreating(true) try { await adminAPI.createUser({ username, password }) toast.success('用户已添加,默认仅允许浏览与播放媒体') setUsername('') setPassword('') - await refresh() } catch (err: unknown) { const msg = userCreateErrorMessage(err) ?? '添加用户失败' toast.error(msg) + } finally { + setCreating(false) } + await refresh().catch(() => undefined) } const startEdit = (u: User) => { @@ -171,6 +176,7 @@ export function AdminUsersPanel() { maxUsers={maxUsers} maxUsersDraft={maxUsersDraft} savingLimit={savingLimit} + creating={creating} username={username} password={password} userLimitReached={userLimitReached} diff --git a/web/src/pages/AdultSettingsPanel.tsx b/web/src/pages/AdultSettingsPanel.tsx index c3a5398..0326f36 100644 --- a/web/src/pages/AdultSettingsPanel.tsx +++ b/web/src/pages/AdultSettingsPanel.tsx @@ -37,6 +37,7 @@ export function AdultSettingsPanel() { const [values, setValues] = useState>({}) const [dirty, setDirty] = useState>(new Set()) const [loading, setLoading] = useState(true) + const [loadError, setLoadError] = useState('') const [saving, setSaving] = useState(false) const [libraries, setLibraries] = useState([]) const [showToken, setShowToken] = useState(false) @@ -79,6 +80,13 @@ export function AdultSettingsPanel() { setValues(idx) setLibraries(libs as Library[]) setDirty(new Set()) + setLoadError('') + } catch (err: unknown) { + // 加载失败时保留错误态,避免把表单默认值误当成已保存的配置 + const msg = + (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '加载成人设置失败' + setLoadError(msg) + toast.error(msg) } finally { setLoading(false) } @@ -196,6 +204,19 @@ export function AdultSettingsPanel() { ) } + if (loadError) { + return ( +
+ +

成人设置加载失败:{loadError}

+

当前展示的并非已保存配置,请重新加载后再修改

+ +
+ ) + } + return (
{/* 1. 全局访问与隔离卡片 */} diff --git a/web/src/pages/PlayerPage.tsx b/web/src/pages/PlayerPage.tsx index c3b163e..d1563dd 100644 --- a/web/src/pages/PlayerPage.tsx +++ b/web/src/pages/PlayerPage.tsx @@ -13,6 +13,7 @@ import type { Media } from '../types' import { getSeriesKey, seriesTitleFromPath } from '../utils/groupSeries' import { isRemoteEmbyID } from '../utils/remoteEmby' import { pickPlayerMode, needsTranscodeForBrowser, isDirectStreamMedia, type PlayerMode } from './playerPageModel' +import { apiErrorMessage } from './StrmManagePage' import { PlayerTopBar } from './PlayerTopBar' import { PlayerVideoStage } from './PlayerVideoStage' import { PlayerDanmakuPanel } from '../components/PlayerDanmakuPanel' @@ -62,6 +63,8 @@ export function PlayerPage() { const [subtitleIndex, setSubtitleIndex] = useState(initialSubtitleIndex) const [hlsUnavailable, setHlsUnavailable] = useState(false) const [playerError, setPlayerError] = useState('') + // 媒体元数据加载失败(404 / 无权限等):舞台区直接展示错误而不是永远「加载中」 + const [loadError, setLoadError] = useState('') // 「客户端直连解码」模式:宿主机不转码,播放器强制 direct play、隐藏 HLS 切换。 const [directOnly, setDirectOnly] = useState(false) const [resumePosition, setResumePosition] = useState(0) @@ -176,6 +179,7 @@ export function PlayerPage() { // 切换视频时重置媒体与弹幕状态,确保新视频自动重新识别并加载弹幕 useEffect(() => { setMedia(null) + setLoadError('') setDanmakuEpisodeId(null) setDanmakuCandidates([]) setDanmakuSearch(null) @@ -184,29 +188,48 @@ export function PlayerPage() { setDanmakuSearching(true) }, [id]) + // 依赖收敛为 mode 参数的字符串值:避免 params 对象引用每次变化都重复拉取元数据 + const modeParam = params.get('mode') as PlayerMode | null + // Load metadata and pick a default mode. useEffect(() => { if (!id) return - mediaAPI.get(id).then((m) => { - setMedia(m) - const isDirect = isDirectStreamMedia(m) - const forced = params.get('mode') as PlayerMode | null - const auto = pickPlayerMode(m) - // 直连解码模式以及 STRM / Emby 挂载等直连媒体,忽略 ?mode=hls,始终 direct play。 - setMode(directOnly || isDirect ? 'direct' : (forced ?? auto)) - setPlayerError('') - }) + let cancelled = false + mediaAPI + .get(id) + .then((m) => { + if (cancelled) return + setMedia(m) + const isDirect = isDirectStreamMedia(m) + const auto = pickPlayerMode(m) + // 直连解码模式以及 STRM / Emby 挂载等直连媒体,忽略 ?mode=hls,始终 direct play。 + setMode(directOnly || isDirect ? 'direct' : (modeParam ?? auto)) + setPlayerError('') + setLoadError('') + }) + .catch((err: unknown) => { + if (cancelled) return + // 404 / 无权限等:给出可见错误提示,避免永久「加载中」 + setLoadError(`无法加载该媒体:${apiErrorMessage(err)}`) + }) subtitlesAPI .list(id) .then((tracks) => { + if (cancelled) return const list = tracks ?? [] setSubs(list) // 记忆的轨道下标可能超出当前媒体的轨道数(不同媒体字幕数量不同), // 越界时回退到第一条;无字幕则关闭。 setSubtitleIndex((cur) => (cur >= list.length ? (list.length > 0 ? 0 : -1) : cur)) }) - .catch(() => setSubs([])) - }, [id, params, directOnly]) + .catch(() => { + if (cancelled) return + setSubs([]) + }) + return () => { + cancelled = true + } + }, [id, modeParam, directOnly]) // Wire up the actual