优化,排查项目问题

This commit is contained in:
truewhile
2026-09-05 12:34:17 +08:00
parent 7fa05391e1
commit 1407b9b5c4
85 changed files with 1940 additions and 638 deletions
+16 -5
View File
@@ -130,26 +130,35 @@ func (m *serverManager) Shutdown(ctx context.Context) error {
// desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。 // desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。
// 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。 // 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。
func (m *serverManager) desiredPair() (*tlsPair, error) { 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 return nil, nil
} }
certPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLCert, m.cfg.App.SSLCertPath, "证书") certPEM, err := service.ResolveSSLMaterial(cert, certPath, "证书")
if err != nil { if err != nil {
return nil, err return nil, err
} }
keyPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLKey, m.cfg.App.SSLKeyPath, "私钥") keyPEM, err := service.ResolveSSLMaterial(key, keyPath, "私钥")
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil { if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil {
return nil, err return nil, err
} }
cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) pairCert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
if err != nil { if err != nil {
return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err) return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err)
} }
return &tlsPair{ return &tlsPair{
cert: cert, cert: pairCert,
certPEM: certPEM, certPEM: certPEM,
keyPEM: keyPEM, keyPEM: keyPEM,
version: certPEM + "\x00" + keyPEM, version: certPEM + "\x00" + keyPEM,
@@ -171,6 +180,8 @@ func (m *serverManager) maybeStartAutoReloadLocked() {
// pathBased 是否至少有一侧证书/私钥通过文件路径配置。 // pathBased 是否至少有一侧证书/私钥通过文件路径配置。
func (m *serverManager) pathBased() bool { 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) != "" return strings.TrimSpace(m.cfg.App.SSLCertPath) != "" || strings.TrimSpace(m.cfg.App.SSLKeyPath) != ""
} }
+1 -1
View File
@@ -8,7 +8,7 @@ require (
github.com/gin-contrib/gzip v1.2.6 github.com/gin-contrib/gzip v1.2.6
github.com/gin-gonic/gin v1.12.0 github.com/gin-gonic/gin v1.12.0
github.com/glebarez/sqlite v1.11.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/google/uuid v1.6.0
github.com/gorilla/websocket v1.5.3 github.com/gorilla/websocket v1.5.3
github.com/redis/go-redis/v9 v9.7.0 github.com/redis/go-redis/v9 v9.7.0
+2 -2
View File
@@ -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-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 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM=
github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= 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.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8=
github.com/golang-jwt/jwt/v5 v5.2.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= 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.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 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
+14 -2
View File
@@ -12,6 +12,7 @@ import (
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"github.com/spf13/viper" "github.com/spf13/viper"
) )
@@ -19,6 +20,12 @@ import (
// EnvPrefix 是所有环境变量驱动的覆盖使用的前缀。 // EnvPrefix 是所有环境变量驱动的覆盖使用的前缀。
const EnvPrefix = "MeBox" const EnvPrefix = "MeBox"
// RuntimeMu 保护运行时热更新配置字段的并发读写:ApplyRuntimeSetting 在
// HTTP goroutine 中写字段,serverManager 的证书轮询等后台协程在无锁读取
// 同一批字段。string 是双字结构,无锁并发读写可读到撕裂的 header。
// 写方在 ApplyRuntimeSetting 内 Lock,读方(cmd/server)在轮询处 RLock。
var RuntimeMu sync.RWMutex
// Load 从默认值 / 文件 / 环境读取配置。 // Load 从默认值 / 文件 / 环境读取配置。
// //
// 即使没有文件也始终返回可用的 Config。 // 即使没有文件也始终返回可用的 Config。
@@ -45,8 +52,13 @@ func Load() (*Config, error) {
} }
s := viper.New() s := viper.New()
s.SetConfigFile(filepath.Join("config", e.Name())) s.SetConfigFile(filepath.Join("config", e.Name()))
if err := s.ReadInConfig(); err == nil { if err := s.ReadInConfig(); err != nil {
_ = v.MergeConfigMap(s.AllSettings()) // 分片解析失败不能静默吞掉: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())
} }
} }
} }
+8 -2
View File
@@ -68,8 +68,14 @@ func (c *Config) normalize() error {
return fmt.Errorf("generate jwt secret: %w", err) return fmt.Errorf("generate jwt secret: %w", err)
} }
c.Secrets.JWTSecret = hex.EncodeToString(buf) c.Secrets.JWTSecret = hex.EncodeToString(buf)
_ = os.MkdirAll(c.App.DataDir, 0o750) // 持久化失败(DataDir 只读/权限异常)会导致每次重启重新生成
_ = os.WriteFile(path, []byte(c.Secrets.JWTSecret), 0o600) // 密钥、全部会话静默失效、多实例各持不同 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 return nil
+9 -2
View File
@@ -34,8 +34,15 @@ func SaveDatabaseConfig(dbType, dsn string) error {
return fmt.Errorf("marshal config.yaml: %w", err) return fmt.Errorf("marshal config.yaml: %w", err)
} }
if err := os.WriteFile(configPath, out, 0644); err != nil { // 原子写:临时文件 + rename,避免进程崩溃/断电留下截断的 config.yaml
return fmt.Errorf("write config.yaml: %w", err) // (下次启动会硬失败);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 return nil
} }
+9
View File
@@ -6,6 +6,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"strings" "strings"
"time"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
"go.uber.org/zap" "go.uber.org/zap"
@@ -73,6 +74,14 @@ func configureConnectionPool(db *gorm.DB, cfg *config.Config) error {
if cfg.Database.MaxIdleConns > 0 { if cfg.Database.MaxIdleConns > 0 {
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns) 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 return nil
} }
+1 -1
View File
@@ -9,7 +9,7 @@ const mediaSearchIndexSchemaVersion = 2
func ensureMediaSearchIndex(db *gorm.DB) error { func ensureMediaSearchIndex(db *gorm.DB) error {
if err := ensureMediaSearchMetaTable(db); err != nil { if err := ensureMediaSearchMetaTable(db); err != nil {
return nil return err // meta 表创建失败必须上抛,不能静默掩盖
} }
version := currentMediaSearchIndexVersion(db) version := currentMediaSearchIndexVersion(db)
if version != mediaSearchIndexSchemaVersion { if version != mediaSearchIndexSchemaVersion {
+10 -4
View File
@@ -132,13 +132,19 @@ func ensureEmbyMountsCompatibility(db *gorm.DB) error {
return err return err
} }
} }
// 针对已有数据:如果存在多个 sort_order=0/NULL 的记录,按创建时间顺序赋予稳定递增的序号 // 针对已有数据:只给 sort_order=0/NULL 的行按创建时间补号(从现有
// 最大值之后递增),不能整表重排——此前无条件按 created_at 从 0 重新
// 编号,会把用户自定义的顺序覆盖掉。
var zeroCount int64 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 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 { 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
} }
} }
} }
+31 -16
View File
@@ -50,28 +50,43 @@ func copyModelTables(src, target *gorm.DB, batchSize int) (map[string]int64, int
if modelType.Kind() != reflect.Ptr { if modelType.Kind() != reflect.Ptr {
return tableCounts, totalCopied, fmt.Errorf("model %T is not a pointer", m) return tableCounts, totalCopied, fmt.Errorf("model %T is not a pointer", m)
} }
sliceType := reflect.SliceOf(modelType.Elem()) var primaryKeySet map[string]struct{}
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()
if targetCount > 0 { if targetCount > 0 {
primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns) primaryKeySet, err = targetPrimaryKeySet(target, table, primaryColumns)
if err != nil { if err != nil {
return tableCounts, totalCopied, err return tableCounts, totalCopied, err
} }
filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
} }
if filtered.Len() == 0 { // 分页流式读取:此前整表一次性 Find 进内存,media 表几十万行、
continue // 每行含 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 tableCounts[table] = copiedForTable
totalCopied += copiedForTable totalCopied += copiedForTable
} }
+76 -24
View File
@@ -5,12 +5,21 @@ import (
"fmt" "fmt"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"sync/atomic"
"time"
"gorm.io/gorm" "gorm.io/gorm"
"github.com/truewhile/MeBox/internal/config" "github.com/truewhile/MeBox/internal/config"
) )
// sqliteGateHoldLimit 是写闸持有者的最长合法持有时长。语句级写闸在 SQL
// 执行 panic 时 After 回调不会运行,令牌会泄漏并让后续所有写入永久等锁;
// 超过该时长的持有者按泄漏强制回收(60s 内单条写语句远未到,正常写路径
// 不受影响)。
const sqliteGateHoldLimit = 60 * time.Second
func installSQLiteWriteGate(db *gorm.DB) { func installSQLiteWriteGate(db *gorm.DB) {
if db == nil { if db == nil {
return return
@@ -22,15 +31,18 @@ func installSQLiteWriteGate(db *gorm.DB) {
if tx.Statement != nil && tx.Statement.Context != nil { if tx.Statement != nil && tx.Statement.Context != nil {
ctx = tx.Statement.Context ctx = tx.Statement.Context
} }
if err := gate.Lock(ctx); err != nil { holder, err := gate.Lock(ctx)
if err != nil {
_ = tx.AddError(err) _ = tx.AddError(err)
return return
} }
tx.InstanceSet(lockedKey, struct{}{}) tx.InstanceSet(lockedKey, holder)
} }
unlock := func(tx *gorm.DB) { unlock := func(tx *gorm.DB) {
if _, ok := tx.InstanceGet(lockedKey); ok { if holder, ok := tx.InstanceGet(lockedKey); ok {
gate.Unlock() if h, ok := holder.(*sqliteGateHolder); ok {
gate.Unlock(h)
}
} }
} }
rawLock := func(tx *gorm.DB) { rawLock := func(tx *gorm.DB) {
@@ -64,38 +76,76 @@ func isReadOnlySQL(sql string) bool {
return false return false
} }
// sqliteWriteGate serializes in-process SQLite writes while respecting the // sqliteWriteGate serializes in-process SQLite writes. 所有权令牌(而非裸
// statement context, so request cancellation can break out of a queued write. // 信号量)保证只有持有者本人能释放;持有超时按泄漏自动回收,避免一次
// panic 让进程的 SQLite 写入半永久性瘫痪。
type sqliteWriteGate struct { 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 { 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 { func (g *sqliteWriteGate) Lock(ctx context.Context) (*sqliteGateHolder, error) {
select { g.mu.Lock()
case g.ch <- struct{}{}: defer g.mu.Unlock()
return nil
default:
}
if ctx == nil { if ctx == nil {
ctx = context.Background() ctx = context.Background()
} }
select { // ctx 取消时唤醒等待者(cond 无法感知 ctx,用旁路 goroutine 广播)。
case g.ch <- struct{}{}: if done := ctx.Done(); done != nil {
return nil stop := make(chan struct{})
case <-ctx.Done(): defer close(stop)
return ctx.Err() 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() { func (g *sqliteWriteGate) Unlock(h *sqliteGateHolder) {
select { g.mu.Lock()
case <-g.ch: defer g.mu.Unlock()
default: if h == nil || g.owner != h {
return
} }
g.owner = nil
g.cond.Broadcast()
} }
func buildSQLiteDSN(cfg *config.Config) string { 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. // keep as-is to respect user-provided relative paths.
dbPath = filepath.Clean(dbPath) 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 { if cfg.Database.WALMode {
dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)" dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)"
} }
+20 -4
View File
@@ -10,6 +10,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"go.uber.org/zap" "go.uber.org/zap"
"github.com/truewhile/MeBox/internal/config"
"github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/model"
"github.com/truewhile/MeBox/internal/service" "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)) 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 { switch key {
case "https.enabled": case "https.enabled":
if svc.Cfg.App.HTTPSEnabled { if httpsEnabled {
if _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath); err != nil { if _, err := service.ResolveSSLKeyPair(cert, certPath, keyMaterial, keyPath); err != nil {
return fmt.Errorf("启用 HTTPS 失败:%v", err) 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 { if err := validateSSLMaterialSource(key, value); err != nil {
return err return err
} }
if !svc.Cfg.App.HTTPSEnabled { if !httpsEnabled {
return nil return nil
} }
if !httpsPairReady(svc) { if !httpsPairReady(svc) {
@@ -144,7 +154,13 @@ func validateSSLMaterialSource(key, value string) error {
// httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。 // httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。
func httpsPairReady(svc *service.Container) bool { 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 return err == nil
} }
+39
View File
@@ -3,6 +3,8 @@ package handler
import ( import (
"net/http" "net/http"
"net/url"
"strings"
"github.com/gin-gonic/gin" "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()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return 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 { if err := svc.DLNA.Cast(c.Request.Context(), req.ControlURL, req.MediaURL); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return return
@@ -40,3 +57,25 @@ func dlnaCastHandler(svc *service.Container) gin.HandlerFunc {
c.Status(http.StatusNoContent) 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
}
+16 -13
View File
@@ -132,22 +132,25 @@ func embyMeHandler(svc *service.Container) gin.HandlerFunc {
func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc { func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { 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 { if err == nil && u != nil {
c.JSON(http.StatusOK, u) c.JSON(http.StatusOK, u)
return return
} }
if authUID := embyUserID(c); authUID != "" && authUID != c.Param("userId") { c.JSON(http.StatusOK, embyFallbackUser(uid))
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")))
} }
} }
// embyFallbackUser 是查库失败时的最后兜底(保持客户端可渲染)。
// Policy 必须是最小权限:不声明管理员/删除内容/控制他人等能力,
// 实际权限始终由服务端各路由的校验决定。
func embyFallbackUser(id string) gin.H { func embyFallbackUser(id string) gin.H {
if strings.TrimSpace(id) == "" { if strings.TrimSpace(id) == "" {
id = "mebox-user" id = "mebox-user"
@@ -161,10 +164,10 @@ func embyFallbackUser(id string) gin.H {
"HasConfiguredEasyPassword": false, "HasConfiguredEasyPassword": false,
"EnableAutoLogin": false, "EnableAutoLogin": false,
"Policy": gin.H{ "Policy": gin.H{
"IsAdministrator": true, "IsAdministrator": false,
"EnableContentDeletion": true, "EnableContentDeletion": false,
"EnableRemoteControlOfOtherUsers": true, "EnableRemoteControlOfOtherUsers": false,
"EnableSharedDeviceControl": true, "EnableSharedDeviceControl": false,
"EnableRemoteAccess": true, "EnableRemoteAccess": true,
"EnableAllDevices": true, "EnableAllDevices": true,
"EnableAllChannels": true, "EnableAllChannels": true,
+5 -1
View File
@@ -368,7 +368,11 @@ func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
} }
uid, _ := c.Get("ctx_user_id") uid, _ := c.Get("ctx_user_id")
svc.Audit.Record(c.Request.Context(), toString(uid), "library.delete", id, c.ClientIP(), "") 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) c.Status(http.StatusNoContent)
} }
} }
+36 -3
View File
@@ -2,9 +2,11 @@
package handler package handler
import ( import (
"errors"
"net/http" "net/http"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/truewhile/MeBox/internal/middleware" "github.com/truewhile/MeBox/internal/middleware"
"github.com/truewhile/MeBox/internal/service" "github.com/truewhile/MeBox/internal/service"
@@ -157,6 +159,25 @@ type playlistItemReq struct {
MediaID string `json:"media_id" binding:"required"` 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 { func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
var req playlistItemReq var req playlistItemReq
@@ -164,8 +185,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return return
} }
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
if !ok {
return
}
if err := svc.Playback.AddToPlaylist( 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 { ); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return return
@@ -176,8 +201,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc { func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
if !ok {
return
}
if err := svc.Playback.RemoveFromPlaylist( 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 { ); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return return
@@ -188,8 +217,12 @@ func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
func deletePlaylistHandler(svc *service.Container) gin.HandlerFunc { func deletePlaylistHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
if !ok {
return
}
if err := svc.Playback.DeletePlaylist( if err := svc.Playback.DeletePlaylist(
c.Request.Context(), c.Param("id"), c.Request.Context(), c.Param("id"), uid, isAdmin,
); err != nil { ); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return return
+8 -1
View File
@@ -27,6 +27,9 @@ func reorderPlaylistHandler(svc *service.Container) gin.HandlerFunc {
return return
} }
pid := c.Param("id") pid := c.Param("id")
if _, _, ok := playlistWriteGuard(c, svc, pid); !ok {
return
}
for i, mid := range req.Order { for i, mid := range req.Order {
if err := svc.Repo.DB.WithContext(c.Request.Context()). if err := svc.Repo.DB.WithContext(c.Request.Context()).
Model(&model.PlaylistItem{}). 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). // /playlists/:id/items/:item_id (vs. the existing /:media_id variant).
func deletePlaylistItemByIDHandler(svc *service.Container) gin.HandlerFunc { func deletePlaylistItemByIDHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { 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()). 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 { Delete(&model.PlaylistItem{}).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return return
+21
View File
@@ -14,9 +14,14 @@ import (
) )
// statsUserHandler returns a watch-time summary for one user. // statsUserHandler returns a watch-time summary for one user.
// 观看统计是隐私数据:仅允许本人或管理员查询。
func statsUserHandler(svc *service.Container) gin.HandlerFunc { func statsUserHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
uid := c.Param("id") uid := c.Param("id")
if !statsCallerAllowed(c, uid) {
c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"})
return
}
var watched int64 var watched int64
_ = svc.Repo.DB.Model(&model.PlaybackHistory{}). _ = svc.Repo.DB.Model(&model.PlaybackHistory{}).
Where("user_id = ?", uid). Where("user_id = ?", uid).
@@ -35,8 +40,14 @@ func statsUserHandler(svc *service.Container) gin.HandlerFunc {
} }
// statsTopUsersHandler returns the most active users by play count. // statsTopUsersHandler returns the most active users by play count.
// 全员排行含用户名与精确时长,仅管理员可查。
func statsTopUsersHandler(svc *service.Container) gin.HandlerFunc { func statsTopUsersHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) { 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")) limit, _ := strconv.Atoi(c.DefaultQuery("limit", "10"))
if limit <= 0 || limit > 50 { if limit <= 0 || limit > 50 {
limit = 10 limit = 10
@@ -109,3 +120,13 @@ func statsPlayHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusOK, gin.H{"ok": true}) 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
}
+6 -1
View File
@@ -39,11 +39,16 @@ func listSystemConfigHandler(svc *service.Container) gin.HandlerFunc {
} }
func isSecretKey(k string) bool { 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) { if endsWith(k, suffix) {
return true return true
} }
} }
// 非后缀型敏感键:可触发服务端任意命令的更新命令等。
switch k {
case "system.update.command":
return true
}
return false return false
} }
+17 -4
View File
@@ -9,6 +9,8 @@ package handler
import ( import (
"encoding/json" "encoding/json"
"net/http" "net/http"
"net/url"
"strings"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -21,10 +23,21 @@ import (
var wsUpgrader = websocket.Upgrader{ var wsUpgrader = websocket.Upgrader{
ReadBufferSize: 1024, ReadBufferSize: 1024,
WriteBufferSize: 1024, WriteBufferSize: 1024,
// Allow any origin: the AuthRequired middleware already validated the // 同源校验:浏览器跨站页面虽读不到 ?token=,但可能借 cookie 通道
// JWT before we got here, and we never serve sensitive cross-domain // (extractToken 接受 msgo_access_token cookie)发起跨站 WebSocket
// state through the socket. // 劫持。放行同源与非浏览器客户端(不发 Origin 头的 App/脚本),
CheckOrigin: func(_ *http.Request) bool { return true }, // 拒绝跨站 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 { func wsHandler(svc *service.Container) gin.HandlerFunc {
+13
View File
@@ -51,6 +51,19 @@ func EmbyAuthRequired(secret string) gin.HandlerFunc {
return 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(EmbyCtxUserID, claims.UserID)
c.Set(CtxUserID, claims.UserID) c.Set(CtxUserID, claims.UserID)
c.Set(CtxUserRole, claims.Role) c.Set(CtxUserRole, claims.Role)
+16 -1
View File
@@ -16,6 +16,8 @@ type RateLimiter struct {
window time.Duration window time.Duration
max int max int
requests map[string][]time.Time requests map[string][]time.Time
stop chan struct{}
stopped sync.Once
} }
// NewRateLimiter creates a rate limiter allowing max requests per window // NewRateLimiter creates a rate limiter allowing max requests per window
@@ -25,14 +27,27 @@ func NewRateLimiter(max int, window time.Duration) *RateLimiter {
window: window, window: window,
max: max, max: max,
requests: make(map[string][]time.Time), requests: make(map[string][]time.Time),
stop: make(chan struct{}),
} }
go rl.cleanup() go rl.cleanup()
return rl return rl
} }
// Close 停止后台清理 goroutine:清理循环此前无停止机制,每建一个实例
// 就永久滞留一条 goroutine(测试场景会随实例创建不断累积)。
func (rl *RateLimiter) Close() {
rl.stopped.Do(func() { close(rl.stop) })
}
func (rl *RateLimiter) cleanup() { func (rl *RateLimiter) cleanup() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for { for {
time.Sleep(5 * time.Minute) select {
case <-rl.stop:
return
case <-ticker.C:
}
rl.mu.Lock() rl.mu.Lock()
now := time.Now() now := time.Now()
for ip, times := range rl.requests { for ip, times := range rl.requests {
+15 -8
View File
@@ -5,15 +5,22 @@ import (
"time" "time"
) )
// ApiConfig 存储第三方 API 密钥和配置信息。 // APIConfig 存储第三方 API 密钥和配置信息。
// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储。 // APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储(AES-GCM 密文,
type ApiConfig struct { // base64 后常超 512 字符,因此必须是 text 而非 varchar(512))。
//
// NOTE: 历史上曾有 APIConfig / ApiConfig 两个结构体映射到同一张 api_configs
// 表,AutoMigrate 每次启动互相改列(provider/api_key 长度来回切换),且
// varchar(512) 收窄会让长密文入库后下一次启动迁移直接失败。现已合并为本
// 结构体,字段取两者并集,请勿再拆分。
type APIConfig struct {
Base Base
Provider string `gorm:"size:64;uniqueIndex;not null" json:"provider"` Provider string `gorm:"uniqueIndex;size:64;not null" json:"provider"`
APIKey string `gorm:"size:512" json:"-"` APIKey string `gorm:"type:text" json:"-"` // ciphertext (never serialised)
BaseURL string `gorm:"size:512" json:"base_url,omitempty"` BaseURL string `gorm:"size:512" json:"base_url,omitempty"`
Extra string `gorm:"type:text" json:"extra,omitempty"` Extra string `gorm:"type:text" json:"extra,omitempty"` // free-form JSON
Enabled bool `gorm:"default:true" json:"enabled"` Enabled bool `gorm:"default:true" json:"enabled"`
Description string `gorm:"size:255" json:"description,omitempty"` Description string `gorm:"size:255" json:"description,omitempty"`
LastTestedAt *time.Time `json:"last_tested_at,omitempty"` LastTestedAt *time.Time `json:"last_tested_at,omitempty"`
TestResult string `gorm:"size:32" json:"test_result,omitempty"` TestResult string `gorm:"size:32" json:"test_result,omitempty"`
+6 -21
View File
@@ -1,23 +1,8 @@
package model package model
// APIConfig stores third-party data-source configuration. The api_key // NOTE: 原 APIConfig(provider varchar(32) / api_key text)与 api_config.go
// column is encrypted with AES-GCM (see internal/service/crypto.go) so an // 里的 ApiConfig(provider varchar(64) / api_key varchar(512))映射到同一张
// SQLite leak does not expose third-party credentials. // api_configs 表,AutoMigrate 每次启动互相改列;且 api_key 被收窄成
// // varchar(512) 后,成人区/豆瓣等存的长 AES-GCM Cookie 密文一旦入库,下次
// Provider values mirror the original Python project: // 启动迁移即失败、服务无法启动。两者已合并为 api_config.go 中唯一的
// // APIConfig 结构体(字段取并集),此处不再定义重复模型。
// 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"`
}
+2 -2
View File
@@ -26,7 +26,7 @@ type LibraryRoot struct {
// Media 是单个可播放项。剧集链接到 SeriesID;电影 SeriesID == ""。 // Media 是单个可播放项。剧集链接到 SeriesID;电影 SeriesID == ""。
type Media struct { type Media struct {
Base 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"` LibraryRootID string `gorm:"index;size:36" json:"library_root_id,omitempty"`
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"` SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
Title string `gorm:"size:255;not null" json:"title"` Title string `gorm:"size:255;not null" json:"title"`
@@ -46,7 +46,7 @@ type Media struct {
Overview string `gorm:"type:text" json:"overview,omitempty"` Overview string `gorm:"type:text" json:"overview,omitempty"`
Rating float32 `json:"rating"` Rating float32 `json:"rating"`
Year int `json:"year"` 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"` SeasonNum int `json:"season_num"`
EpisodeNum int `json:"episode_num"` EpisodeNum int `json:"episode_num"`
ScrapeStatus string `gorm:"size:16;default:pending" json:"scrape_status"` ScrapeStatus string `gorm:"size:16;default:pending" json:"scrape_status"`
-1
View File
@@ -46,7 +46,6 @@ func AllModels() []interface{} {
&APIConfig{}, &APIConfig{},
&UserPermission{}, &UserPermission{},
&RefreshToken{}, &RefreshToken{},
&ApiConfig{},
&PlayProfile{}, &PlayProfile{},
&RegistrationCode{}, &RegistrationCode{},
&SignIn{}, &SignIn{},
+32 -19
View File
@@ -10,17 +10,17 @@ import (
"github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/model"
) )
// ApiConfigRepository persists model.ApiConfig records. // ApiConfigRepository persists model.APIConfig records.
type ApiConfigRepository struct{ db *gorm.DB } type ApiConfigRepository struct{ db *gorm.DB }
// Create inserts a new API config record. // 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 return r.db.WithContext(ctx).Create(c).Error
} }
// FindByProvider returns the API config for a provider, or (nil, nil). // FindByProvider returns the API config for a provider, or (nil, nil).
func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) { func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.APIConfig, error) {
var c model.ApiConfig var c model.APIConfig
err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil return nil, nil
@@ -32,27 +32,40 @@ func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider strin
} }
// List returns all API configs. // List returns all API configs.
func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) { func (r *ApiConfigRepository) List(ctx context.Context) ([]model.APIConfig, error) {
var rows []model.ApiConfig var rows []model.APIConfig
err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error
return rows, err return rows, err
} }
// Upsert creates or updates an API config. // Upsert creates or updates an API config.
func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.ApiConfig) error { // 显式 map 更新:Assign(struct) 会跳过零值字段,导致 Enabled=false、
return r.db.WithContext(ctx).Where("provider = ?", c.Provider). // 清空 BaseURL/Extra 等撤销操作静默失效。
Assign(model.ApiConfig{ func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.APIConfig) error {
Base: model.Base{UpdatedAt: time.Now()}, return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
APIKey: c.APIKey, var existing model.APIConfig
BaseURL: c.BaseURL, err := tx.Where("provider = ?", c.Provider).First(&existing).Error
Extra: c.Extra, if errors.Is(err, gorm.ErrRecordNotFound) {
Enabled: c.Enabled, return tx.Create(c).Error
}).FirstOrCreate(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. // Update updates an API config.
func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) error { func (r *ApiConfigRepository) Update(ctx context.Context, c *model.APIConfig) error {
return r.db.WithContext(ctx).Model(&model.ApiConfig{}). return r.db.WithContext(ctx).Model(&model.APIConfig{}).
Where("provider = ?", c.Provider).Updates(map[string]any{ Where("provider = ?", c.Provider).Updates(map[string]any{
"api_key": c.APIKey, "api_key": c.APIKey,
"base_url": c.BaseURL, "base_url": c.BaseURL,
@@ -64,13 +77,13 @@ func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) er
// Delete 物理删除 API 配置。 // Delete 物理删除 API 配置。
func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error { 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 更新测试结果。 // UpdateTestResult 更新测试结果。
func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error { func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error {
now := time.Now() 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{ Where("provider = ?", provider).Updates(map[string]any{
"test_result": result, "test_result": result,
"last_tested_at": &now, "last_tested_at": &now,
+62 -29
View File
@@ -22,45 +22,93 @@ import (
// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending') // 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
// 永远捞不到数据。 // 永远捞不到数据。
func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
return withSQLiteBusyRetry(ctx, func() error { var indexIDs []string
return r.upsertWithDB(ctx, r.db, m) 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)一次, // UpsertBatch 在单个事务里逐条执行 Upsert:扫描一批只提交(fsync)一次,
// 而不是每条一个隐式事务。任一条目落库失败不影响批内已成功的条目—— // 而不是每条一个隐式事务。任一条目落库失败不影响批内已成功的条目——
// 事务回滚后由调用方退回逐条 Upsert 兜底。 // 事务回滚后由调用方退回逐条 Upsert 兜底。
//
// OpenSearch 索引同步(HTTP,4s 超时)必须在事务提交之后统一执行:放在
// 事务内会把 SQLite 写锁挂起在网络 IO 上,且批内用非事务连接回读只能
// 拿到提交前的旧版本数据,把陈旧内容写进索引。
func (r *MediaRepository) UpsertBatch(ctx context.Context, items []*model.Media) error { func (r *MediaRepository) UpsertBatch(ctx context.Context, items []*model.Media) error {
if len(items) == 0 { if len(items) == 0 {
return nil 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 { return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
for _, m := range items { for _, m := range items {
if m == nil { if m == nil {
continue continue
} }
if err := r.upsertWithDB(ctx, tx, m); err != nil { id, err := r.upsertWithDB(ctx, tx, m)
if err != nil {
return err return err
} }
if id != "" {
indexIDs = append(indexIDs, id)
}
} }
return nil 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 { if err != nil {
return err 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 { if created {
r.indexMediaBestEffort(ctx, *m) return m.ID, nil
return nil
} }
updates := mediaUpsertUpdates(existing, *m) 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) { 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 return *m, true, nil
} else if retryErr := db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil { } else if retryErr := db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil {
return model.Media{}, false, createErr return model.Media{}, false, createErr
} else {
// 并发插入竞态:重查已命中既有行,直接走更新分支。
return existing, false, nil
} }
} }
if err != 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) { func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
if current != next { if current != next {
updates[key] = next updates[key] = next
+11 -3
View File
@@ -105,11 +105,17 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string,
var total int64 var total int64
q := r.db.WithContext(ctx).Model(&model.Media{}) q := r.db.WithContext(ctx).Model(&model.Media{})
q = applyMediaQueryFilter(q, filter) 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) terms := mediaSearchTerms(query)
for _, term := range terms { for _, term := range terms {
like := "%" + escapeLike(term) + "%" like := "%" + escapeLike(term) + "%"
q = q.Where( 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, like, like, like, like,
) )
} }
@@ -120,7 +126,7 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string,
prefix := escapeLike(query) + "%" prefix := escapeLike(query) + "%"
exact := query exact := query
q = q.Order(gorm.Expr( 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, exact, exact, prefix, prefix,
)) ))
} else { } else {
@@ -259,7 +265,9 @@ func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool {
} }
r.searchIndexOnce.Do(func() { r.searchIndexOnce.Do(func() {
var count int64 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'`). Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`).
Scan(&count).Error Scan(&count).Error
r.searchIndexAvailable = err == nil && count > 0 r.searchIndexAvailable = err == nil && count > 0
+36 -2
View File
@@ -3,6 +3,7 @@ package repository
import ( import (
"context" "context"
"errors" "errors"
"time"
"gorm.io/gorm" "gorm.io/gorm"
@@ -44,10 +45,43 @@ func (r *PermissionRepository) Update(ctx context.Context, userID string, update
} }
// Upsert creates or updates a permission record. // Upsert creates or updates a permission record.
// 显式 map 更新:Assign(struct) 会被 GORM 跳过零值字段,导致权限
// "撤销"(false)保存后静默失效且无法重置。
func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error { func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
return withSQLiteBusyRetry(ctx, func() error { return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("user_id = ?", p.UserID). return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
Assign(*p).FirstOrCreate(p).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
})
}) })
} }
@@ -52,6 +52,40 @@ func (r *ScrapeTaskRepository) FindActiveByMediaID(ctx context.Context, mediaID
return &t, err 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) { func (r *ScrapeTaskRepository) List(ctx context.Context, status string, page, pageSize int) ([]model.ScrapeTask, int64, error) {
if page < 1 { if page < 1 {
page = 1 page = 1
+63
View File
@@ -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 { func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) error {
return withSQLiteBusyRetry(ctx, func() error { return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmDownloadTask{}).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 { func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) error {
return withSQLiteBusyRetry(ctx, func() error { return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error
@@ -20,7 +20,7 @@ func TestAdultProviderRouting(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("failed to open sqlite: %v", err) t.Fatalf("failed to open sqlite: %v", err)
} }
_ = db.AutoMigrate(&model.Setting{}, &model.ApiConfig{}) _ = db.AutoMigrate(&model.Setting{}, &model.APIConfig{})
repos := repository.New(db) repos := repository.New(db)
+6 -6
View File
@@ -40,7 +40,7 @@ func (s *ApiConfigService) TestConnection(ctx context.Context, provider string)
} }
// testTMDb 测试 TMDb API 连接。 // testTMDb 测试 TMDb API 连接。
func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) { func (s *ApiConfigService) testTMDb(cfg *model.APIConfig) (string, error) {
if cfg.APIKey == "" { if cfg.APIKey == "" {
return "error", errors.New("API key is required") return "error", errors.New("API key is required")
} }
@@ -74,7 +74,7 @@ func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) {
} }
// testOpenAI 测试 OpenAI API 连接。 // testOpenAI 测试 OpenAI API 连接。
func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) { func (s *ApiConfigService) testOpenAI(cfg *model.APIConfig) (string, error) {
if cfg.APIKey == "" { if cfg.APIKey == "" {
return "error", errors.New("API key is required") return "error", errors.New("API key is required")
} }
@@ -108,7 +108,7 @@ func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) {
} }
// testDeepSeek 测试 DeepSeek API 连接。 // testDeepSeek 测试 DeepSeek API 连接。
func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) { func (s *ApiConfigService) testDeepSeek(cfg *model.APIConfig) (string, error) {
if cfg.APIKey == "" { if cfg.APIKey == "" {
return "error", errors.New("API key is required") return "error", errors.New("API key is required")
} }
@@ -142,7 +142,7 @@ func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) {
} }
// testSiliconFlow 测试 SiliconFlow API 连接。 // testSiliconFlow 测试 SiliconFlow API 连接。
func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error) { func (s *ApiConfigService) testSiliconFlow(cfg *model.APIConfig) (string, error) {
if cfg.APIKey == "" { if cfg.APIKey == "" {
return "error", errors.New("API key is required") 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) 刮削数据源连接与年龄验证。 // 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{} bases := []string{}
if cfg.BaseURL != "" { if cfg.BaseURL != "" {
bases = append(bases, adultConfiguredBases(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。 // 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), "/") serverURL := strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
if serverURL == "" { if serverURL == "" {
serverURL = "http://127.0.0.1:7700" serverURL = "http://127.0.0.1:7700"
+3 -3
View File
@@ -9,7 +9,7 @@ import (
) )
// GetEffectiveConfig 获取生效的 API 配置(数据库配置优先于配置文件)。 // 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) cfg, err := s.GetByProvider(ctx, provider)
if err == nil && cfg != nil { if err == nil && cfg != nil {
@@ -21,7 +21,7 @@ func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider stri
} }
// getConfigFromFile 从配置文件获取 API 配置。 // getConfigFromFile 从配置文件获取 API 配置。
func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig, error) { func (s *ApiConfigService) getConfigFromFile(provider string) (*model.APIConfig, error) {
var apiKey string var apiKey string
var hasKey bool var hasKey bool
@@ -44,7 +44,7 @@ func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig,
return nil, ErrApiConfigNotFound return nil, ErrApiConfigNotFound
} }
return &model.ApiConfig{ return &model.APIConfig{
Provider: provider, Provider: provider,
APIKey: apiKey, APIKey: apiKey,
Enabled: true, Enabled: true,
+5 -5
View File
@@ -33,7 +33,7 @@ var (
) )
// GetByProvider 获取指定提供者的 API 配置。 // 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) cfg, err := s.repo.ApiConfig.FindByProvider(ctx, provider)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -49,7 +49,7 @@ func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (
} }
// List 返回所有 API 配置。 // 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) configs, err := s.repo.ApiConfig.List(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -69,7 +69,7 @@ func (s *ApiConfigService) GetProviders() []model.ApiProvider {
} }
// Upsert 创建或更新 API 配置,自动加密敏感字段。 // 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) { if !s.isValidProvider(provider) {
return nil, ErrInvalidProvider return nil, ErrInvalidProvider
@@ -81,7 +81,7 @@ func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey,
encryptedKey = s.crypto.Encrypt(apiKey) encryptedKey = s.crypto.Encrypt(apiKey)
} }
cfg := &model.ApiConfig{ cfg := &model.APIConfig{
Provider: provider, Provider: provider,
APIKey: encryptedKey, APIKey: encryptedKey,
BaseURL: baseURL, BaseURL: baseURL,
@@ -112,7 +112,7 @@ func (s *ApiConfigService) Update(ctx context.Context, provider string, apiKey,
encryptedKey = s.crypto.Encrypt(apiKey) encryptedKey = s.crypto.Encrypt(apiKey)
} }
cfg := &model.ApiConfig{ cfg := &model.APIConfig{
Provider: provider, Provider: provider,
APIKey: encryptedKey, APIKey: encryptedKey,
BaseURL: baseURL, BaseURL: baseURL,
+10 -1
View File
@@ -8,6 +8,8 @@ import (
"net/url" "net/url"
"path" "path"
"strings" "strings"
"sync"
"time"
) )
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint. // cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
@@ -22,11 +24,18 @@ type cloudDrive2Provider struct {
base *url.URL base *url.URL
username string username string
password string password string
token string token string // 配置的静态令牌(构造后只读)
ua string ua string
apiBase *url.URL apiBase *url.URL
client *http.Client client *http.Client
proxy bool 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 { func newCloudDrive2(cfg map[string]any, client *http.Client) *cloudDrive2Provider {
+3 -2
View File
@@ -36,9 +36,10 @@ func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.decorateDAVStatusError(resp, target) return nil, p.decorateDAVStatusError(resp, target)
} }
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20)) // 流式解码:超大目录(如上万条目的网盘目录)响应可能远超旧 4MB 截断上限,
// 直接 xml.Unmarshal 会截断报错;这里用 LimitReader(64MB) + Decoder 边读边解
var multi cloudDAVMultiStatus 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) return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err)
} }
basePath := strings.TrimRight(p.base.EscapedPath(), "/") basePath := strings.TrimRight(p.base.EscapedPath(), "/")
+46 -18
View File
@@ -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 { func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath string, payload any, action string) error {
token, err := p.openListAPIToken(ctx) _, err := doWithOpenListAPIToken(ctx, p, func(token string) (struct{}, error) {
if err != nil { return struct{}{}, p.openListAPIPostWithToken(ctx, apiPath, payload, action, token)
return err })
} return err
}
func (p *cloudDrive2Provider) openListAPIPostWithToken(ctx context.Context, apiPath string, payload any, action, token string) error {
body, _ := json.Marshal(payload) body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body))
if err != nil { if err != nil {
@@ -151,6 +154,9 @@ func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath strin
return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err) return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err)
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
return errOpenListAPITokenExpired
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("%s: api %s returned http %d", p.name, action, resp.StatusCode) 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 同款契约: // 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 { func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath string, r io.Reader) error {
token, err := p.openListAPIToken(ctx) token, err := p.openListAPIToken(ctx)
if err != nil { if err != nil {
return err return err
} }
encodedPath := openListPathEscape(remotePath) encodedPath := openListPathEscape(remotePath)
body := &bytes.Buffer{}
writer := multipart.NewWriter(body) pr, pw := io.Pipe()
formFile, err := writer.CreateFormFile("file", path.Base(remotePath)) writer := multipart.NewWriter(pw)
if err != nil { go func() {
return err var writeErr error
} defer func() {
if _, err := io.Copy(formFile, r); err != nil { // 读源失败必须传给 pipe 写端,让 HTTP 请求以失败收场而不是静默截断
return err if writeErr != nil {
} _ = pw.CloseWithError(writeErr)
if err := writer.Close(); err != nil { return
return err }
} _ = pw.Close()
req, err := http.NewRequestWithContext(ctx, http.MethodPut, p.openListAPIURL("/api/fs/form"), body) }()
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 { if err != nil {
// 关闭读端以释放仍在等待写入的后台 goroutine(其 Write 会立即失败返回)
_ = pr.Close()
return err return err
} }
req.Header.Set("Authorization", token) req.Header.Set("Authorization", token)
@@ -230,9 +252,15 @@ func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath
req.Header.Set("Overwrite", "true") req.Header.Set("Overwrite", "true")
resp, err := p.client.Do(req) resp, err := p.client.Do(req)
if err != nil { if err != nil {
// 传输层失败(含提前断开)时 net/http 会关闭请求 body,解除后台 goroutine 阻塞
return decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/form"), err) return decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/form"), err)
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
// 流式 body 无法重放,不能自动重试:清除登录 token 缓存让下次上传重新登录,
// 本次返回明确错误交由调用方重试
p.invalidateOpenListAPIToken()
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return p.openListAPIStatusError("upload", remotePath, resp.StatusCode) return p.openListAPIStatusError("upload", remotePath, resp.StatusCode)
} }
+74 -7
View File
@@ -4,18 +4,50 @@ import (
"bytes" "bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"net/url" "net/url"
"strings" "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) token, err := p.openListAPIToken(ctx)
if err != nil { 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 const pageSize = 500
target := normalizeCloudDAVPath(dir) target := normalizeCloudDAVPath(dir)
out := make([]FileEntry, 0, pageSize) out := make([]FileEntry, 0, pageSize)
@@ -45,6 +77,9 @@ func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) (
var decoded openListListResponse var decoded openListListResponse
decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded) decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded)
resp.Body.Close() resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
return nil, errOpenListAPITokenExpired
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.openListAPIStatusError("list", target, resp.StatusCode) 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) { func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
token, err := p.openListAPIToken(ctx) return doWithOpenListAPIToken(ctx, p, func(token string) (*DirectLink, error) {
if err != nil { return p.resolveOpenListAPIDirectWithToken(ctx, fileRef, token)
return nil, err })
} }
func (p *cloudDrive2Provider) resolveOpenListAPIDirectWithToken(ctx context.Context, fileRef, token string) (*DirectLink, error) {
payload, _ := json.Marshal(map[string]string{"path": normalizeCloudDAVPath(fileRef), "password": ""}) 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)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/get"), bytes.NewReader(payload))
if err != nil { 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) return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err)
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
return nil, errOpenListAPITokenExpired
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.openListAPIStatusError("get", fileRef, resp.StatusCode) 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 != "") 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) { func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) {
if token := strings.TrimSpace(p.token); token != "" { if token := strings.TrimSpace(p.token); token != "" {
return token, nil return token, nil
@@ -170,6 +214,30 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err
if strings.TrimSpace(p.username) == "" || p.password == "" { if strings.TrimSpace(p.username) == "" || p.password == "" {
return "", nil 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{ payload, _ := json.Marshal(map[string]string{
"username": p.username, "username": p.username,
"password": p.password, "password": p.password,
@@ -204,7 +272,6 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err
if token == "" { if token == "" {
return "", fmt.Errorf("%s: api login returned empty token", p.name) return "", fmt.Errorf("%s: api login returned empty token", p.name)
} }
p.token = token
return token, nil return token, nil
} }
+8 -5
View File
@@ -48,7 +48,7 @@ func (p *openAPI115Provider) Ping(ctx context.Context) error {
if strings.TrimSpace(p.c.AppID) == "" { if strings.TrimSpace(p.c.AppID) == "" {
return fmt.Errorf("115: 缺少开放平台应用 ID,请重新授权") return fmt.Errorf("115: 缺少开放平台应用 ID,请重新授权")
} }
if strings.TrimSpace(p.c.AccessToken) == "" { if strings.TrimSpace(p.c.CurrentAccessToken()) == "" {
return fmt.Errorf("115: 缺少访问令牌,请重新授权") return fmt.Errorf("115: 缺少访问令牌,请重新授权")
} }
_, _, err := p.c.GetFsList(ctx, "0", 0, 1) _, _, 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, Name: f.FileName,
IsDir: f.Category == cloud115.TypeDir, IsDir: f.Category == cloud115.TypeDir,
Size: f.FileSize, Size: f.FileSize,
MTime: f.Utime, MTime: f.ModifiedAt(),
PickCode: f.PickCode, PickCode: f.PickCode,
}) })
} }
@@ -127,12 +127,15 @@ func (p *openAPI115Provider) PutFileNamed(ctx context.Context, parentCID, fileNa
if err := tmp.Close(); err != nil { if err := tmp.Close(); err != nil {
return fmt.Errorf("115: 关闭临时文件失败:%w", err) return fmt.Errorf("115: 关闭临时文件失败:%w", err)
} }
// 重命名为目标文件名,保证上传到 115 后保留原始文件名 // 重命名为目标文件名,保证上传到 115 后保留原始文件名。
// 重命名失败必须 fail fast:静默用随机临时名上传会导致 115 上的文件名
// 变成 mebox-upload-xxx,破坏元数据文件名契约。
if fileName != "" && fileName != filepath.Base(tmpPath) { if fileName != "" && fileName != filepath.Base(tmpPath) {
namedPath := filepath.Join(filepath.Dir(tmpPath), fileName) namedPath := filepath.Join(filepath.Dir(tmpPath), fileName)
if err := os.Rename(tmpPath, namedPath); err == nil { if err := os.Rename(tmpPath, namedPath); err != nil {
tmpPath = namedPath return fmt.Errorf("115: 重命名临时文件为 %s 失败:%w", fileName, err)
} }
tmpPath = namedPath
} }
_, err = p.c.Upload(ctx, tmpPath, parentCID, "", "") _, err = p.c.Upload(ctx, tmpPath, parentCID, "", "")
if err != nil { if err != nil {
+95 -16
View File
@@ -23,9 +23,14 @@ type OpenClient struct {
RefreshTokenStr string RefreshTokenStr string
executor *QueueExecutor executor *QueueExecutor
// tokenMu 保护令牌刷新:业务请求中途 access_token 失效时自动刷新重试, // OnTokenRefreshed 在 access_token 刷新成功后回调(参数为新令牌对),
// 多 goroutine(同步列表 + 下载队列)并发下只允许一次刷新进行。 // 供上层持久化新令牌使用;nil 安全,且在 tokenMu 释放后调用以避免死锁。
tokenMu sync.Mutex OnTokenRefreshed func(accessToken, refreshToken string)
// tokenMu 保护 AccessToken / RefreshTokenStr 的并发读写:业务请求中途
// access_token 失效时自动刷新重试,多 goroutine(同步列表 + 下载队列)
// 并发下只允许一次刷新进行。
tokenMu sync.RWMutex
} }
// default115HTTPClient 创建带有防 405 重定向保护的 http.Client。 // default115HTTPClient 创建带有防 405 重定向保护的 http.Client。
@@ -57,12 +62,40 @@ func NewOpenClient(appID, accessToken, refreshToken string) *OpenClient {
} }
} }
// SetAuthToken 更新认证令牌。 // SetAuthToken 更新认证令牌(并发安全)。
func (c *OpenClient) SetAuthToken(accessToken, refreshToken string) { 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.AccessToken = accessToken
c.RefreshTokenStr = refreshToken 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 返回数字)。 // RespState 兼容 115 不同端点返回的 state 类型(proapi 返回布尔、passport 返回数字)。
type RespState bool type RespState bool
@@ -191,6 +224,14 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map
if access { if access {
// 刷新失败(或已刷新仍失败)时返回明确错误 // 刷新失败(或已刷新仍失败)时返回明确错误
lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115: access_token 校验失败且刷新未成功") 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 return &base, lastErr
} }
@@ -242,8 +283,11 @@ func (c *OpenClient) buildRequestWithUA(ctx context.Context, method, rawURL stri
if method == http.MethodPost && len(form) > 0 { if method == http.MethodPost && len(form) > 0 {
req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
} }
if access && c.AccessToken != "" { if access {
req.Header.Set("Authorization", "Bearer "+c.AccessToken) // RLock 读取令牌,避免与刷新流程的写入产生数据竞争
if accessToken := c.currentAccessToken(); accessToken != "" {
req.Header.Set("Authorization", "Bearer "+accessToken)
}
} }
return req, nil return req, nil
} }
@@ -261,33 +305,57 @@ func (c *OpenClient) doAuthJSONWithUA(ctx context.Context, method, rawURL string
// tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方 // tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方
// 应使用内存中的新 token 重试原请求)。 // 应使用内存中的新 token 重试原请求)。
// //
// 拿到写锁后在锁内读取 oldAccess,与持锁期间的当前值对比:若已被其他
// goroutine 刷新过则直接复用新 token,避免并发请求连环轮转消耗 115 的
// 一次性 refresh_token。全程持写锁读写 token 字段,无 TOCTOU 窗口。
//
// 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120) // 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120)
// 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。 // 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。
// 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃, // 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃,
// 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。 // 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。
func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool { func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool {
c.tokenMu.Lock() 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++ { for attempt := 0; attempt < refreshAttempts; attempt++ {
token, err := c.RefreshToken(c.RefreshTokenStr) token, err := c.doRefreshToken(refreshToken)
if err == nil { if err == nil {
c.SetAuthToken(token.AccessToken, token.RefreshToken) c.setAuthTokenLocked(token.AccessToken, token.RefreshToken)
return true return token, true
} }
if IsRefreshTokenDead(err) { if IsRefreshTokenDead(err) {
c.SetAuthToken("", "") c.setAuthTokenLocked("", "")
return false return nil, false
} }
// 可恢复失败:退避后重试。ctx 取消时立即放弃。 // 可恢复失败:退避后重试。ctx 取消时立即放弃。
if attempt < refreshAttempts-1 { if attempt < refreshAttempts-1 {
select { select {
case <-ctx.Done(): case <-ctx.Done():
return false return nil, false
case <-time.After(refreshBackoff(attempt)): case <-time.After(refreshBackoff(attempt)):
} }
} }
} }
return false return nil, false
} }
// refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。 // refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。
@@ -312,7 +380,14 @@ func isTokenCode(code int) bool {
} }
// openList 解析 data 为对象或数组(StructOrArray 语义)。 // openList 解析 data 为对象或数组(StructOrArray 语义)。
// 115 部分接口在鉴权/业务异常时会返回 data:null 或 data:{},此时若直接
// 反序列化会得到零值元素 + nil error,调用方会把空数据当成功处理;
// 这里对 null/空对象显式报错。
func openList[T any](raw json.RawMessage) ([]T, error) { 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 var single T
if err := json.Unmarshal(raw, &single); err == nil { if err := json.Unmarshal(raw, &single); err == nil {
return []T{single}, nil return []T{single}, nil
@@ -324,12 +399,16 @@ func openList[T any](raw json.RawMessage) ([]T, error) {
return nil, fmt.Errorf("115: data 既不是对象也不是数组") return nil, fmt.Errorf("115: data 既不是对象也不是数组")
} }
// openFirstList 取 data 的第一个元素。 // openFirstList 取 data 的第一个元素;data 为空(null/空数组)时返回显式错误,
// 避免调用方拿到 (nil, nil) 后解引用空指针。
func openFirstList[T any](raw json.RawMessage) (*T, error) { func openFirstList[T any](raw json.RawMessage) (*T, error) {
items, err := openList[T](raw) items, err := openList[T](raw)
if err != nil || len(items) == 0 { if err != nil {
return nil, err return nil, err
} }
if len(items) == 0 {
return nil, fmt.Errorf("115: data 为空数组")
}
return &items[0], nil return &items[0], nil
} }
+6 -1
View File
@@ -317,13 +317,18 @@ func appendCallbackParams(rawURL string, params url.Values) (string, error) {
return callbackURL.String(), nil 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) { func httpGetJSON(ctx context.Context, endpoint string) (map[string]any, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil { if err != nil {
return nil, err return nil, err
} }
req.Header.Set("User-Agent", DefaultUA) req.Header.Set("User-Agent", DefaultUA)
resp, err := http.DefaultClient.Do(req) resp, err := oauthHTTPClient.Do(req)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+31 -12
View File
@@ -300,6 +300,11 @@ func (c *OpenClient) GetQrCode() (*QrCodeDataReturn, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
// 关键字段缺失时显式报错:空 uid/sign 会导致后续扫码轮询必然失败,
// 不能把残缺响应当成功返回给界面。
if code.Uid == "" || code.Sign == "" {
return nil, fmt.Errorf("115: 设备码响应缺少 uid/sign,无法发起扫码授权")
}
return &QrCodeDataReturn{QrCodeData: *code, CodeVerifier: codeVerifier}, nil return &QrCodeDataReturn{QrCodeData: *code, CodeVerifier: codeVerifier}, nil
} }
@@ -352,6 +357,10 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) {
if err != nil { if err != nil {
return nil, err 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) c.SetAuthToken(token.AccessToken, token.RefreshToken)
return token, nil return token, nil
} }
@@ -359,11 +368,30 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) {
// RefreshToken 刷新访问令牌。 // RefreshToken 刷新访问令牌。
func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) { func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) {
if refreshToken == "" { if refreshToken == "" {
refreshToken = c.RefreshTokenStr refreshToken = c.currentRefreshToken()
} }
if refreshToken == "" { if refreshToken == "" {
return nil, fmt.Errorf("没有可用的 refresh_token") 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} params := map[string]string{"refresh_token": refreshToken}
resp, err := c.doJSON(context.Background(), "POST", PassportAPIBase+"/open/refreshToken", params, false, 0) resp, err := c.doJSON(context.Background(), "POST", PassportAPIBase+"/open/refreshToken", params, false, 0)
if err != nil && resp == nil { if err != nil && resp == nil {
@@ -373,18 +401,9 @@ func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) {
return nil, err return nil, err
} }
if !resp.State { if !resp.State {
apiErr := NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败") return nil, NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败")
if IsRefreshTokenDead(apiErr) {
c.SetAuthToken("", "")
}
return nil, apiErr
} }
token, err := openFirstList[TokenData](resp.Data) return openFirstList[TokenData](resp.Data)
if err != nil {
return nil, err
}
c.SetAuthToken(token.AccessToken, token.RefreshToken)
return token, nil
} }
// ─── 用户信息 ────────────────────────────────────────────────────────────────── // ─── 用户信息 ──────────────────────────────────────────────────────────────────
+76 -37
View File
@@ -10,6 +10,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"io" "io"
"log"
"os" "os"
"sort" "sort"
@@ -106,14 +107,23 @@ func (u *OSSMultipartUploader) UploadFile(ctx context.Context, input OSSMultipar
return result.CallbackResult, nil return result.CallbackResult, nil
} }
// UploadedPart 是 OSS 已上传分片的定位信息(断点续传时复用 ETag 用)。
type UploadedPart struct {
PartNumber int32
Size int64
ETag string
}
// UploadFileWithResult 上传文件并返回 multipart 结果。 // 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 { if input.PartRetryMax <= 0 {
input.PartRetryMax = 3 input.PartRetryMax = 3
} }
partSize := input.PartSize partSize := input.PartSize
totalParts := 0 totalParts := 0
var err error
if partSize <= 0 { if partSize <= 0 {
partSize, totalParts, err = CalculateMultipartPartSize(input.FileSize) partSize, totalParts, err = CalculateMultipartPartSize(input.FileSize)
if err != nil { if err != nil {
@@ -124,28 +134,45 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
} }
uploadId := input.UploadId uploadId := input.UploadId
if uploadId == "" { // ownUploadId 标记 uploadId 是否为本调用 Initiate 出来的:仅自建的
initResult, err := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{ // multipart 在失败时由本函数 Abort;调用方显式传入的 uploadId(断点续传)
// 失败后保留现场,由调用方决定重试或清理。
ownUploadId := uploadId == ""
if ownUploadId {
initResult, initErr := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{
Bucket: oss.Ptr(input.Bucket), Bucket: oss.Ptr(input.Bucket),
Key: oss.Ptr(input.Object), Key: oss.Ptr(input.Object),
RequestCommon: oss.RequestCommon{ RequestCommon: oss.RequestCommon{
Parameters: map[string]string{"sequential": "1"}, Parameters: map[string]string{"sequential": "1"},
}, },
}) })
if err != nil { if initErr != nil {
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", err) return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", initErr)
} }
if initResult.UploadId == nil || *initResult.UploadId == "" { if initResult.UploadId == nil || *initResult.UploadId == "" {
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 返回空 upload_id") return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 返回空 upload_id")
} }
uploadId = *initResult.UploadId 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) existingPartMap := make(map[int32]UploadedPart)
existingParts, err := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId) if existingParts, listErr := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId); listErr == nil {
if err == nil {
for _, part := range existingParts { 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 { if length < 0 {
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 uploadedBytes += length
uploadedParts++ 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) etag, uploadErr := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length)
if err != nil { if uploadErr != nil {
return OSSMultipartUploadResult{}, err return OSSMultipartUploadResult{}, uploadErr
} }
uploadedBytes += length uploadedBytes += length
uploadedParts++ uploadedParts++
@@ -224,29 +258,34 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
}, nil }, nil
} }
// ListUploadedParts 查询 OSS 已上传分片。 // ListUploadedParts 查询 OSS 已上传分片(MaxParts 上限 1000,超过时按
func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]struct { // NextPartNumberMarker 自动翻页取全量,否则断点续传只能看到前 1000 片)。
PartNumber int32 func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]UploadedPart, error) {
Size int64 parts := []UploadedPart{}
}, error) { var marker int32
parts := []struct { for {
PartNumber int32 result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{
Size int64 Bucket: oss.Ptr(bucket),
}{} Key: oss.Ptr(object),
result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{ UploadId: oss.Ptr(uploadId),
Bucket: oss.Ptr(bucket), MaxParts: 1000,
Key: oss.Ptr(object), PartNumberMarker: marker,
UploadId: oss.Ptr(uploadId), })
MaxParts: 1000, if err != nil {
}) return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err)
if err != nil { }
return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err) for _, part := range result.Parts {
} etag := ""
for _, part := range result.Parts { if part.ETag != nil {
parts = append(parts, struct { etag = *part.ETag
PartNumber int32 }
Size int64 parts = append(parts, UploadedPart{PartNumber: part.PartNumber, Size: part.Size, ETag: etag})
}{PartNumber: part.PartNumber, Size: part.Size}) }
if !result.IsTruncated || result.NextPartNumberMarker <= marker {
// 防御:marker 不前进时终止循环,避免异常响应导致死循环
break
}
marker = result.NextPartNumberMarker
} }
return parts, nil return parts, nil
} }
+6 -2
View File
@@ -262,7 +262,10 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s
} }
switch status { switch status {
case UploadInitStatusRapidUploaded: 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 return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil
case UploadInitStatusSignFailed: case UploadInitStatusSignFailed:
return nil, fmt.Errorf("115: 签名验证后失败") return nil, fmt.Errorf("115: 签名验证后失败")
@@ -271,7 +274,8 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s
case UploadInitStatusNeedUpload: case UploadInitStatusNeedUpload:
// 真实上传:OSS multipart // 真实上传:OSS multipart
default: default:
return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil // 未知状态不能当成功返回(会静默丢文件),显式报错便于排查
return nil, fmt.Errorf("115: 未知的 upload/init 状态 %d", status)
} }
if initResult.Bucket == "" || initResult.Object == "" { if initResult.Bucket == "" || initResult.Object == "" {
+15 -3
View File
@@ -1,14 +1,26 @@
package cloud115 package cloud115
import "math/rand" import (
"crypto/rand"
"fmt"
"math/big"
)
const randCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" const randCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
// RandomString 生成指定长度的随机字符串(PKCE code_verifier 等)。 // RandomString 生成指定长度的密码学安全随机字符串(PKCE code_verifier、
// OAuth state 等安全敏感场景)。必须使用 crypto/rand:math/rand 未播种时
// 序列可预测,会造成 PKCE 防御失效。
func RandomString(length int) string { func RandomString(length int) string {
b := make([]byte, length) b := make([]byte, length)
max := big.NewInt(int64(len(randCharset)))
for i := range b { 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) return string(b)
} }
+6 -2
View File
@@ -36,6 +36,10 @@ type DLNAService struct {
cachedAt time.Time cachedAt time.Time
} }
// dlnaHTTPClient 是 DLNA 专用 HTTP 客户端:SSDP 描述拉取与 SOAP 投递
// 都应快速失败,不占用全局 DefaultClient,也不无限悬挂。
var dlnaHTTPClient = &http.Client{Timeout: 15 * time.Second}
// NewDLNAService is the constructor. // NewDLNAService is the constructor.
func NewDLNAService(log *zap.Logger) *DLNAService { func NewDLNAService(log *zap.Logger) *DLNAService {
return &DLNAService{log: log} return &DLNAService{log: log}
@@ -153,7 +157,7 @@ func (d *DLNAService) fetchDescription(ctx context.Context, location string) (*D
if err != nil { if err != nil {
return nil, err return nil, err
} }
resp, err := http.DefaultClient.Do(req) resp, err := dlnaHTTPClient.Do(req)
if err != nil { if err != nil {
return nil, err 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("Content-Type", `text/xml; charset="utf-8"`)
req.Header.Set("SOAPAction", req.Header.Set("SOAPAction",
fmt.Sprintf(`"urn:schemas-upnp-org:service:AVTransport:1#%s"`, action)) 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 { if err != nil {
return err return err
} }
+72 -28
View File
@@ -18,6 +18,7 @@ import (
"time" "time"
"github.com/truewhile/MeBox/internal/config" "github.com/truewhile/MeBox/internal/config"
"github.com/truewhile/MeBox/internal/model"
"github.com/truewhile/MeBox/internal/repository" "github.com/truewhile/MeBox/internal/repository"
"go.uber.org/zap" "go.uber.org/zap"
) )
@@ -267,12 +268,19 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[
if err != nil { if err != nil {
return nil, err return nil, err
} }
type remoteResult struct { type remoteReply struct {
items []any acct *model.StrmAccount
envelope map[string]any
} }
mounts, aerr := e.remote.ListMounts(ctx) mounts, aerr := e.remote.ListMounts(ctx)
results := make([]remoteResult, 0, len(mounts)) replies := make([]*remoteReply, 0, len(mounts))
if aerr == nil { if aerr == nil {
type mountSearchJob struct {
idx int
mount *model.EmbyMount
acct *model.StrmAccount
}
jobs := make([]*mountSearchJob, 0, len(mounts))
for i := range mounts { for i := range mounts {
m := mounts[i] m := mounts[i]
if !m.Enabled { if !m.Enabled {
@@ -285,40 +293,76 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[
if acct == nil { if acct == nil {
continue continue
} }
// 按挂载逐个搜索:搜索结果归属明确(伪装 ID 正确),也天然只搜已 // idx 使用 jobs 内的序号(而非 mounts 下标):fetched 按
// 挂载的媒体库。 // len(jobs) 分配,必须与 jobs 下标对齐,否则越界 panic。
searchParams := p jobs = append(jobs, &mountSearchJob{idx: len(jobs), mount: &mounts[i], acct: acct})
searchParams.ParentID = "" // RemoteSearchMount 内部设 ParentId }
remote, rerr := e.remote.RemoteSearchMount(ctx, &m, acct, p) // 并发搜索各挂载(限并发 + 单挂载超时):串行时每挂载最多
if rerr != nil { // 15s×线路数,多挂载下首屏延迟被成倍放大。结果按挂载顺序合并。
if e.log != nil { sem := make(chan struct{}, 4)
e.log.Warn("remote emby search failed", var wg sync.WaitGroup
zap.String("account", acct.Name), zap.Error(rerr)) 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 fetched[job.idx] = &remoteReply{acct: job.acct, envelope: remote}
} }(job)
if err := e.mergeRemoteUserData(ctx, p.UserID, remote); err != nil { }
return nil, err wg.Wait()
} for _, r := range fetched {
if raw, ok := remote["Items"].([]any); ok { if r != nil {
results = append(results, remoteResult{items: raw}) replies = append(replies, r)
} 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})
} }
} }
} }
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)...) items = append(items, localItemsAsAny(local)...)
for _, res := range results { for _, reply := range replies {
items = append(items, res.items...) 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 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 { func localItemsAsAny(envelope map[string]any) []any {
if envelope == nil { if envelope == nil {
return nil return nil
+19 -4
View File
@@ -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 return map[string]any{"Items": []any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
} }
// 历史记录限行:此前无上限全量加载,远程条目多时既拖慢 SQL 也放大
// 下面的远程详情请求量。
var hist []model.PlaybackHistory var hist []model.PlaybackHistory
if err := e.repo.DB.WithContext(ctx). if err := e.repo.DB.WithContext(ctx).
Where("user_id = ? AND completed = ? AND position_ms > 0", p.UserID, false). 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 return nil, err
} }
if len(hist) == 0 { 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 { for _, h := range hist {
if m, ok := byID[h.MediaID]; ok { if m, ok := byID[h.MediaID]; ok {
if p.ParentID != "" && m.LibraryID != p.ParentID && m.SeriesID != p.ParentID { if p.ParentID != "" && m.LibraryID != p.ParentID && m.SeriesID != p.ParentID {
continue 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 continue
} }
if e.remote == nil || !IsEmbyRemoteID(h.MediaID) { if e.remote == nil || !IsEmbyRemoteID(h.MediaID) {
continue continue
} }
remoteTotal++
if len(items) >= needed {
continue
}
mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID) mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID)
mount, acct, err := e.remote.ResolveMount(ctx, mountID) mount, acct, err := e.remote.ResolveMount(ctx, mountID)
if err != nil || mount == nil || acct == nil { 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) items = append(items, item)
} }
total := int64(len(items)) total := int64(localTotal + remoteTotal)
if p.StartIndex >= len(items) { if p.StartIndex >= len(items) {
return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
} }
+58 -27
View File
@@ -26,6 +26,7 @@ import (
"regexp" "regexp"
"strconv" "strconv"
"strings" "strings"
"sync"
"time" "time"
"go.uber.org/zap" "go.uber.org/zap"
@@ -73,6 +74,7 @@ type EmbyRemoteService struct {
repo *repository.Container repo *repository.Container
crypto *CryptoService crypto *CryptoService
http *http.Client http *http.Client
stream *http.Client // 流式代理专用(视频/字幕),无整体 Timeout
cache *RuntimeCacheService cache *RuntimeCacheService
} }
@@ -87,6 +89,12 @@ func NewEmbyRemoteService(cfg *config.Config, log *zap.Logger, repo *repository.
Timeout: embyRemoteHTTPTimeout, Timeout: embyRemoteHTTPTimeout,
Transport: &embyRemoteTransport{base: http.DefaultTransport}, 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 return
} }
cfg.RemoteUserID = uid cfg.RemoteUserID = uid
raw := map[string]string{} _ = r.updateAccountConfig(ctx, acct, func(raw map[string]string) {
_ = json.Unmarshal([]byte(acct.Config), &raw) raw["remote_user_id"] = uid
raw["remote_user_id"] = uid })
data, err := json.Marshal(raw)
if err != nil {
return
}
acct.Config = string(data)
_ = r.repo.StrmAccount.Update(ctx, acct)
} }
// CleanupOrphanMounts 清理账号已删除的残留挂载(老版本删除账号未级联), // CleanupOrphanMounts 清理账号已删除的残留挂载(老版本删除账号未级联),
@@ -495,19 +497,28 @@ func (r *EmbyRemoteService) ensureTokenOnLine(ctx context.Context, acct *model.S
return nil return nil
} }
// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。 // acctCfgMu 序列化对账号 Config 的读-改-写。并发请求若各自基于请求开始
func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error { // 时的快照做整包覆盖,会互相丢失更新(刚持久化的 token / active_line 被
if acct == nil { // 旧快照覆盖回去)。
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 return nil
} }
acctCfgMu.Lock()
defer acctCfgMu.Unlock()
raw := map[string]string{} 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) != "" { if strings.TrimSpace(acct.Config) != "" {
_ = json.Unmarshal([]byte(acct.Config), &raw) _ = json.Unmarshal([]byte(acct.Config), &raw)
} }
raw["api_key"] = r.crypto.Encrypt(cfg.Token) if mutate != nil {
raw["remote_user_id"] = cfg.RemoteUserID mutate(raw)
if strings.TrimSpace(raw["username"]) == "" {
raw["username"] = cfg.Username
} }
data, err := json.Marshal(raw) data, err := json.Marshal(raw)
if err != nil { if err != nil {
@@ -517,6 +528,20 @@ func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAc
return r.repo.StrmAccount.Update(ctx, acct) 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。 // doGet 向远程 Emby 发起带 api_key 的 GET,把响应 JSON 解码到 out。
// 401 时自动重认证一次再重试(凭据过期场景)。连接失败时按线路优先级自动切换。 // 401 时自动重认证一次再重试(凭据过期场景)。连接失败时按线路优先级自动切换。
func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, path string, q url.Values, out any) error { 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 { if err != nil {
return fmt.Errorf("请求远程 Emby 失败: %w", err) 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() resp.Body.Close()
if readErr != nil { if readErr != nil {
return readErr return readErr
} }
if len(data) > 8<<20 {
return fmt.Errorf("远程 Emby 响应超过 8MB 上限(路径 %s):请减小分页或 Fields 字段", path)
}
if resp.StatusCode == http.StatusUnauthorized && attempt == 0 { if resp.StatusCode == http.StatusUnauthorized && attempt == 0 {
// 401:只清当前线路的内存 token 并立即重认证;不在此时删除
// DB 里的 api_key——①外层还会按线路故障转移(其他线路可能
// 存有自己的 token);②纯 api_key 账号删除后无法再认证,一次
// 线路误报就会把账号“砖化”。重认证成功后 persistToken 会用
// 新 token 覆盖 api_key。
cfg.Token = "" cfg.Token = ""
master.Token = "" if err := r.ensureTokenOnLine(ctx, acct, cfg); err != nil {
if acct != nil { return fmt.Errorf("认证重试失败: %w", err)
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)
} }
master.Token = cfg.Token
master.RemoteUserID = cfg.RemoteUserID
continue continue
} }
if resp.StatusCode >= 300 { 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 != "" { if rangeHeader := req.Header.Get("Range"); rangeHeader != "" {
upstream.Header.Set("Range", rangeHeader) upstream.Header.Set("Range", rangeHeader)
} }
resp, err := r.http.Do(upstream) resp, err := r.stream.Do(upstream)
if err != nil { if err != nil {
return fmt.Errorf("连接远程 Emby 视频流失败: %w", err) return fmt.Errorf("连接远程 Emby 视频流失败: %w", err)
} }
@@ -1028,7 +1059,7 @@ func (r *EmbyRemoteService) proxySubtitleOnLine(ctx context.Context, w http.Resp
return err return err
} }
upstream.Header.Set("X-Emby-Token", cfg.Token) upstream.Header.Set("X-Emby-Token", cfg.Token)
resp, err := r.http.Do(upstream) resp, err := r.stream.Do(upstream)
if err != nil { if err != nil {
return fmt.Errorf("连接远程 Emby 字幕流失败: %w", err) return fmt.Errorf("连接远程 Emby 字幕流失败: %w", err)
} }
+9 -15
View File
@@ -124,10 +124,11 @@ func isEmbyLineFailoverError(err error) bool {
return false return false
} }
msg := strings.ToLower(err.Error()) msg := strings.ToLower(err.Error())
// 注意:认证类错误(重认证失败 / 缺少凭据)不在此排除——401 后清空
// 内存 token 重认证失败时应继续按线路故障转移,其他线路可能存有
// 自己的 token。仅“登录失败”(密码错误)是账号级问题,无需换线。
if strings.Contains(msg, "登录失败") || if strings.Contains(msg, "登录失败") ||
strings.Contains(msg, "未返回 accesstoken") || strings.Contains(msg, "未返回 accesstoken") {
strings.Contains(msg, "缺少 emby 凭据") ||
strings.Contains(msg, "认证重试失败") {
return false return false
} }
var urlErr *url.Error 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) { if acct == nil || cfg == nil || lineIndex < 0 || lineIndex >= len(cfg.Lines) {
return nil return nil
} }
raw := map[string]string{} err := r.updateAccountConfig(ctx, acct, func(raw map[string]string) {
if strings.TrimSpace(acct.Config) != "" { raw["active_line"] = strconv.Itoa(lineIndex)
_ = json.Unmarshal([]byte(acct.Config), &raw) raw["url"] = cfg.Lines[lineIndex].URL
} })
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)
cfg.ActiveLine = lineIndex cfg.ActiveLine = lineIndex
cfg.BaseURL = normalizeEmbyRemoteURL(cfg.Lines[lineIndex].URL) 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) { func (r *EmbyRemoteService) adoptWorkingLine(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, lineIndex int) {
+53 -23
View File
@@ -371,37 +371,52 @@ func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model
// RemoteMediaDetail 拉远程单条目映射为 Media(网页详情页)。 // RemoteMediaDetail 拉远程单条目映射为 Media(网页详情页)。
func (r *EmbyRemoteService) RemoteMediaDetail(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, error) { 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) cfg, err := r.remoteConfigWithToken(ctx, acct)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID) 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" path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating"
var out map[string]any var out map[string]any
if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil { if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil {
return nil, err return nil, nil, err
} }
RewriteEmbyRemoteIDs(out, mount.ID) RewriteEmbyRemoteIDs(out, mount.ID)
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, out) m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, out)
return &m, nil return &m, out, nil
} }
// RemoteEpisodes 拉远程条目下的集列表(Series/Season/Folder→子集;Episode→同系列; // RemoteEpisodes 拉远程条目下的集列表(Series/Season/Folder→子集;Episode→同系列;
// Movie→自身单条),按季/集排序,与本地 ListMediaEpisodes 行为一致。 // Movie→自身单条),按季/集排序,与本地 ListMediaEpisodes 行为一致。
func (r *EmbyRemoteService) RemoteEpisodes(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) ([]model.Media, error) { 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 { if err != nil {
return nil, err return nil, err
} }
// 用远程详情载荷精判类型(Episode→同系列;Series/Season/Folder→子集;Movie→单条)。 // 用远程详情载荷精判类型(Episode→同系列;Series/Season/Folder→子集;Movie→单条)。
itemType := r.remoteItemType(ctx, acct, remoteID) // Type/SeriesId 都在详情载荷里现成可用,不再为判定类型/系列额外发起
// 两次重复的远程全量 GET(远程慢时页面延迟直接×3)。
itemType := remoteItemString(rawDetail, "Type")
if itemType == "" { if itemType == "" {
itemType = remoteItemTypeOf(detail) itemType = remoteItemTypeOf(detail)
} }
var parentID string var parentID string
switch itemType { switch itemType {
case "Episode": 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 == "" { if parentID == "" {
parentID = remoteID parentID = remoteID
} }
@@ -435,23 +450,36 @@ func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.E
q.Set("ParentId", parentID) q.Set("ParentId", parentID)
q.Set("IncludeItemTypes", "Episode") q.Set("IncludeItemTypes", "Episode")
q.Set("Recursive", "true") 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") q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
var body struct { items := make([]model.Media, 0, 64)
Items []map[string]any `json:"Items"` total := int64(0)
TotalRecordCount int64 `json:"TotalRecordCount"` // 每页 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 items, total, 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
} }
// RemoteSeriesCards 远程剧集库的系列卡片(ChildCount 作为集数)。 // RemoteSeriesCards 远程剧集库的系列卡片(ChildCount 作为集数)。
@@ -477,14 +505,16 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.
q.Set("Recursive", "false") q.Set("Recursive", "false")
q.Set("SortBy", "DateLastContentAdded") q.Set("SortBy", "DateLastContentAdded")
q.Set("SortOrder", "Descending") 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") q.Set("Fields", "Overview,Genres,ProviderIds,Path,RecursiveItemCount,SeriesPrimaryImage,DateCreated,DateLastMediaAdded,PremiereDate,ProductionYear,CommunityRating,CriticRating")
var body struct { var body struct {
Items []map[string]any `json:"Items"` Items []map[string]any `json:"Items"`
TotalRecordCount int64 `json:"TotalRecordCount"` TotalRecordCount int64 `json:"TotalRecordCount"`
} }
cards := make([]SeriesCard, 0) cards := make([]SeriesCard, 0)
for startIndex := 0; ; startIndex += 1000 { for startIndex := 0; ; startIndex += 200 {
q.Set("StartIndex", strconv.Itoa(startIndex)) q.Set("StartIndex", strconv.Itoa(startIndex))
body.Items = nil body.Items = nil
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil { if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
+16
View File
@@ -5,6 +5,7 @@ import (
"errors" "errors"
"strconv" "strconv"
"strings" "strings"
"sync"
"time" "time"
"github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/model"
@@ -218,7 +219,22 @@ func mergedRemoteUserData(raw any, history *model.PlaybackHistory) map[string]an
return userData 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) { 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 { if e.cache != nil {
e.cache.DeletePrefix(ctx, "media:emby:") e.cache.DeletePrefix(ctx, "media:emby:")
} }
+40
View File
@@ -13,9 +13,12 @@
package service package service
import ( import (
"errors"
"net"
"net/http" "net/http"
"path/filepath" "path/filepath"
"sync" "sync"
"syscall"
"time" "time"
"go.uber.org/zap" "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 // from image.tmdb.org via their HTTP proxy without extra config. On
// Windows we also honor the current user's system proxy settings. // Windows we also honor the current user's system proxy settings.
transport := NewExternalTransport() 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{ return &ImageProxy{
cfg: cfg, cfg: cfg,
log: log, 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 // SetLibraryRootsProvider injects a callback that returns the current set of
// media library root directories. Sidecar posters live under these roots // media library root directories. Sidecar posters live under these roots
// (which are arbitrary, user-defined, and not necessarily under the // (which are arbitrary, user-defined, and not necessarily under the
+8 -3
View File
@@ -35,13 +35,18 @@ func isPrivateHost(host string) bool {
if host == "" { if host == "" {
return true return true
} }
ip := net.ParseIP(host) if ip := net.ParseIP(host); ip != nil {
if ip != nil { return isPrivateIP(ip)
return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified()
} }
return false 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. // isAllowedLocalPath restricts local file reads to known-safe roots.
func (p *ImageProxy) isAllowedLocalPath(abs string) bool { func (p *ImageProxy) isAllowedLocalPath(abs string) bool {
roots := []string{p.cfg.App.DataDir, p.cfg.Cache.CacheDir, p.cfg.Media.MoviesDir, p.cfg.Media.TVDir, p.cfg.Media.AnimeDir} roots := []string{p.cfg.App.DataDir, p.cfg.Cache.CacheDir, p.cfg.Media.MoviesDir, p.cfg.Media.TVDir, p.cfg.Media.AnimeDir}
@@ -2,6 +2,7 @@ package service
import ( import (
"context" "context"
"fmt"
"os" "os"
"path/filepath" "path/filepath"
"strconv" "strconv"
@@ -219,24 +220,50 @@ func (o *OrganizerService) replaceVersions(ctx context.Context, src string, exis
o.log.Warn("organize replace sidecar artwork failed", o.log.Warn("organize replace sidecar artwork failed",
zap.String("from", src), zap.String("to", dst), zap.Error(err)) 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 // New file is safely staged. 先把现有 dst 改名为备份、rename stage→dst
// supersede the existing lower-res versions. // 成功后,才删除旧版本:此前顺序是先删旧版本再 rename,一旦 rename
for _, e := range existing { // 失败(Windows 下 dst 被播放器/杀软占用很常见),cleanup 会删掉
if nfo := nfoPath(e); nfo != "" { // stage——旧版本已删、move 模式下源已不在、新文件也删,数据彻底丢失。
_ = os.Remove(nfo) 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", if err := os.Rename(stage, dst); err != nil {
zap.String("path", e), zap.Error(err)) 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 { if o.repo != nil && o.repo.DB != nil {
_ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", e).Delete(&model.Media{}).Error _ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", e).Delete(&model.Media{}).Error
} }
} }
// Move staged file + sidecars into the final path. if backup != "" {
if err := os.Rename(stage, dst); err != nil { // 备份文件即被取代的旧 dst 内容,新文件已成功落位后移除。
cleanup() if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
return err o.log.Warn("organize replace remove backup failed",
zap.String("path", backup), zap.Error(err))
}
} }
moveSidecarRename(nfoPath(stage), nfoPath(dst)) moveSidecarRename(nfoPath(stage), nfoPath(dst))
moveStagedArtwork(stage, dst) moveStagedArtwork(stage, dst)
+29 -4
View File
@@ -291,10 +291,29 @@ func (p *PlaybackService) GetPlaylist(ctx context.Context, playlistID string) (*
return &PlaylistDetail{Playlist: pl, Items: ordered}, nil 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. // 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 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 { Where("playlist_id = ?", playlistID).Count(&count).Error; err != nil {
return err return err
} }
@@ -307,14 +326,20 @@ func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID
} }
// RemoveFromPlaylist 物理删除播放列表项(幂等)。 // 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(). return p.repo.DB.WithContext(ctx).Unscoped().
Where("playlist_id = ? AND media_id = ?", playlistID, mediaID). Where("playlist_id = ? AND media_id = ?", playlistID, mediaID).
Delete(&model.PlaylistItem{}).Error Delete(&model.PlaylistItem{}).Error
} }
// DeletePlaylist 物理删除播放列表及其全部条目。 // 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). if err := p.repo.DB.WithContext(ctx).Unscoped().Where("playlist_id = ?", playlistID).
Delete(&model.PlaylistItem{}).Error; err != nil { Delete(&model.PlaylistItem{}).Error; err != nil {
return err return err
+2
View File
@@ -28,6 +28,8 @@ func ApplyRuntimeSettings(ctx context.Context, cfg *config.Config, repos *reposi
} }
func ApplyRuntimeSetting(cfg *config.Config, key, value string) { func ApplyRuntimeSetting(cfg *config.Config, key, value string) {
config.RuntimeMu.Lock()
defer config.RuntimeMu.Unlock()
if cfg == nil { if cfg == nil {
return return
} }
+26 -1
View File
@@ -16,8 +16,33 @@ func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, er
if _, err := os.Stat(path); err == nil { if _, err := os.Stat(path); err == nil {
return 0, nil // still exists; nothing to remove 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(). res := s.repo.DB.WithContext(ctx).Unscoped().
Where("path = ?", path). Where("id IN ?", ids).
Delete(&model.Media{}) Delete(&model.Media{})
if res.Error == nil && res.RowsAffected > 0 { if res.Error == nil && res.RowsAffected > 0 {
s.invalidateMediaCache(ctx) s.invalidateMediaCache(ctx)
+43 -1
View File
@@ -35,6 +35,13 @@ func (s *ScraperService) Start(ctx context.Context) {
if s == nil { if s == nil {
return 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) go s.queueWorker(ctx)
} }
@@ -69,6 +76,9 @@ func (s *ScraperService) queueWorker(ctx context.Context) {
defer wg.Done() defer wg.Done()
select { select {
case <-ctx.Done(): case <-ctx.Done():
// 任务已被 Claim 置为 running:停机前回写 pending,
// 避免留下永久卡死的任务。
s.requeueClaimedScrapeTask(t)
return return
case sem <- struct{}{}: 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) { func (s *ScraperService) processScrapeTask(ctx context.Context, task *model.ScrapeTask) {
media, err := s.repo.Media.FindByID(ctx, task.MediaID) media, err := s.repo.Media.FindByID(ctx, task.MediaID)
if err != nil || media == nil { if err != nil || media == nil {
@@ -220,8 +239,22 @@ func (s *ScraperService) EnqueueLibrary(ctx context.Context, libraryID string, o
return 0, nil 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)) tasks := make([]model.ScrapeTask, 0, len(rows))
for _, m := range rows { for _, m := range rows {
if activeByMedia != nil && activeByMedia[m.ID] {
continue
}
tasks = append(tasks, model.ScrapeTask{ tasks = append(tasks, model.ScrapeTask{
MediaID: m.ID, MediaID: m.ID,
LibraryID: lib.ID, LibraryID: lib.ID,
@@ -310,7 +343,16 @@ func (s *ScraperService) RetryScrapeTask(ctx context.Context, id string) error {
return errors.New("刮削任务不存在") return errors.New("刮削任务不存在")
} }
if task.Status != model.ScrapeTaskFailed && task.Status != model.ScrapeTaskCanceled { 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.Status = model.ScrapeTaskPending
task.Error = "" task.Error = ""
+43
View File
@@ -14,6 +14,7 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"strconv"
"strings" "strings"
"time" "time"
@@ -371,6 +372,16 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) {
if err != nil || cfg["access_token"] == "" || cfg["refresh_token"] == "" { if err != nil || cfg["access_token"] == "" || cfg["refresh_token"] == "" {
continue 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"]) client := cloud115.NewOpenClient(cfg["app_id"], cfg["access_token"], cfg["refresh_token"])
token, err := client.RefreshToken(cfg["refresh_token"]) token, err := client.RefreshToken(cfg["refresh_token"])
if err != nil { if err != nil {
@@ -395,6 +406,7 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) {
} }
cfg["access_token"] = s.crypto.Encrypt(token.AccessToken) cfg["access_token"] = s.crypto.Encrypt(token.AccessToken)
cfg["refresh_token"] = s.crypto.Encrypt(token.RefreshToken) cfg["refresh_token"] = s.crypto.Encrypt(token.RefreshToken)
cfg["token_refreshed_at"] = strconv.FormatInt(time.Now().Unix(), 10)
enc, err := s.strmAccountConfigJSON(cfg, false) enc, err := s.strmAccountConfigJSON(cfg, false)
if err != nil { if err != nil {
continue 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(启动与设置保存时调用)。 // sync115RelayKey 把设置里的中继密钥同步给 cloud115(启动与设置保存时调用)。
func (s *StrmService) sync115RelayKey(ctx context.Context) { func (s *StrmService) sync115RelayKey(ctx context.Context) {
cloud115.RelayEncryptionKey = s.strmSetting(ctx, Strm115RelayKeySetting) cloud115.RelayEncryptionKey = s.strmSetting(ctx, Strm115RelayKeySetting)
+81 -7
View File
@@ -67,9 +67,17 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
return return
} }
defer s.releaseDownloadSlot(task.Provider) defer s.releaseDownloadSlot(task.Provider)
// 单个任务 panic 不应拖垮整个下载 worker。 // 单个任务 panic 不应拖垮整个下载 worker,且 panic 时任务
// 会永远停在 running:兜底走失败重试路径。
completed := false
helper.Run(s.log, "strm.downloadTask", func() { helper.Run(s.log, "strm.downloadTask", func() {
defer func() {
if !completed {
s.downloadTaskFailWithRetry(task, "任务执行异常中断")
}
}()
s.processDownloadTask(ctx, task) s.processDownloadTask(ctx, task)
completed = true
}) })
}(i) }(i)
} }
@@ -78,9 +86,19 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
} }
// requeueDownloadTask 把已认领但未实际执行的任务退回 pending,避免长期停留在 running。 // requeueDownloadTask 把已认领但未实际执行的任务退回 pending,避免长期停留在 running。
// 退回时必须设置 NextTryAt(WAF 冷却剩余时间):claim 只过滤 next_try_at
// 已过期的任务,不设会让同一批任务被立刻再认领,形成 claim/requeue
// 热循环(占用 SQLite 写锁并饿死上传队列)。
func (s *StrmService) requeueDownloadTask(task *model.StrmDownloadTask) { func (s *StrmService) requeueDownloadTask(task *model.StrmDownloadTask) {
task.Status = model.StrmTaskPending task.Status = model.StrmTaskPending
task.StartedAt = nil 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 { 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)) 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.Status = status
task.Error = message task.Error = message
task.FinishedAt = &now 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)) 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) acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID)
@@ -151,7 +177,19 @@ func (s *StrmService) uploadWorker(ctx context.Context) {
continue continue
} }
for i := range tasks { 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.Status = status
task.Error = message task.Error = message
task.FinishedAt = &now 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)) 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 { if task.Provider == model.StrmProvider115 {
@@ -213,8 +258,15 @@ func (s *StrmService) processUpload115(ctx context.Context, task *model.StrmUplo
task.Status = status task.Status = status
task.Error = message task.Error = message
task.FinishedAt = &now 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)) 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) 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) { if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) {
return 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。 // 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) { if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) {
return 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。 // retryTask 失败状态机:重试次数不足则回 pending 并设置退避时间,否则 failed。
+27
View File
@@ -28,6 +28,7 @@ import (
"github.com/truewhile/MeBox/internal/model" "github.com/truewhile/MeBox/internal/model"
"github.com/truewhile/MeBox/internal/repository" "github.com/truewhile/MeBox/internal/repository"
"github.com/truewhile/MeBox/internal/service/cloud" "github.com/truewhile/MeBox/internal/service/cloud"
"github.com/truewhile/MeBox/internal/service/cloud115"
) )
// strm 全局设置键(存于 Setting 表,strm.* 前缀)。 // strm 全局设置键(存于 Setting 表,strm.* 前缀)。
@@ -169,6 +170,11 @@ func NewStrmService(cfg *config.Config, log *zap.Logger, repos *repository.Conta
// Start 启动下载/上传队列 worker、定时同步巡检、115 token 刷新与队列清理。 // Start 启动下载/上传队列 worker、定时同步巡检、115 token 刷新与队列清理。
func (s *StrmService) Start(ctx context.Context) { func (s *StrmService) Start(ctx context.Context) {
// baseCtx 挂到服务生命周期 ctx 上(Start 由启动流程传入 stopCtx):
// 此前硬编码 context.Background(),Stop() 关 stopCh 后 worker 会退出,
// 但进行中的全量同步(可能持续数小时)完全不受停机控制,优雅停机
// 窗口内仍在批量写库/写盘。
s.baseCtx = ctx
s.sync115RelayKey(ctx) s.sync115RelayKey(ctx)
s.recoverInterruptedSyncs(ctx) s.recoverInterruptedSyncs(ctx)
downloadThreads := s.strmIntSetting(ctx, StrmSettingDownloadThreads, 6) downloadThreads := s.strmIntSetting(ctx, StrmSettingDownloadThreads, 6)
@@ -192,6 +198,18 @@ func (s *StrmService) Start(ctx context.Context) {
for i := 0; i < uploadThreads; i++ { for i := 0; i < uploadThreads; i++ {
helper.Go(s.log, "strm.uploadWorker", func() { s.uploadWorker(ctx) }) 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.cronLoop", func() { s.cronLoop(ctx) })
helper.Go(s.log, "strm.queueCleanupLoop", func() { s.queueCleanupLoop(ctx) }) helper.Go(s.log, "strm.queueCleanupLoop", func() { s.queueCleanupLoop(ctx) })
helper.Go(s.log, "strm.refresh115TokensLoop", func() { s.refresh115TokensLoop(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 { if err != nil {
return nil, err 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 return provider, nil
} }
+101 -37
View File
@@ -6,6 +6,7 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"reflect"
"net/url" "net/url"
"os" "os"
"path/filepath" "path/filepath"
@@ -97,7 +98,8 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...
StartedAt: &now, StartedAt: &now,
} }
if err := s.repo.StrmSyncRecord.Create(ctx, rec); err != nil { if err := s.repo.StrmSyncRecord.Create(ctx, rec); err != nil {
s.clearRunning(pathID) s.clearRunning(pathID, cancel)
cancel()
return err return err
} }
status := model.StrmSyncRecordRunning status := model.StrmSyncRecordRunning
@@ -106,7 +108,7 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...
p.LastSyncMessage = "同步进行中" p.LastSyncMessage = "同步进行中"
_ = s.repo.StrmSyncPath.Update(ctx, p) _ = 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 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 { func (s *StrmService) CancelSync(ctx context.Context, pathID string) error {
s.mu.Lock() s.mu.Lock()
cancel, exists := s.running[pathID] cancel, exists := s.running[pathID]
if exists {
delete(s.running, pathID)
}
s.mu.Unlock() s.mu.Unlock()
// 不在这里预删 running 标记:runSync 退出时的 clearRunning 会按
// cancel 身份校验后删除,避免旧同步收尾误删新同步的标记。
if exists && cancel != nil { if exists && cancel != nil {
cancel() cancel()
} }
@@ -142,12 +143,24 @@ func (s *StrmService) IsSyncRunning(pathID string) bool {
return exists return exists
} }
func (s *StrmService) clearRunning(pathID string) { // clearRunning 清除同步的运行标记;仅当 map 中登记的 cancel 与本次同步
// 一致时才删除,防止慢收尾的旧同步把随后启动的新同步标记误删掉。
func (s *StrmService) clearRunning(pathID string, cancel context.CancelFunc) {
s.mu.Lock() 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() 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 列出网盘账号某目录下的条目(供前端目录选择器使用)。 // ListRemoteDir 列出网盘账号某目录下的条目(供前端目录选择器使用)。
func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string) ([]cloud.FileEntry, error) { func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string) ([]cloud.FileEntry, error) {
acct, err := s.repo.StrmAccount.FindByID(ctx, accountID) acct, err := s.repo.StrmAccount.FindByID(ctx, accountID)
@@ -169,8 +182,9 @@ func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string)
} }
// runSync 执行同步主体;结束时更新记录与目录状态。 // runSync 执行同步主体;结束时更新记录与目录状态。
func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord) { func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord, cancel context.CancelFunc) {
defer s.clearRunning(p.ID) defer s.clearRunning(p.ID, cancel)
defer cancel()
cfg, err := s.strmEffectiveConfig(ctx, p) cfg, err := s.strmEffectiveConfig(ctx, p)
if err != nil { if err != nil {
@@ -334,24 +348,34 @@ func (st *strmSyncState) walkRemote() error {
ctx, cancel := context.WithCancel(st.ctx) ctx, cancel := context.WithCancel(st.ctx)
defer cancel() defer cancel()
queue := make(chan dirTask, 512) // 工作队列用「互斥锁 + 条件变量 + 动态 slice」实现,而不是有界
var pending atomic.Int64 // 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) push(dirTask{id: root, rel: ""})
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)
}
}()
var ( var (
wg sync.WaitGroup wg sync.WaitGroup
@@ -362,11 +386,27 @@ func (st *strmSyncState) walkRemote() error {
wg.Add(1) wg.Add(1)
go func() { go func() {
defer wg.Done() defer wg.Done()
// worker 解析远端响应 panic 时取消整个同步,让 closer 与其余 // worker 解析远端响应 panic 时取消整个同步,让其余 worker
// worker 正常收尾,避免队列与 pending 计数卡死;正常退出不取消。 // 正常收尾;正常退出不取消。
if err := helper.Recover(st.s.log, "strm.sync.walkRemote", func() error { 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 { if ctx.Err() != nil {
walkMu.Lock()
pending--
walkCond.Broadcast()
walkMu.Unlock()
return nil return nil
} }
entries, err := st.provider.List(ctx, task.id) 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) firstErr = fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err)
} }
errMu.Unlock() errMu.Unlock()
walkMu.Lock()
pending--
walkCond.Broadcast()
walkMu.Unlock()
cancel() cancel()
return nil return nil
} }
@@ -386,19 +430,18 @@ func (st *strmSyncState) walkRemote() error {
rel = task.rel + "/" + cleanName rel = task.rel + "/" + cleanName
} }
if entry.IsDir { if entry.IsDir {
pending.Add(1) push(dirTask{id: entry.ID, rel: rel})
select {
case queue <- dirTask{id: entry.ID, rel: rel}:
case <-ctx.Done():
pending.Add(-1)
}
} else { } else {
st.processRemoteFile(entry, rel) st.processRemoteFile(entry, rel)
} }
} }
pending.Add(-1) walkMu.Lock()
pending--
if pending == 0 {
walkCond.Broadcast()
}
walkMu.Unlock()
} }
return nil
}); err != nil { }); err != nil {
cancel() cancel()
} }
@@ -1329,6 +1372,10 @@ func (st *strmSyncState) updateSyncMessage(msg string) {
func (s *StrmService) cronLoop(ctx context.Context) { func (s *StrmService) cronLoop(ctx context.Context) {
ticker := time.NewTicker(60 * time.Second) ticker := time.NewTicker(60 * time.Second)
defer ticker.Stop() defer ticker.Stop()
// 记录上次检查到的分钟:一轮循环若被慢操作拖过 60s(远端 List 慢、
// 串行 StartSync、DB 忙),ticker 会丢掉中间的 tick,命中排程的分钟
// 若只按"当前分钟相等"判定就会被静默跳过。逐分钟回放补触。
last := time.Now().Truncate(time.Minute)
for { for {
select { select {
case <-ctx.Done(): case <-ctx.Done():
@@ -1336,8 +1383,18 @@ func (s *StrmService) cronLoop(ctx context.Context) {
case <-s.stopCh: case <-s.stopCh:
return return
case now := <-ticker.C: case now := <-ticker.C:
now = now.Truncate(time.Minute)
paths, err := s.repo.StrmSyncPath.List(ctx) paths, err := s.repo.StrmSyncPath.List(ctx)
if err != nil { 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 continue
} }
for i := range paths { for i := range paths {
@@ -1345,7 +1402,14 @@ func (s *StrmService) cronLoop(ctx context.Context) {
if !p.Enabled || !p.EnableCron || strings.TrimSpace(p.Cron) == "" { if !p.Enabled || !p.EnableCron || strings.TrimSpace(p.Cron) == "" {
continue continue
} }
if !cronMatches(p.Cron, now) { matched := false
for _, m := range due {
if cronMatches(p.Cron, m) {
matched = true
break
}
}
if !matched {
continue continue
} }
s.mu.Lock() s.mu.Lock()
+19 -2
View File
@@ -279,12 +279,29 @@ func (w *WatcherService) process(ctx context.Context, d duePath) {
if removed, derr := w.scanner.RemovePath(ctx, d.path); derr != nil { 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)) w.log.Warn("watcher remove failed", zap.String("path", d.path), zap.Error(derr))
} else if removed > 0 { } 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 return
} }
if fi.IsDir() { 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 { 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)) w.log.Warn("watcher ingest failed", zap.String("path", d.path), zap.Error(ierr))
+7
View File
@@ -92,6 +92,13 @@ func (h *Hub) Subscribe(id string, topics []string) *Subscriber {
sub.topics[t] = struct{}{} sub.topics[t] = struct{}{}
} }
h.mu.Lock() 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.subs[id] = sub
h.mu.Unlock() h.mu.Unlock()
return sub return sub
+7 -2
View File
@@ -114,26 +114,31 @@ function LayoutHeaderSearch() {
const [loading, setLoading] = useState(false) const [loading, setLoading] = useState(false)
const [results, setResults] = useState<Media[]>([]) const [results, setResults] = useState<Media[]>([])
const containerRef = useRef<HTMLDivElement>(null) const containerRef = useRef<HTMLDivElement>(null)
// 递增序号守卫:快速连续输入时丢弃过期响应
const searchSeqRef = useRef(0)
const navigate = useNavigate() const navigate = useNavigate()
useEffect(() => { useEffect(() => {
const trimmed = query.trim() const trimmed = query.trim()
if (!trimmed) { if (!trimmed) {
searchSeqRef.current += 1
setResults([]) setResults([])
setLoading(false) setLoading(false)
return return
} }
const seq = ++searchSeqRef.current
setLoading(true) setLoading(true)
const timer = setTimeout(async () => { const timer = setTimeout(async () => {
try { try {
const res = await mediaAPI.search(trimmed, 8) const res = await mediaAPI.search(trimmed, 8)
if (seq !== searchSeqRef.current) return
setResults(res.items || []) setResults(res.items || [])
setIsOpen(true) setIsOpen(true)
} catch { } catch {
setResults([]) // 请求失败时保留旧结果,避免网络抖动清空下拉
} finally { } finally {
setLoading(false) if (seq === searchSeqRef.current) setLoading(false)
} }
}, 250) }, 250)
-123
View File
@@ -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 <div>SSE Demo</div>
* }
* ```
*/
export function useSSE(
onEvent: SSEEventHandler,
options: { autoConnect?: boolean } = {}
) {
const { autoConnect = true } = options
const onEventRef = useRef(onEvent)
onEventRef.current = onEvent
const eventSourceRef = useRef<EventSource | null>(null)
const reconnectTimeoutRef = useRef<ReturnType<typeof setTimeout> | 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,
}
}
+12 -5
View File
@@ -2,12 +2,15 @@ import { useEffect, useRef } from 'react'
import { useAuthStore } from '../stores/auth' 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 // useWebSocket opens a single connection to /api/ws and dispatches every
// message to the supplied handler. Auto-reconnects with back-off while the // 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 // auth token is present; after the fast retries are exhausted it keeps a
// cannot create an endless /api/ws 401 loop. // slow 60s retry loop instead of giving up permanently.
export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) { export function useWebSocket(onEvent: (topic: string, payload: unknown) => void) {
const ref = useRef<WebSocket | null>(null) const ref = useRef<WebSocket | null>(null)
const token = useAuthStore((s) => s.token) const token = useAuthStore((s) => s.token)
@@ -22,7 +25,6 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void)
const open = () => { const open = () => {
if (closed) return if (closed) return
if (reconnectAttempts >= MAX_RECONNECT_ATTEMPTS) return
const proto = window.location.protocol === 'https:' ? 'wss:' : 'ws:' const proto = window.location.protocol === 'https:' ? 'wss:' : 'ws:'
const url = `${proto}//${window.location.host}/api/ws?token=${encodeURIComponent(token)}` const url = `${proto}//${window.location.host}/api/ws?token=${encodeURIComponent(token)}`
const ws = new WebSocket(url) const ws = new WebSocket(url)
@@ -43,7 +45,12 @@ export function useWebSocket(onEvent: (topic: string, payload: unknown) => void)
ws.onclose = () => { ws.onclose = () => {
if (closed) return if (closed) return
reconnectAttempts += 1 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) timer = window.setTimeout(open, delay)
} }
} }
+1
View File
@@ -14,6 +14,7 @@ export function AdminLibraryPanel() {
coverURL={createForm.coverURL} coverURL={createForm.coverURL}
roots={createForm.roots} roots={createForm.roots}
createPerSubfolder={createForm.createPerSubfolder} createPerSubfolder={createForm.createPerSubfolder}
creating={createForm.creating}
onNameChange={createForm.setName} onNameChange={createForm.setName}
onTypeChange={createForm.setType} onTypeChange={createForm.setType}
onCoverURLChange={createForm.setCoverURL} onCoverURLChange={createForm.setCoverURL}
+6 -3
View File
@@ -1,5 +1,5 @@
import { FormEvent, useState } from 'react' 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 { LocalDirBrowserDialog } from '../components/LocalDirBrowserDialog'
import type { RootDraft } from './adminLibraryPanelModel' import type { RootDraft } from './adminLibraryPanelModel'
@@ -10,6 +10,7 @@ type CreateFormProps = {
coverURL: string coverURL: string
roots: RootDraft[] roots: RootDraft[]
createPerSubfolder: boolean createPerSubfolder: boolean
creating: boolean
onNameChange: (value: string) => void onNameChange: (value: string) => void
onTypeChange: (value: string) => void onTypeChange: (value: string) => void
onCoverURLChange: (value: string) => void onCoverURLChange: (value: string) => void
@@ -26,6 +27,7 @@ export function AdminLibraryCreateForm({
coverURL, coverURL,
roots, roots,
createPerSubfolder, createPerSubfolder,
creating,
onNameChange, onNameChange,
onTypeChange, onTypeChange,
onCoverURLChange, onCoverURLChange,
@@ -108,8 +110,9 @@ export function AdminLibraryCreateForm({
批处理模式:仅取上方第一个路径作为父级目录,会为其中每个子文件夹分别创建媒体库,可自选类型用于整体推断。 批处理模式:仅取上方第一个路径作为父级目录,会为其中每个子文件夹分别创建媒体库,可自选类型用于整体推断。
</p> </p>
)} )}
<button type="submit" className="neon-button md:col-span-4"> <button type="submit" className="neon-button md:col-span-4 disabled:opacity-50" disabled={creating}>
{createPerSubfolder ? '按目录批量创建' : '新建 / 追加路径'} {creating && <Loader2 size={16} className="animate-spin" />}
{creating ? '创建中…' : createPerSubfolder ? '按目录批量创建' : '新建 / 追加路径'}
</button> </button>
</form> </form>
+10 -4
View File
@@ -1,11 +1,12 @@
import { FormEvent } from 'react' import { FormEvent } from 'react'
import { Plus, Save } from 'lucide-react' import { Loader2, Plus, Save } from 'lucide-react'
type AdminUsersFormProps = { type AdminUsersFormProps = {
usersCount: number usersCount: number
maxUsers: number maxUsers: number
maxUsersDraft: string maxUsersDraft: string
savingLimit: boolean savingLimit: boolean
creating: boolean
username: string username: string
password: string password: string
userLimitReached: boolean userLimitReached: boolean
@@ -21,6 +22,7 @@ export function AdminUsersForm({
maxUsers, maxUsers,
maxUsersDraft, maxUsersDraft,
savingLimit, savingLimit,
creating,
username, username,
password, password,
userLimitReached, userLimitReached,
@@ -91,9 +93,13 @@ export function AdminUsersForm({
onChange={(e) => onPasswordChange(e.target.value)} onChange={(e) => onPasswordChange(e.target.value)}
disabled={userLimitReached} disabled={userLimitReached}
/> />
<button type="submit" className="neon-button inline-flex items-center justify-center gap-2" disabled={userLimitReached}> <button
<Plus size={16} /> type="submit"
添加用户 className="neon-button inline-flex items-center justify-center gap-2 disabled:opacity-50"
disabled={userLimitReached || creating}
>
{creating ? <Loader2 size={16} className="animate-spin" /> : <Plus size={16} />}
{creating ? '添加中…' : '添加用户'}
</button> </button>
</form> </form>
) )
+7 -1
View File
@@ -16,6 +16,7 @@ export function AdminUsersPanel() {
const [maxUsers, setMaxUsers] = useState(DEFAULT_MAX_USERS) const [maxUsers, setMaxUsers] = useState(DEFAULT_MAX_USERS)
const [maxUsersDraft, setMaxUsersDraft] = useState(String(DEFAULT_MAX_USERS)) const [maxUsersDraft, setMaxUsersDraft] = useState(String(DEFAULT_MAX_USERS))
const [savingLimit, setSavingLimit] = useState(false) const [savingLimit, setSavingLimit] = useState(false)
const [creating, setCreating] = useState(false)
const [username, setUsername] = useState('') const [username, setUsername] = useState('')
const [password, setPassword] = useState('') const [password, setPassword] = useState('')
const [editingID, setEditingID] = useState<string | null>(null) const [editingID, setEditingID] = useState<string | null>(null)
@@ -69,18 +70,22 @@ export function AdminUsersPanel() {
const handleCreate = async (e: FormEvent) => { const handleCreate = async (e: FormEvent) => {
e.preventDefault() e.preventDefault()
if (creating) return
setCreating(true)
try { try {
await adminAPI.createUser({ username, password }) await adminAPI.createUser({ username, password })
toast.success('用户已添加,默认仅允许浏览与播放媒体') toast.success('用户已添加,默认仅允许浏览与播放媒体')
setUsername('') setUsername('')
setPassword('') setPassword('')
await refresh()
} catch (err: unknown) { } catch (err: unknown) {
const msg = const msg =
userCreateErrorMessage(err) ?? userCreateErrorMessage(err) ??
'添加用户失败' '添加用户失败'
toast.error(msg) toast.error(msg)
} finally {
setCreating(false)
} }
await refresh().catch(() => undefined)
} }
const startEdit = (u: User) => { const startEdit = (u: User) => {
@@ -171,6 +176,7 @@ export function AdminUsersPanel() {
maxUsers={maxUsers} maxUsers={maxUsers}
maxUsersDraft={maxUsersDraft} maxUsersDraft={maxUsersDraft}
savingLimit={savingLimit} savingLimit={savingLimit}
creating={creating}
username={username} username={username}
password={password} password={password}
userLimitReached={userLimitReached} userLimitReached={userLimitReached}
+21
View File
@@ -37,6 +37,7 @@ export function AdultSettingsPanel() {
const [values, setValues] = useState<Record<string, string>>({}) const [values, setValues] = useState<Record<string, string>>({})
const [dirty, setDirty] = useState<Set<string>>(new Set()) const [dirty, setDirty] = useState<Set<string>>(new Set())
const [loading, setLoading] = useState(true) const [loading, setLoading] = useState(true)
const [loadError, setLoadError] = useState('')
const [saving, setSaving] = useState(false) const [saving, setSaving] = useState(false)
const [libraries, setLibraries] = useState<Library[]>([]) const [libraries, setLibraries] = useState<Library[]>([])
const [showToken, setShowToken] = useState(false) const [showToken, setShowToken] = useState(false)
@@ -79,6 +80,13 @@ export function AdultSettingsPanel() {
setValues(idx) setValues(idx)
setLibraries(libs as Library[]) setLibraries(libs as Library[])
setDirty(new Set()) setDirty(new Set())
setLoadError('')
} catch (err: unknown) {
// 加载失败时保留错误态,避免把表单默认值误当成已保存的配置
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '加载成人设置失败'
setLoadError(msg)
toast.error(msg)
} finally { } finally {
setLoading(false) setLoading(false)
} }
@@ -196,6 +204,19 @@ export function AdultSettingsPanel() {
) )
} }
if (loadError) {
return (
<div className="glass-panel flex flex-col items-center gap-3 py-10 text-center">
<XCircle className="text-rose-500" size={24} />
<p className="text-sm text-ink-100">成人设置加载失败:{loadError}</p>
<p className="text-xs text-sand-500">当前展示的并非已保存配置,请重新加载后再修改</p>
<button type="button" onClick={() => refresh()} className="neon-button !px-4 !py-1.5 !text-xs">
重试
</button>
</div>
)
}
return ( return (
<form onSubmit={onSave} className="space-y-6"> <form onSubmit={onSave} className="space-y-6">
{/* 1. 全局访问与隔离卡片 */} {/* 1. 全局访问与隔离卡片 */}
+35 -11
View File
@@ -13,6 +13,7 @@ import type { Media } from '../types'
import { getSeriesKey, seriesTitleFromPath } from '../utils/groupSeries' import { getSeriesKey, seriesTitleFromPath } from '../utils/groupSeries'
import { isRemoteEmbyID } from '../utils/remoteEmby' import { isRemoteEmbyID } from '../utils/remoteEmby'
import { pickPlayerMode, needsTranscodeForBrowser, isDirectStreamMedia, type PlayerMode } from './playerPageModel' import { pickPlayerMode, needsTranscodeForBrowser, isDirectStreamMedia, type PlayerMode } from './playerPageModel'
import { apiErrorMessage } from './StrmManagePage'
import { PlayerTopBar } from './PlayerTopBar' import { PlayerTopBar } from './PlayerTopBar'
import { PlayerVideoStage } from './PlayerVideoStage' import { PlayerVideoStage } from './PlayerVideoStage'
import { PlayerDanmakuPanel } from '../components/PlayerDanmakuPanel' import { PlayerDanmakuPanel } from '../components/PlayerDanmakuPanel'
@@ -62,6 +63,8 @@ export function PlayerPage() {
const [subtitleIndex, setSubtitleIndex] = useState<number>(initialSubtitleIndex) const [subtitleIndex, setSubtitleIndex] = useState<number>(initialSubtitleIndex)
const [hlsUnavailable, setHlsUnavailable] = useState(false) const [hlsUnavailable, setHlsUnavailable] = useState(false)
const [playerError, setPlayerError] = useState('') const [playerError, setPlayerError] = useState('')
// 媒体元数据加载失败(404 / 无权限等):舞台区直接展示错误而不是永远「加载中」
const [loadError, setLoadError] = useState('')
// 「客户端直连解码」模式:宿主机不转码,播放器强制 direct play、隐藏 HLS 切换。 // 「客户端直连解码」模式:宿主机不转码,播放器强制 direct play、隐藏 HLS 切换。
const [directOnly, setDirectOnly] = useState(false) const [directOnly, setDirectOnly] = useState(false)
const [resumePosition, setResumePosition] = useState(0) const [resumePosition, setResumePosition] = useState(0)
@@ -176,6 +179,7 @@ export function PlayerPage() {
// 切换视频时重置媒体与弹幕状态,确保新视频自动重新识别并加载弹幕 // 切换视频时重置媒体与弹幕状态,确保新视频自动重新识别并加载弹幕
useEffect(() => { useEffect(() => {
setMedia(null) setMedia(null)
setLoadError('')
setDanmakuEpisodeId(null) setDanmakuEpisodeId(null)
setDanmakuCandidates([]) setDanmakuCandidates([])
setDanmakuSearch(null) setDanmakuSearch(null)
@@ -184,29 +188,48 @@ export function PlayerPage() {
setDanmakuSearching(true) setDanmakuSearching(true)
}, [id]) }, [id])
// 依赖收敛为 mode 参数的字符串值:避免 params 对象引用每次变化都重复拉取元数据
const modeParam = params.get('mode') as PlayerMode | null
// Load metadata and pick a default mode. // Load metadata and pick a default mode.
useEffect(() => { useEffect(() => {
if (!id) return if (!id) return
mediaAPI.get(id).then((m) => { let cancelled = false
setMedia(m) mediaAPI
const isDirect = isDirectStreamMedia(m) .get(id)
const forced = params.get('mode') as PlayerMode | null .then((m) => {
const auto = pickPlayerMode(m) if (cancelled) return
// 直连解码模式以及 STRM / Emby 挂载等直连媒体,忽略 ?mode=hls,始终 direct play。 setMedia(m)
setMode(directOnly || isDirect ? 'direct' : (forced ?? auto)) const isDirect = isDirectStreamMedia(m)
setPlayerError('') 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 subtitlesAPI
.list(id) .list(id)
.then((tracks) => { .then((tracks) => {
if (cancelled) return
const list = tracks ?? [] const list = tracks ?? []
setSubs(list) setSubs(list)
// 记忆的轨道下标可能超出当前媒体的轨道数(不同媒体字幕数量不同), // 记忆的轨道下标可能超出当前媒体的轨道数(不同媒体字幕数量不同),
// 越界时回退到第一条;无字幕则关闭。 // 越界时回退到第一条;无字幕则关闭。
setSubtitleIndex((cur) => (cur >= list.length ? (list.length > 0 ? 0 : -1) : cur)) setSubtitleIndex((cur) => (cur >= list.length ? (list.length > 0 ? 0 : -1) : cur))
}) })
.catch(() => setSubs([])) .catch(() => {
}, [id, params, directOnly]) if (cancelled) return
setSubs([])
})
return () => {
cancelled = true
}
}, [id, modeParam, directOnly])
// Wire up the actual <video> element when we know the mode. // Wire up the actual <video> element when we know the mode.
useEffect(() => { useEffect(() => {
@@ -544,6 +567,7 @@ export function PlayerPage() {
/> />
<PlayerVideoStage <PlayerVideoStage
media={media} media={media}
loadError={loadError}
playerError={playerError} playerError={playerError}
subs={subs} subs={subs}
subtitleIndex={subtitleIndex} subtitleIndex={subtitleIndex}
+5
View File
@@ -9,6 +9,8 @@ import { PlayerControls } from '../components/PlayerControls'
type PlayerVideoStageProps = { type PlayerVideoStageProps = {
media: Media | null media: Media | null
/** 媒体元数据加载失败提示(非空时替代「加载中」展示)。 */
loadError?: string
playerError: string playerError: string
subs: SubtitleTrack[] subs: SubtitleTrack[]
/** 当前激活字幕轨道:-1=关闭,0..n-1=对应轨道。 */ /** 当前激活字幕轨道:-1=关闭,0..n-1=对应轨道。 */
@@ -43,6 +45,7 @@ type PlayerVideoStageProps = {
export function PlayerVideoStage({ export function PlayerVideoStage({
media, media,
loadError,
playerError, playerError,
subs, subs,
subtitleIndex, subtitleIndex,
@@ -317,6 +320,8 @@ export function PlayerVideoStage({
{danmakuPanel} {danmakuPanel}
{playlistPanel} {playlistPanel}
</> </>
) : loadError ? (
<p className="max-w-[92%] text-center text-sm text-rose-400">{loadError}</p>
) : ( ) : (
<p className="text-sand-500">加载中…</p> <p className="text-sand-500">加载中…</p>
)} )}
+22 -8
View File
@@ -1,4 +1,4 @@
import { useCallback, useEffect, useMemo, useState, type ReactNode } from 'react' import { useCallback, useEffect, useMemo, useRef, useState, type ReactNode } from 'react'
import toast from 'react-hot-toast' import toast from 'react-hot-toast'
import { import {
AlertCircle, AlertCircle,
@@ -56,13 +56,20 @@ export function ScraperQueuePage({ embedded = false }: { embedded?: boolean }) {
const [batchBusy, setBatchBusy] = useState(false) const [batchBusy, setBatchBusy] = useState(false)
const { selectedIds, setSelectedIds, reset: clearSelection, toggleRow: toggleSelectRow, toggleAll: toggleAllIds } = useTaskSelection() const { selectedIds, setSelectedIds, reset: clearSelection, toggleRow: toggleSelectRow, toggleAll: toggleAllIds } = useTaskSelection()
const [detailTask, setDetailTask] = useState<ScrapeTask | null>(null) const [detailTask, setDetailTask] = useState<ScrapeTask | null>(null)
// 轮询/翻页/切筛选并发时用递增序号丢弃过期响应
const refreshSeqRef = useRef(0)
const [pollFailures, setPollFailures] = useState(0)
const refresh = useCallback( const refresh = useCallback(
async (showLoading = false) => { async (showLoading = false) => {
if (showLoading) setIsRefreshing(true) if (showLoading) setIsRefreshing(true)
const seq = ++refreshSeqRef.current
try { try {
const status = filter === 'all' ? undefined : filter const status = filter === 'all' ? undefined : filter
const data = await scraperAPI.queue(status, page, PAGE_SIZE) const data = await scraperAPI.queue(status, page, PAGE_SIZE)
// 序号不符说明已有更新的请求发出(翻页/切筛选/轮询并发),丢弃旧响应
if (seq !== refreshSeqRef.current) return
setPollFailures(0)
const tp = Math.max(1, Math.ceil((data.total ?? data.tasks.length) / PAGE_SIZE)) const tp = Math.max(1, Math.ceil((data.total ?? data.tasks.length) / PAGE_SIZE))
if (page > tp) { if (page > tp) {
setPage(tp) setPage(tp)
@@ -71,7 +78,8 @@ export function ScraperQueuePage({ embedded = false }: { embedded?: boolean }) {
setTotalPages(tp) setTotalPages(tp)
setSnapshot(data) setSnapshot(data)
} catch { } catch {
/* keep existing */ // 保留旧数据;连续失败 ≥3 次时页头徽标切换为「连接失败,重试中」
if (seq === refreshSeqRef.current) setPollFailures((n) => n + 1)
} finally { } finally {
setLoading(false) setLoading(false)
if (showLoading) setIsRefreshing(false) if (showLoading) setIsRefreshing(false)
@@ -232,12 +240,18 @@ export function ScraperQueuePage({ embedded = false }: { embedded?: boolean }) {
<div> <div>
<div className="flex items-center gap-2"> <div className="flex items-center gap-2">
<h1 className="font-display text-2xl font-bold text-ink-600 sm:text-3xl">刮削队列</h1> <h1 className="font-display text-2xl font-bold text-ink-600 sm:text-3xl">刮削队列</h1>
{autoRefresh && ( {autoRefresh &&
<span className="inline-flex items-center gap-1 rounded-full border border-emerald-300/40 bg-emerald-500/10 px-2 py-0.5 text-[11px] font-semibold text-emerald-600"> (pollFailures >= 3 ? (
<span className="h-1.5 w-1.5 animate-pulse rounded-full bg-emerald-500" /> <span className="inline-flex items-center gap-1 rounded-full border border-rose-300/40 bg-rose-500/10 px-2 py-0.5 text-[11px] font-semibold text-rose-600">
实时同步 <span className="h-1.5 w-1.5 animate-pulse rounded-full bg-rose-500" />
</span> 连接失败,重试中
)} </span>
) : (
<span className="inline-flex items-center gap-1 rounded-full border border-emerald-300/40 bg-emerald-500/10 px-2 py-0.5 text-[11px] font-semibold text-emerald-600">
<span className="h-1.5 w-1.5 animate-pulse rounded-full bg-emerald-500" />
实时同步
</span>
))}
</div> </div>
<p className="text-xs text-sand-500 mt-0.5"> <p className="text-xs text-sand-500 mt-0.5">
媒体元数据在线识别与海报/剧照下载进度(TMDb / 豆瓣 / Bangumi / TheTVDB) 媒体元数据在线识别与海报/剧照下载进度(TMDb / 豆瓣 / Bangumi / TheTVDB)
+6 -1
View File
@@ -899,18 +899,23 @@ function StrmDirBrowserDialog({
const [crumbs, setCrumbs] = useState<string[]>([]) const [crumbs, setCrumbs] = useState<string[]>([])
const [entries, setEntries] = useState<StrmRemoteEntry[]>([]) const [entries, setEntries] = useState<StrmRemoteEntry[]>([])
const [loading, setLoading] = useState(true) const [loading, setLoading] = useState(true)
// 递增序号守卫:快速连续进入目录时丢弃过期目录响应
const loadSeqRef = useRef(0)
const load = async (target: string) => { const load = async (target: string) => {
const seq = ++loadSeqRef.current
setLoading(true) setLoading(true)
try { try {
const list = await strmAPI.listRemoteDir(accountId, target) const list = await strmAPI.listRemoteDir(accountId, target)
if (seq !== loadSeqRef.current) return
setEntries(list) setEntries(list)
setDir(target) setDir(target)
setCrumbs(target ? target.split('/').filter(Boolean) : []) setCrumbs(target ? target.split('/').filter(Boolean) : [])
} catch (err) { } catch (err) {
if (seq !== loadSeqRef.current) return
toast.error(apiErrorMessage(err)) toast.error(apiErrorMessage(err))
} finally { } finally {
setLoading(false) if (seq === loadSeqRef.current) setLoading(false)
} }
} }
+22 -8
View File
@@ -1,4 +1,4 @@
import { useCallback, useEffect, useMemo, useState } from 'react' import { useCallback, useEffect, useMemo, useRef, useState } from 'react'
import toast from 'react-hot-toast' import toast from 'react-hot-toast'
import { import {
AlertCircle, AlertCircle,
@@ -53,6 +53,9 @@ export function StrmQueuePanel({
const [batchBusy, setBatchBusy] = useState(false) const [batchBusy, setBatchBusy] = useState(false)
const { selectedIds, setSelectedIds, reset: clearSelection, toggleRow: toggleSelectRow, toggleAll: toggleAllIds } = useTaskSelection() const { selectedIds, setSelectedIds, reset: clearSelection, toggleRow: toggleSelectRow, toggleAll: toggleAllIds } = useTaskSelection()
const [detailTask, setDetailTask] = useState<StrmTask | null>(null) const [detailTask, setDetailTask] = useState<StrmTask | null>(null)
// 轮询/翻页/切筛选并发时用递增序号丢弃过期响应
const refreshSeqRef = useRef(0)
const [pollFailures, setPollFailures] = useState(0)
const isDownload = kind === 'download' const isDownload = kind === 'download'
const Icon = isDownload ? Download : Upload const Icon = isDownload ? Download : Upload
@@ -60,11 +63,15 @@ export function StrmQueuePanel({
const refresh = useCallback( const refresh = useCallback(
async (showLoading = false) => { async (showLoading = false) => {
if (showLoading) setIsRefreshing(true) if (showLoading) setIsRefreshing(true)
const seq = ++refreshSeqRef.current
try { try {
const status = filter === 'all' ? undefined : filter const status = filter === 'all' ? undefined : filter
const data = isDownload const data = isDownload
? await strmAPI.downloads(status, page, PAGE_SIZE) ? await strmAPI.downloads(status, page, PAGE_SIZE)
: await strmAPI.uploads(status, page, PAGE_SIZE) : await strmAPI.uploads(status, page, PAGE_SIZE)
// 序号不符说明已有更新的请求发出(翻页/切筛选/轮询并发),丢弃旧响应
if (seq !== refreshSeqRef.current) return
setPollFailures(0)
const tp = Math.max(1, Math.ceil((data.total ?? data.tasks.length) / PAGE_SIZE)) const tp = Math.max(1, Math.ceil((data.total ?? data.tasks.length) / PAGE_SIZE))
if (page > tp) { if (page > tp) {
setPage(tp) setPage(tp)
@@ -73,7 +80,8 @@ export function StrmQueuePanel({
setTotalPages(tp) setTotalPages(tp)
setSnapshot(data) setSnapshot(data)
} catch { } catch {
/* keep existing data */ // 保留旧数据;连续失败 ≥3 次时页头徽标切换为「连接失败,重试中」
if (seq === refreshSeqRef.current) setPollFailures((n) => n + 1)
} finally { } finally {
setLoading(false) setLoading(false)
if (showLoading) setIsRefreshing(false) if (showLoading) setIsRefreshing(false)
@@ -227,12 +235,18 @@ export function StrmQueuePanel({
<h1 className="font-display text-2xl font-bold text-ink-600 sm:text-3xl"> <h1 className="font-display text-2xl font-bold text-ink-600 sm:text-3xl">
{isDownload ? '下载队列' : '上传队列'} {isDownload ? '下载队列' : '上传队列'}
</h1> </h1>
{autoRefresh && ( {autoRefresh &&
<span className="inline-flex items-center gap-1 rounded-full border border-emerald-300/40 bg-emerald-500/10 px-2 py-0.5 text-[11px] font-semibold text-emerald-600"> (pollFailures >= 3 ? (
<span className="h-1.5 w-1.5 animate-pulse rounded-full bg-emerald-500" /> <span className="inline-flex items-center gap-1 rounded-full border border-rose-300/40 bg-rose-500/10 px-2 py-0.5 text-[11px] font-semibold text-rose-600">
实时同步 <span className="h-1.5 w-1.5 animate-pulse rounded-full bg-rose-500" />
</span> 连接失败,重试中
)} </span>
) : (
<span className="inline-flex items-center gap-1 rounded-full border border-emerald-300/40 bg-emerald-500/10 px-2 py-0.5 text-[11px] font-semibold text-emerald-600">
<span className="h-1.5 w-1.5 animate-pulse rounded-full bg-emerald-500" />
实时同步
</span>
))}
</div> </div>
<p className="text-xs text-sand-500 mt-0.5"> <p className="text-xs text-sand-500 mt-0.5">
{isDownload {isDownload
+10
View File
@@ -29,6 +29,10 @@ export function WatchHistoryPage() {
historyAPI historyAPI
.list(200) .list(200)
.then(setItems) .then(setItems)
.catch((err: unknown) => {
const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '加载观看历史失败'
toast.error(msg)
})
.finally(() => setLoading(false)) .finally(() => setLoading(false))
} }
@@ -41,6 +45,9 @@ export function WatchHistoryPage() {
await historyAPI.remove(id) await historyAPI.remove(id)
setItems((prev) => prev.filter((item) => item.id !== id)) setItems((prev) => prev.filter((item) => item.id !== id))
toast.success('已移除观看历史') toast.success('已移除观看历史')
} catch (err: unknown) {
const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '移除观看历史失败'
toast.error(msg)
} finally { } finally {
setBusy('') setBusy('')
} }
@@ -54,6 +61,9 @@ export function WatchHistoryPage() {
await historyAPI.clear(undefined, status) await historyAPI.clear(undefined, status)
setItems((prev) => prev.filter((item) => status === 'completed' ? !item.completed : item.completed)) setItems((prev) => prev.filter((item) => status === 'completed' ? !item.completed : item.completed))
toast.success(`已清除${label}记录`) toast.success(`已清除${label}记录`)
} catch (err: unknown) {
const msg = (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? `清除${label}记录失败`
toast.error(msg)
} finally { } finally {
setBusy('') setBusy('')
} }
@@ -158,16 +158,30 @@ export function Strm115AuthPanel({
: { sessionId: result.session_id, accountId: account.id, mode: 'url', authUrl: result.auth_url }, : { sessionId: result.session_id, accountId: account.id, mode: 'url', authUrl: result.auth_url },
) )
stopPolling() stopPolling()
// confirmed / expired 只处理一次,防止重叠的轮询回调重复触发
let settled = false
pollRef.current = setInterval(async () => { pollRef.current = setInterval(async () => {
if (settled) return
try { try {
const status = await strmAPI.poll115OAuth(account.id, result.session_id) const status = await strmAPI.poll115OAuth(account.id, result.session_id)
if (settled) return
setAuthStatus(status.tip) setAuthStatus(status.tip)
if (status.status === 'confirmed') { if (status.status === 'confirmed') {
settled = true
stopPolling() stopPolling()
const updated = await strmAPI.testAccount(account.id) setAuthStatus('授权成功,正在验证凭据…')
onAuthed(updated) try {
const updated = await strmAPI.testAccount(account.id)
onAuthed(updated)
} catch (err) {
// 授权已确认、令牌已保存;凭据验证失败不阻塞授权完成,
// 避免轮询已停且 onAuthed 未回调导致弹窗卡死。
toast.error(`凭据验证失败:${apiErrorMessage(err)},授权已保存,可稍后在账号列表重试`)
onAuthed(account)
}
} }
if (status.status === 'expired') { if (status.status === 'expired') {
settled = true
stopPolling() stopPolling()
setAuthUI(null) setAuthUI(null)
toast.error('授权已过期,请重新发起') toast.error('授权已过期,请重新发起')
+8 -1
View File
@@ -41,9 +41,13 @@ function useCreateLibraryForm(refresh: () => Promise<void>) {
const [type, setType] = useState('movie') const [type, setType] = useState('movie')
const [coverURL, setCoverURL] = useState('') const [coverURL, setCoverURL] = useState('')
const [createPerSubfolder, setCreatePerSubfolder] = useState(false) const [createPerSubfolder, setCreatePerSubfolder] = useState(false)
// 提交中标记:防止重复点击创建出多个媒体库
const [creating, setCreating] = useState(false)
const handleCreate = async (e: FormEvent) => { const handleCreate = async (e: FormEvent) => {
e.preventDefault() e.preventDefault()
if (creating) return
setCreating(true)
try { try {
if (createPerSubfolder) { if (createPerSubfolder) {
const parentPath = roots[0]?.path?.trim() const parentPath = roots[0]?.path?.trim()
@@ -66,10 +70,12 @@ function useCreateLibraryForm(refresh: () => Promise<void>) {
setRoots([emptyRootDraft()]) setRoots([emptyRootDraft()])
setCoverURL('') setCoverURL('')
setCreatePerSubfolder(false) setCreatePerSubfolder(false)
await refresh()
} catch (err: unknown) { } catch (err: unknown) {
toast.error(apiErrorMessage(err, '创建失败')) toast.error(apiErrorMessage(err, '创建失败'))
} finally {
setCreating(false)
} }
await refresh().catch(() => undefined)
} }
const updateRoot = (index: number, patch: Partial<RootDraft>) => { const updateRoot = (index: number, patch: Partial<RootDraft>) => {
@@ -82,6 +88,7 @@ function useCreateLibraryForm(refresh: () => Promise<void>) {
coverURL, coverURL,
roots, roots,
createPerSubfolder, createPerSubfolder,
creating,
setName, setName,
setType, setType,
setCoverURL, setCoverURL,
+10 -2
View File
@@ -72,7 +72,7 @@ export function useLibraryData(libraryID: string, selectedSeries: SeriesCard | n
setServerSeriesCards(next.items) setServerSeriesCards(next.items)
setLoading(false) setLoading(false)
} }
}) }, () => cancelled)
if (!cancelled) setServerSeriesCards(collected.items) if (!cancelled) setServerSeriesCards(collected.items)
return return
} }
@@ -84,7 +84,7 @@ export function useLibraryData(libraryID: string, selectedSeries: SeriesCard | n
setItems(next.items) setItems(next.items)
setLoading(false) setLoading(false)
} }
}) }, () => cancelled)
if (!cancelled) setItems(collected.items) if (!cancelled) setItems(collected.items)
} catch { } catch {
if (!cancelled) toast.error('媒体库加载失败') if (!cancelled) toast.error('媒体库加载失败')
@@ -165,12 +165,16 @@ async function loadAllSeriesCards(
libraryID: string, libraryID: string,
isRemoteEmby: boolean | undefined, isRemoteEmby: boolean | undefined,
onPage: (state: { items: SeriesCard[]; total: number; firstPage: boolean }) => void, onPage: (state: { items: SeriesCard[]; total: number; firstPage: boolean }) => void,
isCancelled: () => boolean,
) { ) {
const pageSize = isRemoteEmby ? 100 : 500 const pageSize = isRemoteEmby ? 100 : 500
let page = 1 let page = 1
let collected: SeriesCard[] = [] let collected: SeriesCard[] = []
for (;;) { for (;;) {
// 切库/卸载后取消标志置位,立即停止继续拉取剩余页
if (isCancelled()) return { items: collected }
const data = await libraryAPI.listSeries(libraryID, page, pageSize) const data = await libraryAPI.listSeries(libraryID, page, pageSize)
if (isCancelled()) return { items: collected }
// 后端对空库可能返回 items: null(Go nil slice);不兜底会 concat 出 [null] 并崩溃。 // 后端对空库可能返回 items: null(Go nil slice);不兜底会 concat 出 [null] 并崩溃。
const pageItems = data.items ?? [] const pageItems = data.items ?? []
collected = collected.concat(pageItems) collected = collected.concat(pageItems)
@@ -186,12 +190,16 @@ async function loadAllMedia(
libraryID: string, libraryID: string,
isRemoteEmby: boolean | undefined, isRemoteEmby: boolean | undefined,
onPage: (state: { items: Media[]; total: number; firstPage: boolean }) => void, onPage: (state: { items: Media[]; total: number; firstPage: boolean }) => void,
isCancelled: () => boolean,
) { ) {
const pageSize = isRemoteEmby ? 100 : 2000 const pageSize = isRemoteEmby ? 100 : 2000
let page = 1 let page = 1
let collected: Media[] = [] let collected: Media[] = []
for (;;) { for (;;) {
// 切库/卸载后取消标志置位,立即停止继续拉取剩余页
if (isCancelled()) return { items: collected }
const data = await libraryAPI.listMedia(libraryID, page, pageSize) const data = await libraryAPI.listMedia(libraryID, page, pageSize)
if (isCancelled()) return { items: collected }
// 后端对空库可能返回 items: null(Go nil slice);不兜底会 concat 出 [null] 并崩溃。 // 后端对空库可能返回 items: null(Go nil slice);不兜底会 concat 出 [null] 并崩溃。
const pageItems = data.items ?? [] const pageItems = data.items ?? []
collected = collected.concat(pageItems) collected = collected.concat(pageItems)
+24 -10
View File
@@ -207,9 +207,13 @@ async function toggleMediaFavourite(
setFavourite: Dispatch<SetStateAction<boolean>>, setFavourite: Dispatch<SetStateAction<boolean>>,
): Promise<void> { ): Promise<void> {
if (!media) return if (!media) return
const state = await playbackAPI.toggleFavourite(media.id) try {
setFavourite(state) const state = await playbackAPI.toggleFavourite(media.id)
toast.success(state ? '已加入我的收藏' : '已取消收藏') setFavourite(state)
toast.success(state ? '已加入我的收藏' : '已取消收藏')
} catch (err: unknown) {
toast.error(apiErrorMessage(err, '收藏操作失败'))
}
} }
async function rescrapeMedia( async function rescrapeMedia(
@@ -218,13 +222,18 @@ async function rescrapeMedia(
refresh: () => Promise<void>, refresh: () => Promise<void>,
): Promise<void> { ): Promise<void> {
if (!media) return if (!media) return
await api.post(`/media/${media.id}/scrape`, { try {
episode_images: scrapeEpisodeArtwork, await api.post(`/media/${media.id}/scrape`, {
refresh_matched: true, episode_images: scrapeEpisodeArtwork,
include_matched: true, refresh_matched: true,
}) include_matched: true,
})
} catch (err: unknown) {
toast.error(apiErrorMessage(err, '触发重新刮削失败'))
return
}
toast.success('已触发重新刮削') toast.success('已触发重新刮削')
await refresh() await refresh().catch(() => undefined)
} }
async function reprobeMedia(media: Media | null, refresh: () => Promise<void>): Promise<void> { async function reprobeMedia(media: Media | null, refresh: () => Promise<void>): Promise<void> {
@@ -258,7 +267,12 @@ async function deleteMediaFromLibrary(media: Media | null, navigate: NavigateFun
checkboxLabel: '同时删除本地文件(含同名 NFO)', checkboxLabel: '同时删除本地文件(含同名 NFO)',
}) })
if (!result.confirmed) return if (!result.confirmed) return
await mediaAPI.delete(media.id, { deleteFiles: result.checked }) try {
await mediaAPI.delete(media.id, { deleteFiles: result.checked })
} catch (err: unknown) {
toast.error(apiErrorMessage(err, '删除媒体失败'))
return
}
toast.success(result.checked ? '已删除媒体及本地文件' : '已从媒体库删除') toast.success(result.checked ? '已删除媒体及本地文件' : '已从媒体库删除')
goBackFromMediaDetail(media, navigate, true) goBackFromMediaDetail(media, navigate, true)
} }