mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-28 11:16:37 +08:00
Compare commits
20 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| d00b14df55 | |||
| 1407b9b5c4 | |||
| 7fa05391e1 | |||
| 4790f7753e | |||
| a311438aa1 | |||
| 203abd106a | |||
| cc40169616 | |||
| cd720ae879 | |||
| fc84291346 | |||
| c37e936f48 | |||
| 389cb99bcf | |||
| 1b611a6181 | |||
| 25c03f2b0d | |||
| a56b1801f9 | |||
| 99c755dc29 | |||
| 086c0307c3 | |||
| 51d0f5010e | |||
| b9dd09a5d2 | |||
| e79f393969 | |||
| e422ecce53 |
+2
-1
@@ -23,6 +23,7 @@ import (
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/database"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
@@ -121,7 +122,7 @@ func main() {
|
||||
)
|
||||
}
|
||||
}()
|
||||
go services.Boot()
|
||||
helper.Go(logger, "services.boot", services.Boot)
|
||||
|
||||
// Graceful shutdown.
|
||||
stop := make(chan os.Signal, 1)
|
||||
|
||||
@@ -130,26 +130,35 @@ func (m *serverManager) Shutdown(ctx context.Context) error {
|
||||
// desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。
|
||||
// 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。
|
||||
func (m *serverManager) desiredPair() (*tlsPair, error) {
|
||||
if m.cfg == nil || !m.cfg.App.HTTPSEnabled {
|
||||
// 与 ApplyRuntimeSetting 的写锁配对:HTTPS 相关字段可能被运行时设置
|
||||
// 热更新,无锁读存在数据竞争(string 撕裂)。
|
||||
config.RuntimeMu.RLock()
|
||||
httpsEnabled := m.cfg != nil && m.cfg.App.HTTPSEnabled
|
||||
cert := m.cfg.App.SSLCert
|
||||
certPath := m.cfg.App.SSLCertPath
|
||||
key := m.cfg.App.SSLKey
|
||||
keyPath := m.cfg.App.SSLKeyPath
|
||||
config.RuntimeMu.RUnlock()
|
||||
if !httpsEnabled {
|
||||
return nil, nil
|
||||
}
|
||||
certPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLCert, m.cfg.App.SSLCertPath, "证书")
|
||||
certPEM, err := service.ResolveSSLMaterial(cert, certPath, "证书")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLKey, m.cfg.App.SSLKeyPath, "私钥")
|
||||
keyPEM, err := service.ResolveSSLMaterial(key, keyPath, "私钥")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
pairCert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err)
|
||||
}
|
||||
return &tlsPair{
|
||||
cert: cert,
|
||||
cert: pairCert,
|
||||
certPEM: certPEM,
|
||||
keyPEM: keyPEM,
|
||||
version: certPEM + "\x00" + keyPEM,
|
||||
@@ -171,6 +180,8 @@ func (m *serverManager) maybeStartAutoReloadLocked() {
|
||||
|
||||
// pathBased 是否至少有一侧证书/私钥通过文件路径配置。
|
||||
func (m *serverManager) pathBased() bool {
|
||||
config.RuntimeMu.RLock()
|
||||
defer config.RuntimeMu.RUnlock()
|
||||
return strings.TrimSpace(m.cfg.App.SSLCertPath) != "" || strings.TrimSpace(m.cfg.App.SSLKeyPath) != ""
|
||||
}
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ require (
|
||||
github.com/gin-contrib/gzip v1.2.6
|
||||
github.com/gin-gonic/gin v1.12.0
|
||||
github.com/glebarez/sqlite v1.11.0
|
||||
github.com/golang-jwt/jwt/v5 v5.2.0
|
||||
github.com/golang-jwt/jwt/v5 v5.2.2
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
github.com/redis/go-redis/v9 v9.7.0
|
||||
|
||||
@@ -52,8 +52,8 @@ github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
|
||||
github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
|
||||
github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM=
|
||||
github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.0 h1:d/ix8ftRUorsN+5eMIlF4T6J8CAt9rch3My2winC1Jw=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||
github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
@@ -19,6 +20,12 @@ import (
|
||||
// EnvPrefix 是所有环境变量驱动的覆盖使用的前缀。
|
||||
const EnvPrefix = "MeBox"
|
||||
|
||||
// RuntimeMu 保护运行时热更新配置字段的并发读写:ApplyRuntimeSetting 在
|
||||
// HTTP goroutine 中写字段,serverManager 的证书轮询等后台协程在无锁读取
|
||||
// 同一批字段。string 是双字结构,无锁并发读写可读到撕裂的 header。
|
||||
// 写方在 ApplyRuntimeSetting 内 Lock,读方(cmd/server)在轮询处 RLock。
|
||||
var RuntimeMu sync.RWMutex
|
||||
|
||||
// Load 从默认值 / 文件 / 环境读取配置。
|
||||
//
|
||||
// 即使没有文件也始终返回可用的 Config。
|
||||
@@ -45,8 +52,13 @@ func Load() (*Config, error) {
|
||||
}
|
||||
s := viper.New()
|
||||
s.SetConfigFile(filepath.Join("config", e.Name()))
|
||||
if err := s.ReadInConfig(); err == nil {
|
||||
_ = v.MergeConfigMap(s.AllSettings())
|
||||
if err := s.ReadInConfig(); err != nil {
|
||||
// 分片解析失败不能静默吞掉:database.yaml 语法错误会让
|
||||
// database.dsn 缺失 → type=auto 静默回退 SQLite,新数据
|
||||
// 全部写进一个空库而用户无感知。
|
||||
fmt.Fprintf(os.Stderr, "warning: parse config/%s failed: %v\n", e.Name(), err)
|
||||
} else {
|
||||
v.MergeConfigMap(s.AllSettings())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -68,8 +68,14 @@ func (c *Config) normalize() error {
|
||||
return fmt.Errorf("generate jwt secret: %w", err)
|
||||
}
|
||||
c.Secrets.JWTSecret = hex.EncodeToString(buf)
|
||||
_ = os.MkdirAll(c.App.DataDir, 0o750)
|
||||
_ = os.WriteFile(path, []byte(c.Secrets.JWTSecret), 0o600)
|
||||
// 持久化失败(DataDir 只读/权限异常)会导致每次重启重新生成
|
||||
// 密钥、全部会话静默失效、多实例各持不同 secret——必须让
|
||||
// 操作员感知。
|
||||
if mkErr := os.MkdirAll(c.App.DataDir, 0o750); mkErr != nil {
|
||||
fmt.Fprintf(os.Stderr, "warning: persist jwt secret failed (mkdir): %v\n", mkErr)
|
||||
} else if wErr := os.WriteFile(path, []byte(c.Secrets.JWTSecret), 0o600); wErr != nil {
|
||||
fmt.Fprintf(os.Stderr, "warning: persist jwt secret failed (write): %v\n", wErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
|
||||
@@ -34,8 +34,15 @@ func SaveDatabaseConfig(dbType, dsn string) error {
|
||||
return fmt.Errorf("marshal config.yaml: %w", err)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(configPath, out, 0644); err != nil {
|
||||
return fmt.Errorf("write config.yaml: %w", err)
|
||||
// 原子写:临时文件 + rename,避免进程崩溃/断电留下截断的 config.yaml
|
||||
// (下次启动会硬失败);DSN 含数据库密码,权限收窄到 0600。
|
||||
tmp := configPath + ".tmp"
|
||||
if err := os.WriteFile(tmp, out, 0o600); err != nil {
|
||||
return fmt.Errorf("write config.yaml.tmp: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmp, configPath); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return fmt.Errorf("replace config.yaml: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
@@ -73,6 +74,14 @@ func configureConnectionPool(db *gorm.DB, cfg *config.Config) error {
|
||||
if cfg.Database.MaxIdleConns > 0 {
|
||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||
}
|
||||
// 连接生命周期:默认 0 意味着 Postgres 重启/故障切换后的陈旧连接
|
||||
// 永不过期,首次复用才报错,运行期断连恢复慢且可能批量报错。
|
||||
if isPostgres(db) {
|
||||
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||
sqlDB.SetConnMaxIdleTime(10 * time.Minute)
|
||||
} else if isSQLite(db) {
|
||||
sqlDB.SetConnMaxLifetime(24 * time.Hour)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ const mediaSearchIndexSchemaVersion = 2
|
||||
|
||||
func ensureMediaSearchIndex(db *gorm.DB) error {
|
||||
if err := ensureMediaSearchMetaTable(db); err != nil {
|
||||
return nil
|
||||
return err // meta 表创建失败必须上抛,不能静默掩盖
|
||||
}
|
||||
version := currentMediaSearchIndexVersion(db)
|
||||
if version != mediaSearchIndexSchemaVersion {
|
||||
|
||||
@@ -8,6 +8,11 @@ import (
|
||||
|
||||
// AutoMigrate creates tables for every model registered in the model package.
|
||||
func AutoMigrate(db *gorm.DB) error {
|
||||
// 必须先于 AutoMigrate:旧库中可能已有重复的 (user_id, media_id) 历史行,
|
||||
// 不去重会导致唯一索引 uniq_user_history 创建失败。
|
||||
if err := dedupePlaybackHistories(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := db.AutoMigrate(model.AllModels()...); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -37,6 +42,27 @@ func ensureSQLiteQueryOptimizer(db *gorm.DB) error {
|
||||
return db.Exec("ANALYZE").Error
|
||||
}
|
||||
|
||||
// dedupePlaybackHistories removes duplicate (user_id, media_id) rows left by
|
||||
// the former read-then-write upsert, so the uniq_user_history composite unique
|
||||
// index can be created on existing databases. Keeps the most recent row per
|
||||
// pair, preferring live rows over soft-deleted ones.
|
||||
func dedupePlaybackHistories(db *gorm.DB) error {
|
||||
if !db.Migrator().HasTable("playback_histories") {
|
||||
return nil
|
||||
}
|
||||
return db.Exec(`
|
||||
DELETE FROM playback_histories WHERE id IN (
|
||||
SELECT id FROM (
|
||||
SELECT id, ROW_NUMBER() OVER (
|
||||
PARTITION BY user_id, media_id
|
||||
ORDER BY deleted_at IS NULL DESC, watched_at DESC, id DESC
|
||||
) AS rn
|
||||
FROM playback_histories
|
||||
) ranked
|
||||
WHERE ranked.rn > 1
|
||||
)`).Error
|
||||
}
|
||||
|
||||
func ensurePostgresColumnCompatibility(db *gorm.DB) error {
|
||||
if !isPostgres(db) {
|
||||
return nil
|
||||
@@ -106,13 +132,19 @@ func ensureEmbyMountsCompatibility(db *gorm.DB) error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// 针对已有数据:如果存在多个 sort_order=0/NULL 的记录,按创建时间顺序赋予稳定递增的序号
|
||||
// 针对已有数据:只给 sort_order=0/NULL 的行按创建时间补号(从现有
|
||||
// 最大值之后递增),不能整表重排——此前无条件按 created_at 从 0 重新
|
||||
// 编号,会把用户自定义的顺序覆盖掉。
|
||||
var zeroCount int64
|
||||
if err := db.Model(&model.EmbyMount{}).Where("sort_order = 0 OR sort_order IS NULL").Count(&zeroCount).Error; err == nil && zeroCount > 1 {
|
||||
if err := db.Model(&model.EmbyMount{}).Where("sort_order = 0 OR sort_order IS NULL").Count(&zeroCount).Error; err == nil && zeroCount > 0 {
|
||||
// max 只统计非 0 行:sort_order=0 与 NULL 同样视为“未分配”,
|
||||
// 全部为 0 时从 0 开始编号(与迁移前的初始化语义一致)。
|
||||
var maxOrder int
|
||||
_ = db.Raw("SELECT COALESCE(MAX(sort_order), -1) FROM emby_mounts WHERE sort_order > 0").Scan(&maxOrder).Error
|
||||
var mounts []model.EmbyMount
|
||||
if err := db.Order("created_at asc, id asc").Find(&mounts).Error; err == nil {
|
||||
if err := db.Where("sort_order = 0 OR sort_order IS NULL").Order("created_at asc, id asc").Find(&mounts).Error; err == nil {
|
||||
for i, m := range mounts {
|
||||
_ = db.Exec("UPDATE emby_mounts SET sort_order = ? WHERE id = ?", i, m.ID).Error
|
||||
_ = db.Exec("UPDATE emby_mounts SET sort_order = ? WHERE id = ?", maxOrder+1+i, m.ID).Error
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// TestAutoMigrateDedupesPlaybackHistories reproduces the upgrade path: a legacy
|
||||
// database contains duplicate (user_id, media_id) history rows created by the
|
||||
// old read-then-write upsert. AutoMigrate must merge them before creating the
|
||||
// uniq_user_history composite unique index, otherwise the upgrade fails.
|
||||
func TestAutoMigrateDedupesPlaybackHistories(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 旧 schema:无 uniq_user_history 唯一索引。
|
||||
if err := db.Exec(`CREATE TABLE playback_histories (
|
||||
id varchar(36) PRIMARY KEY,
|
||||
created_at datetime,
|
||||
updated_at datetime,
|
||||
deleted_at datetime,
|
||||
user_id varchar(36) NOT NULL,
|
||||
media_id varchar(128) NOT NULL,
|
||||
position_ms integer,
|
||||
duration_ms integer,
|
||||
watched_at datetime,
|
||||
completed numeric
|
||||
)`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
base := time.Now()
|
||||
rows := []struct {
|
||||
id string
|
||||
position int64
|
||||
watchedAt time.Time
|
||||
}{
|
||||
{"h-old", 1_000, base.Add(-2 * time.Hour)},
|
||||
{"h-mid", 2_000, base.Add(-1 * time.Hour)},
|
||||
{"h-new", 3_000, base},
|
||||
}
|
||||
for _, r := range rows {
|
||||
if err := db.Exec(
|
||||
`INSERT INTO playback_histories (id, user_id, media_id, position_ms, watched_at, created_at, updated_at)
|
||||
VALUES (?, 'u-1', 'm-1', ?, ?, ?, ?)`,
|
||||
r.id, r.position, r.watchedAt, r.watchedAt, r.watchedAt,
|
||||
).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := AutoMigrate(db); err != nil {
|
||||
t.Fatalf("auto migrate with duplicate histories: %v", err)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.Table("playback_histories").Where("user_id = ? AND media_id = ?", "u-1", "m-1").Count(&count).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Fatalf("expected duplicate rows merged to 1, got %d", count)
|
||||
}
|
||||
var position int64
|
||||
if err := db.Table("playback_histories").
|
||||
Where("user_id = ? AND media_id = ?", "u-1", "m-1").
|
||||
Select("position_ms").Scan(&position).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if position != 3_000 {
|
||||
t.Fatalf("dedupe should keep the most recent row, got position_ms=%d", position)
|
||||
}
|
||||
|
||||
// 唯一索引存在时,重复插入同一 (user_id, media_id) 应触发冲突而非新增行。
|
||||
if err := db.Exec(
|
||||
`INSERT INTO playback_histories (id, user_id, media_id, position_ms, watched_at, created_at, updated_at)
|
||||
VALUES ('h-dup', 'u-1', 'm-1', 4_000, ?, ?, ?)`,
|
||||
base, base, base,
|
||||
).Error; err == nil {
|
||||
t.Fatal("insert violating uniq_user_history should fail")
|
||||
}
|
||||
}
|
||||
@@ -50,28 +50,43 @@ func copyModelTables(src, target *gorm.DB, batchSize int) (map[string]int64, int
|
||||
if modelType.Kind() != reflect.Ptr {
|
||||
return tableCounts, totalCopied, fmt.Errorf("model %T is not a pointer", m)
|
||||
}
|
||||
sliceType := reflect.SliceOf(modelType.Elem())
|
||||
slicePtr := reflect.New(sliceType)
|
||||
if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
|
||||
return tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err)
|
||||
}
|
||||
filtered := slicePtr.Elem()
|
||||
var primaryKeySet map[string]struct{}
|
||||
if targetCount > 0 {
|
||||
primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
|
||||
primaryKeySet, err = targetPrimaryKeySet(target, table, primaryColumns)
|
||||
if err != nil {
|
||||
return tableCounts, totalCopied, err
|
||||
}
|
||||
filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
|
||||
}
|
||||
if filtered.Len() == 0 {
|
||||
continue
|
||||
// 分页流式读取:此前整表一次性 Find 进内存,media 表几十万行、
|
||||
// 每行含 overview/genres 等长文本时可达数百 MB,迁移过程有 OOM
|
||||
// 风险。源库在迁移期间是静态的,offset 分页安全。
|
||||
const readBatch = 1000
|
||||
copiedForTable := int64(0)
|
||||
for offset := 0; ; offset += readBatch {
|
||||
batchPtr := reflect.New(reflect.SliceOf(modelType.Elem()))
|
||||
if err := src.Unscoped().Limit(readBatch).Offset(offset).Find(batchPtr.Interface()).Error; err != nil {
|
||||
return tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err)
|
||||
}
|
||||
batch := batchPtr.Elem()
|
||||
if batch.Len() == 0 {
|
||||
break
|
||||
}
|
||||
filtered := batch
|
||||
if primaryKeySet != nil {
|
||||
filtered = filterRowsMissingInTarget(target, table, primaryColumns, batch, primaryKeySet)
|
||||
}
|
||||
if filtered.Len() > 0 {
|
||||
filteredPtr := reflect.New(filtered.Type())
|
||||
filteredPtr.Elem().Set(filtered)
|
||||
if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
|
||||
return tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err)
|
||||
}
|
||||
copiedForTable += int64(filtered.Len())
|
||||
}
|
||||
if batch.Len() < readBatch {
|
||||
break
|
||||
}
|
||||
}
|
||||
filteredPtr := reflect.New(filtered.Type())
|
||||
filteredPtr.Elem().Set(filtered)
|
||||
if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
|
||||
return tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err)
|
||||
}
|
||||
copiedForTable := int64(filtered.Len())
|
||||
tableCounts[table] = copiedForTable
|
||||
totalCopied += copiedForTable
|
||||
}
|
||||
|
||||
@@ -5,12 +5,21 @@ import (
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
)
|
||||
|
||||
// sqliteGateHoldLimit 是写闸持有者的最长合法持有时长。语句级写闸在 SQL
|
||||
// 执行 panic 时 After 回调不会运行,令牌会泄漏并让后续所有写入永久等锁;
|
||||
// 超过该时长的持有者按泄漏强制回收(60s 内单条写语句远未到,正常写路径
|
||||
// 不受影响)。
|
||||
const sqliteGateHoldLimit = 60 * time.Second
|
||||
|
||||
func installSQLiteWriteGate(db *gorm.DB) {
|
||||
if db == nil {
|
||||
return
|
||||
@@ -22,15 +31,18 @@ func installSQLiteWriteGate(db *gorm.DB) {
|
||||
if tx.Statement != nil && tx.Statement.Context != nil {
|
||||
ctx = tx.Statement.Context
|
||||
}
|
||||
if err := gate.Lock(ctx); err != nil {
|
||||
holder, err := gate.Lock(ctx)
|
||||
if err != nil {
|
||||
_ = tx.AddError(err)
|
||||
return
|
||||
}
|
||||
tx.InstanceSet(lockedKey, struct{}{})
|
||||
tx.InstanceSet(lockedKey, holder)
|
||||
}
|
||||
unlock := func(tx *gorm.DB) {
|
||||
if _, ok := tx.InstanceGet(lockedKey); ok {
|
||||
gate.Unlock()
|
||||
if holder, ok := tx.InstanceGet(lockedKey); ok {
|
||||
if h, ok := holder.(*sqliteGateHolder); ok {
|
||||
gate.Unlock(h)
|
||||
}
|
||||
}
|
||||
}
|
||||
rawLock := func(tx *gorm.DB) {
|
||||
@@ -64,38 +76,76 @@ func isReadOnlySQL(sql string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// sqliteWriteGate serializes in-process SQLite writes while respecting the
|
||||
// statement context, so request cancellation can break out of a queued write.
|
||||
// sqliteWriteGate serializes in-process SQLite writes. 所有权令牌(而非裸
|
||||
// 信号量)保证只有持有者本人能释放;持有超时按泄漏自动回收,避免一次
|
||||
// panic 让进程的 SQLite 写入半永久性瘫痪。
|
||||
type sqliteWriteGate struct {
|
||||
ch chan struct{}
|
||||
mu sync.Mutex
|
||||
cond *sync.Cond
|
||||
owner *sqliteGateHolder
|
||||
}
|
||||
|
||||
type sqliteGateHolder struct {
|
||||
id uint64
|
||||
acquired time.Time
|
||||
}
|
||||
|
||||
var sqliteGateHolderSeq atomic.Uint64
|
||||
|
||||
func newSQLiteWriteGate() *sqliteWriteGate {
|
||||
return &sqliteWriteGate{ch: make(chan struct{}, 1)}
|
||||
g := &sqliteWriteGate{}
|
||||
g.cond = sync.NewCond(&g.mu)
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *sqliteWriteGate) Lock(ctx context.Context) error {
|
||||
select {
|
||||
case g.ch <- struct{}{}:
|
||||
return nil
|
||||
default:
|
||||
}
|
||||
func (g *sqliteWriteGate) Lock(ctx context.Context) (*sqliteGateHolder, error) {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
select {
|
||||
case g.ch <- struct{}{}:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
// ctx 取消时唤醒等待者(cond 无法感知 ctx,用旁路 goroutine 广播)。
|
||||
if done := ctx.Done(); done != nil {
|
||||
stop := make(chan struct{})
|
||||
defer close(stop)
|
||||
go func() {
|
||||
select {
|
||||
case <-done:
|
||||
g.cond.Broadcast()
|
||||
case <-stop:
|
||||
}
|
||||
}()
|
||||
}
|
||||
for {
|
||||
if g.owner == nil {
|
||||
holder := &sqliteGateHolder{
|
||||
id: sqliteGateHolderSeq.Add(1),
|
||||
acquired: time.Now(),
|
||||
}
|
||||
g.owner = holder
|
||||
return holder, nil
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
if time.Since(g.owner.acquired) > sqliteGateHoldLimit {
|
||||
// 持有者疑似 panic 泄漏(After 回调未执行):强制回收。
|
||||
g.owner = nil
|
||||
g.cond.Broadcast()
|
||||
continue
|
||||
}
|
||||
g.cond.Wait()
|
||||
}
|
||||
}
|
||||
|
||||
func (g *sqliteWriteGate) Unlock() {
|
||||
select {
|
||||
case <-g.ch:
|
||||
default:
|
||||
func (g *sqliteWriteGate) Unlock(h *sqliteGateHolder) {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
if h == nil || g.owner != h {
|
||||
return
|
||||
}
|
||||
g.owner = nil
|
||||
g.cond.Broadcast()
|
||||
}
|
||||
|
||||
func buildSQLiteDSN(cfg *config.Config) string {
|
||||
@@ -104,7 +154,9 @@ func buildSQLiteDSN(cfg *config.Config) string {
|
||||
// keep as-is to respect user-provided relative paths.
|
||||
dbPath = filepath.Clean(dbPath)
|
||||
}
|
||||
dsn := dbPath + "?_pragma=foreign_keys(1)"
|
||||
// _txlock=immediate:事务以写锁开始。此前 deferred BEGIN 在并发事务
|
||||
// 升级写锁时会绕过 busy_timeout 直接报 SQLITE_BUSY。
|
||||
dsn := dbPath + "?_txlock=immediate&_pragma=foreign_keys(1)"
|
||||
if cfg.Database.WALMode {
|
||||
dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)"
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
@@ -86,10 +87,19 @@ func applyHTTPSetting(svc *service.Container, key, value string) error {
|
||||
svc.Log.Warn("https setting saved but not applied yet", zap.String("key", key), zap.String("reason", reason))
|
||||
}
|
||||
}
|
||||
|
||||
config.RuntimeMu.RLock()
|
||||
httpsEnabled := svc.Cfg.App.HTTPSEnabled
|
||||
cert := svc.Cfg.App.SSLCert
|
||||
certPath := svc.Cfg.App.SSLCertPath
|
||||
keyMaterial := svc.Cfg.App.SSLKey
|
||||
keyPath := svc.Cfg.App.SSLKeyPath
|
||||
config.RuntimeMu.RUnlock()
|
||||
|
||||
switch key {
|
||||
case "https.enabled":
|
||||
if svc.Cfg.App.HTTPSEnabled {
|
||||
if _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath); err != nil {
|
||||
if httpsEnabled {
|
||||
if _, err := service.ResolveSSLKeyPair(cert, certPath, keyMaterial, keyPath); err != nil {
|
||||
return fmt.Errorf("启用 HTTPS 失败:%v", err)
|
||||
}
|
||||
}
|
||||
@@ -97,7 +107,7 @@ func applyHTTPSetting(svc *service.Container, key, value string) error {
|
||||
if err := validateSSLMaterialSource(key, value); err != nil {
|
||||
return err
|
||||
}
|
||||
if !svc.Cfg.App.HTTPSEnabled {
|
||||
if !httpsEnabled {
|
||||
return nil
|
||||
}
|
||||
if !httpsPairReady(svc) {
|
||||
@@ -144,7 +154,13 @@ func validateSSLMaterialSource(key, value string) error {
|
||||
|
||||
// httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。
|
||||
func httpsPairReady(svc *service.Container) bool {
|
||||
_, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath)
|
||||
config.RuntimeMu.RLock()
|
||||
cert := svc.Cfg.App.SSLCert
|
||||
certPath := svc.Cfg.App.SSLCertPath
|
||||
keyMaterial := svc.Cfg.App.SSLKey
|
||||
keyPath := svc.Cfg.App.SSLKeyPath
|
||||
config.RuntimeMu.RUnlock()
|
||||
_, err := service.ResolveSSLKeyPair(cert, certPath, keyMaterial, keyPath)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@ package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -33,6 +35,21 @@ func dlnaCastHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// SSRF 防护:control_url 必须命中本服务发现到的真实渲染设备,
|
||||
// 防止登录用户借 cast 接口向任意内网地址发起 POST。
|
||||
// 优先用 30s 缓存;未命中时强制重扫一次再校验(设备可能刚上线)。
|
||||
devices, err := svc.DLNA.Discover(c.Request.Context(), false)
|
||||
if err == nil && !dlnaControlURLKnown(devices, req.ControlURL) {
|
||||
devices, err = svc.DLNA.Discover(c.Request.Context(), true)
|
||||
}
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if !dlnaControlURLKnown(devices, req.ControlURL) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unknown DLNA device: control_url must come from /api/dlna discovery"})
|
||||
return
|
||||
}
|
||||
if err := svc.DLNA.Cast(c.Request.Context(), req.ControlURL, req.MediaURL); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -40,3 +57,25 @@ func dlnaCastHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// dlnaControlURLKnown 判断 control_url 是否属于发现列表中的设备。
|
||||
// 按解析后的 host:port+path 精确比对,容忍大小写与尾斜杠差异。
|
||||
func dlnaControlURLKnown(devices []service.DLNADevice, controlURL string) bool {
|
||||
want, err := url.Parse(strings.TrimSpace(controlURL))
|
||||
if err != nil || want.Host == "" {
|
||||
return false
|
||||
}
|
||||
for _, dev := range devices {
|
||||
for _, candidate := range []string{dev.ControlURL, dev.Location} {
|
||||
u, err := url.Parse(strings.TrimSpace(candidate))
|
||||
if err != nil || u.Host == "" {
|
||||
continue
|
||||
}
|
||||
if strings.EqualFold(u.Host, want.Host) &&
|
||||
strings.EqualFold(strings.TrimRight(u.Path, "/"), strings.TrimRight(want.Path, "/")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -132,22 +132,25 @@ func embyMeHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
u, err := svc.Emby.FindUser(c.Request.Context(), c.Param("userId"))
|
||||
uid := embyUserID(c)
|
||||
if uid == "" {
|
||||
embyError(c, http.StatusUnauthorized, "not authenticated")
|
||||
return
|
||||
}
|
||||
// 只返回调用者自己的用户对象:客户端误传其他 userId 时回退到
|
||||
// 调用者自身(保留旧行为的兼容语义),但绝不返回他人数据。
|
||||
u, err := svc.Emby.FindUser(c.Request.Context(), uid)
|
||||
if err == nil && u != nil {
|
||||
c.JSON(http.StatusOK, u)
|
||||
return
|
||||
}
|
||||
if authUID := embyUserID(c); authUID != "" && authUID != c.Param("userId") {
|
||||
u, err = svc.Emby.FindUser(c.Request.Context(), authUID)
|
||||
if err == nil && u != nil {
|
||||
c.JSON(http.StatusOK, u)
|
||||
return
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, embyFallbackUser(c.Param("userId")))
|
||||
c.JSON(http.StatusOK, embyFallbackUser(uid))
|
||||
}
|
||||
}
|
||||
|
||||
// embyFallbackUser 是查库失败时的最后兜底(保持客户端可渲染)。
|
||||
// Policy 必须是最小权限:不声明管理员/删除内容/控制他人等能力,
|
||||
// 实际权限始终由服务端各路由的校验决定。
|
||||
func embyFallbackUser(id string) gin.H {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
id = "mebox-user"
|
||||
@@ -161,10 +164,10 @@ func embyFallbackUser(id string) gin.H {
|
||||
"HasConfiguredEasyPassword": false,
|
||||
"EnableAutoLogin": false,
|
||||
"Policy": gin.H{
|
||||
"IsAdministrator": true,
|
||||
"EnableContentDeletion": true,
|
||||
"EnableRemoteControlOfOtherUsers": true,
|
||||
"EnableSharedDeviceControl": true,
|
||||
"IsAdministrator": false,
|
||||
"EnableContentDeletion": false,
|
||||
"EnableRemoteControlOfOtherUsers": false,
|
||||
"EnableSharedDeviceControl": false,
|
||||
"EnableRemoteAccess": true,
|
||||
"EnableAllDevices": true,
|
||||
"EnableAllChannels": true,
|
||||
|
||||
+80
-33
@@ -11,8 +11,10 @@ import (
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
@@ -58,17 +60,37 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("include_hidden") == "true" || c.Query("all") == "1")
|
||||
if !includeHidden {
|
||||
libs = service.FilterDisplayCloudLibraries(ctx, svc.Repo, libs)
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
filtered := libs[:0]
|
||||
for _, lib := range libs {
|
||||
if service.LibraryVisibleForUser(ctx, svc.Repo, lib, visibility) {
|
||||
filtered = append(filtered, lib)
|
||||
if !includeHidden {
|
||||
libs = service.FilterDisplayCloudLibraries(ctx, svc.Repo, libs)
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
filtered := libs[:0]
|
||||
for _, lib := range libs {
|
||||
if service.LibraryVisibleForUser(ctx, svc.Repo, lib, visibility) {
|
||||
filtered = append(filtered, lib)
|
||||
}
|
||||
}
|
||||
libs = filtered
|
||||
}
|
||||
rawIDs := strings.TrimSpace(c.Query("ids"))
|
||||
var targetSet map[string]struct{}
|
||||
if rawIDs != "" {
|
||||
targetSet = make(map[string]struct{})
|
||||
for _, id := range strings.Split(rawIDs, ",") {
|
||||
id = strings.TrimSpace(id)
|
||||
if id != "" {
|
||||
targetSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
libs = filtered
|
||||
}
|
||||
if len(targetSet) > 0 {
|
||||
filtered := libs[:0]
|
||||
for _, lib := range libs {
|
||||
if _, ok := targetSet[lib.ID]; ok {
|
||||
filtered = append(filtered, lib)
|
||||
}
|
||||
}
|
||||
libs = filtered
|
||||
}
|
||||
withPreview := c.Query("with_preview") == "1" || c.Query("with_preview") == "true"
|
||||
limit := 10
|
||||
if withPreview {
|
||||
@@ -89,22 +111,41 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
for _, p := range previews {
|
||||
out = append(out, webLibraryPayload{Library: p.Library, Total: p.Total, Cards: p.Cards})
|
||||
}
|
||||
} else {
|
||||
for _, l := range libs {
|
||||
out = append(out, webLibraryPayload{Library: l})
|
||||
} else {
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
libIDs := make([]string, len(libs))
|
||||
for i, l := range libs {
|
||||
libIDs[i] = l.ID
|
||||
}
|
||||
counts, _ := svc.Repo.Media.CountByLibraries(ctx, libIDs, repository.MediaQueryFilter{
|
||||
IncludeNSFW: visibility.IncludeNSFW,
|
||||
AllowedLibraryIDs: visibility.AllowedLibraryIDs,
|
||||
HiddenLibraryIDs: visibility.HiddenLibraryIDs,
|
||||
})
|
||||
for _, l := range libs {
|
||||
var total int64
|
||||
if counts != nil {
|
||||
total = counts[l.ID]
|
||||
}
|
||||
out = append(out, webLibraryPayload{Library: l, Total: total})
|
||||
}
|
||||
}
|
||||
}
|
||||
// 远程 Emby 挂载库追加在本地库之后(非管理员视图仍受 allowed_library_ids 约束)。
|
||||
if svc.EmbyRemote != nil {
|
||||
if views, err := svc.EmbyRemote.RemoteLibraries(ctx); err == nil {
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
allowedViews := make([]service.RemoteLibraryView, 0, len(views))
|
||||
for _, v := range views {
|
||||
if !includeHidden && !service.LibraryVisibleForUser(ctx, svc.Repo, v.Library, visibility) {
|
||||
continue
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
allowedViews := make([]service.RemoteLibraryView, 0, len(views))
|
||||
for _, v := range views {
|
||||
if !includeHidden && !service.LibraryVisibleForUser(ctx, svc.Repo, v.Library, visibility) {
|
||||
continue
|
||||
}
|
||||
if len(targetSet) > 0 {
|
||||
if _, ok := targetSet[v.Library.ID]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
allowedViews = append(allowedViews, v)
|
||||
}
|
||||
allowedViews = append(allowedViews, v)
|
||||
}
|
||||
remotePayloads := make([]webLibraryPayload, len(allowedViews))
|
||||
for i, v := range allowedViews {
|
||||
remotePayloads[i] = webLibraryPayload{Library: v.Library, IsRemoteEmby: true, RemoteSource: v.AccountName}
|
||||
@@ -124,18 +165,20 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
acct := svc.EmbyRemote.AccountByID(ctx, v.AccountID)
|
||||
if acct == nil {
|
||||
return
|
||||
}
|
||||
tmpMount := &model.EmbyMount{Base: model.Base{ID: v.MountID}}
|
||||
itemTypes := remoteLibraryItemTypes(v.CollectionType)
|
||||
if _, total, err := svc.EmbyRemote.RemoteLibraryMedia(ctx, tmpMount, acct, v.RemoteID, itemTypes, 0, 1); err == nil {
|
||||
remotePayloads[i].Total = total
|
||||
}
|
||||
if cards, err := svc.EmbyRemote.RemoteLatestCards(ctx, tmpMount, acct, v.RemoteID, limit); err == nil {
|
||||
remotePayloads[i].Cards = cards
|
||||
}
|
||||
helper.Run(svc.Log, "media.remotePreview", func() {
|
||||
acct := svc.EmbyRemote.AccountByID(ctx, v.AccountID)
|
||||
if acct == nil {
|
||||
return
|
||||
}
|
||||
tmpMount := &model.EmbyMount{Base: model.Base{ID: v.MountID}}
|
||||
itemTypes := remoteLibraryItemTypes(v.CollectionType)
|
||||
if _, total, err := svc.EmbyRemote.RemoteLibraryMedia(ctx, tmpMount, acct, v.RemoteID, itemTypes, 0, 1); err == nil {
|
||||
remotePayloads[i].Total = total
|
||||
}
|
||||
if cards, err := svc.EmbyRemote.RemoteLatestCards(ctx, tmpMount, acct, v.RemoteID, limit); err == nil {
|
||||
remotePayloads[i].Cards = cards
|
||||
}
|
||||
})
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
@@ -325,7 +368,11 @@ func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
uid, _ := c.Get("ctx_user_id")
|
||||
svc.Audit.Record(c.Request.Context(), toString(uid), "library.delete", id, c.ClientIP(), "")
|
||||
go func() { _ = svc.Watcher.Refresh(context.Background()) }()
|
||||
// goroutine 内的 panic 无法被 gin.Recovery 捕获,会直接崩掉进程:
|
||||
// 与其他调用点一致先判空。
|
||||
if svc.Watcher != nil {
|
||||
go func() { _ = svc.Watcher.Refresh(context.Background()) }()
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
@@ -41,7 +42,11 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
task := startScanHTTPTask(svc, "手动扫描入库", lib.Name, lib.Path)
|
||||
go func(libraryID string, task *service.TaskHandle, finish func()) {
|
||||
defer finish()
|
||||
res, err := svc.Scan.ScanLibrary(context.Background(), libraryID)
|
||||
var res *service.ScanResult
|
||||
var err error
|
||||
helper.Run(svc.Log, "scan.library", func() {
|
||||
res, err = svc.Scan.ScanLibrary(context.Background(), libraryID)
|
||||
})
|
||||
if err != nil {
|
||||
finishHTTPTask(task, err, "scan", "手动扫描入库失败", scanTaskMetrics(res), scanTaskDetails(res, 20))
|
||||
return
|
||||
@@ -75,7 +80,11 @@ func scanLibraryRootHandler(svc *service.Container) gin.HandlerFunc {
|
||||
task := startScanHTTPTask(svc, "手动扫描媒体库路径", id, rootID)
|
||||
go func(libraryID, libraryRootID string, task *service.TaskHandle, finish func()) {
|
||||
defer finish()
|
||||
res, err := svc.Scan.ScanLibraryRoot(context.Background(), libraryID, libraryRootID)
|
||||
var res *service.ScanResult
|
||||
var err error
|
||||
helper.Run(svc.Log, "scan.libraryRoot", func() {
|
||||
res, err = svc.Scan.ScanLibraryRoot(context.Background(), libraryID, libraryRootID)
|
||||
})
|
||||
if err != nil {
|
||||
finishHTTPTask(task, err, "scan", "手动扫描路径失败", scanTaskMetrics(res), scanTaskDetails(res, 20))
|
||||
return
|
||||
@@ -110,11 +119,13 @@ func queueLibraryRootScan(svc *service.Container, libraryID, rootID string) {
|
||||
}
|
||||
go func() {
|
||||
defer finish()
|
||||
if strings.TrimSpace(rootID) == "" {
|
||||
_, _ = svc.Scan.ScanLibrary(context.Background(), libraryID)
|
||||
return
|
||||
}
|
||||
_, _ = svc.Scan.ScanLibraryRoot(context.Background(), libraryID, rootID)
|
||||
helper.Run(svc.Log, "scan.queuedRoot", func() {
|
||||
if strings.TrimSpace(rootID) == "" {
|
||||
_, _ = svc.Scan.ScanLibrary(context.Background(), libraryID)
|
||||
return
|
||||
}
|
||||
_, _ = svc.Scan.ScanLibraryRoot(context.Background(), libraryID, rootID)
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
|
||||
@@ -104,11 +104,16 @@ func TestListLibrariesHidesAdultDirectoriesUnlessAdminRequestsAll(t *testing.T)
|
||||
t.Fatalf("watching library list should hide adult directories, got %#v", visible)
|
||||
}
|
||||
|
||||
all := requestLibraries(t, svc, viewer.ID, "admin", "/api/libraries?include_hidden=1")
|
||||
if len(all) != 2 {
|
||||
t.Fatalf("admin include_hidden list should keep management access, got %#v", all)
|
||||
all := requestLibraries(t, svc, viewer.ID, "admin", "/api/libraries?include_hidden=1")
|
||||
if len(all) != 2 {
|
||||
t.Fatalf("admin include_hidden list should keep management access, got %#v", all)
|
||||
}
|
||||
|
||||
filtered := requestLibraries(t, svc, viewer.ID, "admin", "/api/libraries?include_hidden=1&ids="+safe.ID)
|
||||
if len(filtered) != 1 || filtered[0].ID != safe.ID {
|
||||
t.Fatalf("ids filter should return only requested library, got %#v", filtered)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLibraryAllowsEmptyLibrary(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -2,9 +2,11 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
@@ -157,6 +159,25 @@ type playlistItemReq struct {
|
||||
MediaID string `json:"media_id" binding:"required"`
|
||||
}
|
||||
|
||||
// playlistWriteGuard 校验当前用户对播放列表的写权限(属主或 admin)。
|
||||
// 校验失败时已写入错误响应,调用方直接 return。
|
||||
func playlistWriteGuard(c *gin.Context, svc *service.Container, playlistID string) (string, bool, bool) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
isAdmin := role == "admin"
|
||||
if err := svc.Playback.EnsurePlaylistOwner(c.Request.Context(), playlistID, uid.(string), isAdmin); err != nil {
|
||||
if errors.Is(err, service.ErrPlaylistForbidden) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"})
|
||||
} else if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "playlist not found"})
|
||||
} else {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
}
|
||||
return "", isAdmin, false
|
||||
}
|
||||
return uid.(string), isAdmin, true
|
||||
}
|
||||
|
||||
func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req playlistItemReq
|
||||
@@ -164,8 +185,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := svc.Playback.AddToPlaylist(
|
||||
c.Request.Context(), c.Param("id"), req.MediaID,
|
||||
c.Request.Context(), c.Param("id"), uid, req.MediaID, isAdmin,
|
||||
); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -176,8 +201,12 @@ func addPlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := svc.Playback.RemoveFromPlaylist(
|
||||
c.Request.Context(), c.Param("id"), c.Param("media_id"),
|
||||
c.Request.Context(), c.Param("id"), uid, c.Param("media_id"), isAdmin,
|
||||
); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -188,8 +217,12 @@ func removePlaylistItemHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func deletePlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, isAdmin, ok := playlistWriteGuard(c, svc, c.Param("id"))
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := svc.Playback.DeletePlaylist(
|
||||
c.Request.Context(), c.Param("id"),
|
||||
c.Request.Context(), c.Param("id"), uid, isAdmin,
|
||||
); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
|
||||
@@ -27,6 +27,9 @@ func reorderPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
pid := c.Param("id")
|
||||
if _, _, ok := playlistWriteGuard(c, svc, pid); !ok {
|
||||
return
|
||||
}
|
||||
for i, mid := range req.Order {
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Model(&model.PlaylistItem{}).
|
||||
@@ -44,8 +47,12 @@ func reorderPlaylistHandler(svc *service.Container) gin.HandlerFunc {
|
||||
// /playlists/:id/items/:item_id (vs. the existing /:media_id variant).
|
||||
func deletePlaylistItemByIDHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
pid := c.Param("id")
|
||||
if _, _, ok := playlistWriteGuard(c, svc, pid); !ok {
|
||||
return
|
||||
}
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Where("playlist_id = ? AND id = ?", c.Param("id"), c.Param("item_id")).
|
||||
Where("playlist_id = ? AND id = ?", pid, c.Param("item_id")).
|
||||
Delete(&model.PlaylistItem{}).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
|
||||
@@ -14,9 +14,14 @@ import (
|
||||
)
|
||||
|
||||
// statsUserHandler returns a watch-time summary for one user.
|
||||
// 观看统计是隐私数据:仅允许本人或管理员查询。
|
||||
func statsUserHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid := c.Param("id")
|
||||
if !statsCallerAllowed(c, uid) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"})
|
||||
return
|
||||
}
|
||||
var watched int64
|
||||
_ = svc.Repo.DB.Model(&model.PlaybackHistory{}).
|
||||
Where("user_id = ?", uid).
|
||||
@@ -35,8 +40,14 @@ func statsUserHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
|
||||
// statsTopUsersHandler returns the most active users by play count.
|
||||
// 全员排行含用户名与精确时长,仅管理员可查。
|
||||
func statsTopUsersHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
if role != "admin" {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"})
|
||||
return
|
||||
}
|
||||
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "10"))
|
||||
if limit <= 0 || limit > 50 {
|
||||
limit = 10
|
||||
@@ -109,3 +120,13 @@ func statsPlayHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
// statsCallerAllowed 判断当前调用者是否允许查看 uid 的观看统计。
|
||||
func statsCallerAllowed(c *gin.Context, uid string) bool {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
if role == "admin" {
|
||||
return true
|
||||
}
|
||||
caller, _ := c.Get(middleware.CtxUserID)
|
||||
return toString(caller) == uid
|
||||
}
|
||||
|
||||
@@ -39,11 +39,16 @@ func listSystemConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
|
||||
func isSecretKey(k string) bool {
|
||||
for _, suffix := range []string{".token", ".secret", ".password", ".api_key", ".cookie"} {
|
||||
for _, suffix := range []string{".token", ".secret", ".password", ".api_key", ".cookie", ".pin"} {
|
||||
if endsWith(k, suffix) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
// 非后缀型敏感键:可触发服务端任意命令的更新命令等。
|
||||
switch k {
|
||||
case "system.update.command":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
|
||||
+17
-4
@@ -9,6 +9,8 @@ package handler
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -21,10 +23,21 @@ import (
|
||||
var wsUpgrader = websocket.Upgrader{
|
||||
ReadBufferSize: 1024,
|
||||
WriteBufferSize: 1024,
|
||||
// Allow any origin: the AuthRequired middleware already validated the
|
||||
// JWT before we got here, and we never serve sensitive cross-domain
|
||||
// state through the socket.
|
||||
CheckOrigin: func(_ *http.Request) bool { return true },
|
||||
// 同源校验:浏览器跨站页面虽读不到 ?token=,但可能借 cookie 通道
|
||||
// (extractToken 接受 msgo_access_token cookie)发起跨站 WebSocket
|
||||
// 劫持。放行同源与非浏览器客户端(不发 Origin 头的 App/脚本),
|
||||
// 拒绝跨站 Origin。
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
origin := strings.TrimSpace(r.Header.Get("Origin"))
|
||||
if origin == "" {
|
||||
return true
|
||||
}
|
||||
u, err := url.Parse(origin)
|
||||
if err != nil || u.Host == "" {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(u.Host, r.Host)
|
||||
},
|
||||
}
|
||||
|
||||
func wsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
// Package helper provides shared utilities.
|
||||
package helper
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime/debug"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// Go runs fn in a new goroutine and recovers from panics so that a failure in
|
||||
// a background task (scraper parsing remote responses, cloud-drive sync, ...)
|
||||
// is logged instead of crashing the whole process. log may be nil.
|
||||
func Go(log *zap.Logger, name string, fn func()) {
|
||||
go Run(log, name, fn)
|
||||
}
|
||||
|
||||
// Run executes fn and recovers from panics, logging the task name and stack.
|
||||
// Use it as the first statement inside goroutines spawned elsewhere, or wrap
|
||||
// loop bodies so one bad iteration cannot kill a long-running worker.
|
||||
func Run(log *zap.Logger, name string, fn func()) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logPanic(log, name, r)
|
||||
}
|
||||
}()
|
||||
fn()
|
||||
}
|
||||
|
||||
// Recover runs fn and converts a panic into an error so callers can run their
|
||||
// own deferred cleanup (releasing locks, updating job state) before unwinding.
|
||||
func Recover(log *zap.Logger, name string, fn func() error) (err error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logPanic(log, name, r)
|
||||
err = fmt.Errorf("%s panicked: %v", name, r)
|
||||
}
|
||||
}()
|
||||
return fn()
|
||||
}
|
||||
|
||||
func logPanic(log *zap.Logger, name string, r any) {
|
||||
if log == nil {
|
||||
fmt.Fprintf(os.Stderr, "background task panicked: task=%s panic=%v\n%s\n", name, r, debug.Stack())
|
||||
return
|
||||
}
|
||||
log.Error("background task panicked",
|
||||
zap.String("task", name),
|
||||
zap.Any("panic", r),
|
||||
zap.ByteString("stack", debug.Stack()),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zaptest/observer"
|
||||
)
|
||||
|
||||
func newObservedLogger(t *testing.T) (*zap.Logger, *observer.ObservedLogs) {
|
||||
t.Helper()
|
||||
core, logs := observer.New(zap.ErrorLevel)
|
||||
return zap.New(core), logs
|
||||
}
|
||||
|
||||
func waitForLogs(t *testing.T, logs *observer.ObservedLogs, n int) []observer.LoggedEntry {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if entries := logs.All(); len(entries) >= n {
|
||||
return entries
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("timed out waiting for %d log entries, got %d", n, logs.Len())
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestRunRecoversPanic(t *testing.T) {
|
||||
log, logs := newObservedLogger(t)
|
||||
ran := false
|
||||
Run(log, "unit.panic", func() {
|
||||
ran = true
|
||||
panic("boom")
|
||||
})
|
||||
if !ran {
|
||||
t.Fatal("fn should have run before panicking")
|
||||
}
|
||||
entries := waitForLogs(t, logs, 1)
|
||||
if entries[0].Message != "background task panicked" {
|
||||
t.Fatalf("unexpected message: %s", entries[0].Message)
|
||||
}
|
||||
found := false
|
||||
for _, f := range entries[0].Context {
|
||||
if f.Key == "task" && f.String == "unit.panic" {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("expected task name in log context: %v", entries[0].Context)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunNoPanicNoLog(t *testing.T) {
|
||||
log, logs := newObservedLogger(t)
|
||||
Run(log, "unit.ok", func() {})
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
if logs.Len() != 0 {
|
||||
t.Fatalf("expected no error log, got %d", logs.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoverConvertsPanicToError(t *testing.T) {
|
||||
log, _ := newObservedLogger(t)
|
||||
err := Recover(log, "unit.recover", func() error {
|
||||
panic("kaboom")
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected error from recovered panic")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "kaboom") {
|
||||
t.Fatalf("panic value should be in error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoverReturnsFnError(t *testing.T) {
|
||||
sentinel := errors.New("plain failure")
|
||||
err := Recover(nil, "unit.err", func() error { return sentinel })
|
||||
if !errors.Is(err, sentinel) {
|
||||
t.Fatalf("expected fn error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunWithNilLoggerDoesNotCrash(t *testing.T) {
|
||||
Run(nil, "unit.nillog", func() { panic("still caught") })
|
||||
}
|
||||
|
||||
func TestGoLogsPanicFromSpawnedGoroutine(t *testing.T) {
|
||||
log, logs := newObservedLogger(t)
|
||||
Go(log, "unit.go", func() { panic("async boom") })
|
||||
waitForLogs(t, logs, 1)
|
||||
}
|
||||
@@ -51,6 +51,19 @@ func EmbyAuthRequired(secret string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 用途限定令牌(如 external_play,签发给外链播放器且绑定单一
|
||||
// media)只允许走 /api/stream|/hls|/cloud/play,绝不能作为全功能
|
||||
// 凭据访问 Emby 兼容面;否则外链 URL 一旦泄漏,持有者可获得
|
||||
// 该用户最长 24h 的全部 Emby API 权限。
|
||||
if strings.TrimSpace(claims.Purpose) != "" {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
"Code": 40101,
|
||||
"Message": "Invalid token",
|
||||
})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
c.Set(EmbyCtxUserID, claims.UserID)
|
||||
c.Set(CtxUserID, claims.UserID)
|
||||
c.Set(CtxUserRole, claims.Role)
|
||||
|
||||
@@ -16,6 +16,8 @@ type RateLimiter struct {
|
||||
window time.Duration
|
||||
max int
|
||||
requests map[string][]time.Time
|
||||
stop chan struct{}
|
||||
stopped sync.Once
|
||||
}
|
||||
|
||||
// NewRateLimiter creates a rate limiter allowing max requests per window
|
||||
@@ -25,14 +27,27 @@ func NewRateLimiter(max int, window time.Duration) *RateLimiter {
|
||||
window: window,
|
||||
max: max,
|
||||
requests: make(map[string][]time.Time),
|
||||
stop: make(chan struct{}),
|
||||
}
|
||||
go rl.cleanup()
|
||||
return rl
|
||||
}
|
||||
|
||||
// Close 停止后台清理 goroutine:清理循环此前无停止机制,每建一个实例
|
||||
// 就永久滞留一条 goroutine(测试场景会随实例创建不断累积)。
|
||||
func (rl *RateLimiter) Close() {
|
||||
rl.stopped.Do(func() { close(rl.stop) })
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) cleanup() {
|
||||
ticker := time.NewTicker(5 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
time.Sleep(5 * time.Minute)
|
||||
select {
|
||||
case <-rl.stop:
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
rl.mu.Lock()
|
||||
now := time.Now()
|
||||
for ip, times := range rl.requests {
|
||||
|
||||
@@ -5,15 +5,22 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// ApiConfig 存储第三方 API 密钥和配置信息。
|
||||
// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储。
|
||||
type ApiConfig struct {
|
||||
// APIConfig 存储第三方 API 密钥和配置信息。
|
||||
// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储(AES-GCM 密文,
|
||||
// base64 后常超 512 字符,因此必须是 text 而非 varchar(512))。
|
||||
//
|
||||
// NOTE: 历史上曾有 APIConfig / ApiConfig 两个结构体映射到同一张 api_configs
|
||||
// 表,AutoMigrate 每次启动互相改列(provider/api_key 长度来回切换),且
|
||||
// varchar(512) 收窄会让长密文入库后下一次启动迁移直接失败。现已合并为本
|
||||
// 结构体,字段取两者并集,请勿再拆分。
|
||||
type APIConfig struct {
|
||||
Base
|
||||
Provider string `gorm:"size:64;uniqueIndex;not null" json:"provider"`
|
||||
APIKey string `gorm:"size:512" json:"-"`
|
||||
BaseURL string `gorm:"size:512" json:"base_url,omitempty"`
|
||||
Extra string `gorm:"type:text" json:"extra,omitempty"`
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
Provider string `gorm:"uniqueIndex;size:64;not null" json:"provider"`
|
||||
APIKey string `gorm:"type:text" json:"-"` // ciphertext (never serialised)
|
||||
BaseURL string `gorm:"size:512" json:"base_url,omitempty"`
|
||||
Extra string `gorm:"type:text" json:"extra,omitempty"` // free-form JSON
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
|
||||
Description string `gorm:"size:255" json:"description,omitempty"`
|
||||
LastTestedAt *time.Time `json:"last_tested_at,omitempty"`
|
||||
TestResult string `gorm:"size:32" json:"test_result,omitempty"`
|
||||
|
||||
@@ -1,23 +1,8 @@
|
||||
package model
|
||||
|
||||
// APIConfig stores third-party data-source configuration. The api_key
|
||||
// column is encrypted with AES-GCM (see internal/service/crypto.go) so an
|
||||
// SQLite leak does not expose third-party credentials.
|
||||
//
|
||||
// Provider values mirror the original Python project:
|
||||
//
|
||||
// tmdb — themoviedb.org
|
||||
// bangumi — bgm.tv
|
||||
// thetvdb — thetvdb.com
|
||||
// fanart — fanart.tv
|
||||
// douban — douban.com (cookie)
|
||||
// openai — OpenAI / DeepSeek / Qwen / Ollama (compatible)
|
||||
type APIConfig struct {
|
||||
Base
|
||||
Provider string `gorm:"uniqueIndex;size:32;not null" json:"provider"`
|
||||
APIKey string `gorm:"type:text" json:"-"` // ciphertext (never serialised)
|
||||
BaseURL string `gorm:"size:512" json:"base_url,omitempty"`
|
||||
Extra string `gorm:"type:text" json:"extra,omitempty"` // free-form JSON
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
Description string `gorm:"size:255" json:"description,omitempty"`
|
||||
}
|
||||
// NOTE: 原 APIConfig(provider varchar(32) / api_key text)与 api_config.go
|
||||
// 里的 ApiConfig(provider varchar(64) / api_key varchar(512))映射到同一张
|
||||
// api_configs 表,AutoMigrate 每次启动互相改列;且 api_key 被收窄成
|
||||
// varchar(512) 后,成人区/豆瓣等存的长 AES-GCM Cookie 密文一旦入库,下次
|
||||
// 启动迁移即失败、服务无法启动。两者已合并为 api_config.go 中唯一的
|
||||
// APIConfig 结构体(字段取并集),此处不再定义重复模型。
|
||||
|
||||
@@ -26,7 +26,7 @@ type LibraryRoot struct {
|
||||
// Media 是单个可播放项。剧集链接到 SeriesID;电影 SeriesID == ""。
|
||||
type Media struct {
|
||||
Base
|
||||
LibraryID string `gorm:"index;size:36" json:"library_id"`
|
||||
LibraryID string `gorm:"index;size:36;index:idx_media_library_release,priority:1" json:"library_id"`
|
||||
LibraryRootID string `gorm:"index;size:36" json:"library_root_id,omitempty"`
|
||||
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
|
||||
Title string `gorm:"size:255;not null" json:"title"`
|
||||
@@ -46,7 +46,7 @@ type Media struct {
|
||||
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
||||
Rating float32 `json:"rating"`
|
||||
Year int `json:"year"`
|
||||
ReleaseDate string `gorm:"size:10;index" json:"release_date,omitempty"`
|
||||
ReleaseDate string `gorm:"size:10;index:idx_media_library_release,priority:2" json:"release_date,omitempty"`
|
||||
SeasonNum int `json:"season_num"`
|
||||
EpisodeNum int `json:"episode_num"`
|
||||
ScrapeStatus string `gorm:"size:16;default:pending" json:"scrape_status"`
|
||||
|
||||
@@ -46,7 +46,6 @@ func AllModels() []interface{} {
|
||||
&APIConfig{},
|
||||
&UserPermission{},
|
||||
&RefreshToken{},
|
||||
&ApiConfig{},
|
||||
&PlayProfile{},
|
||||
&RegistrationCode{},
|
||||
&SignIn{},
|
||||
|
||||
@@ -3,10 +3,12 @@ package model
|
||||
import "time"
|
||||
|
||||
// PlaybackHistory 记录当前播放位置以支持续播。
|
||||
// (user_id, media_id) 唯一:播放进度每几秒上报一次,唯一索引保证并发上报
|
||||
// 不会插入重复行(否则续播列表会出现重复卡片),也让 upsert 单语句完成。
|
||||
type PlaybackHistory struct {
|
||||
Base
|
||||
UserID string `gorm:"index;size:36;not null" json:"user_id"`
|
||||
MediaID string `gorm:"index;size:128;not null" json:"media_id"`
|
||||
UserID string `gorm:"index;size:36;not null;uniqueIndex:uniq_user_history" json:"user_id"`
|
||||
MediaID string `gorm:"index;size:128;not null;uniqueIndex:uniq_user_history" json:"media_id"`
|
||||
PositionMs int64 `json:"position_ms"`
|
||||
DurationMs int64 `json:"duration_ms"`
|
||||
WatchedAt time.Time `json:"watched_at"`
|
||||
|
||||
@@ -10,17 +10,17 @@ import (
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
// ApiConfigRepository persists model.ApiConfig records.
|
||||
// ApiConfigRepository persists model.APIConfig records.
|
||||
type ApiConfigRepository struct{ db *gorm.DB }
|
||||
|
||||
// Create inserts a new API config record.
|
||||
func (r *ApiConfigRepository) Create(ctx context.Context, c *model.ApiConfig) error {
|
||||
func (r *ApiConfigRepository) Create(ctx context.Context, c *model.APIConfig) error {
|
||||
return r.db.WithContext(ctx).Create(c).Error
|
||||
}
|
||||
|
||||
// FindByProvider returns the API config for a provider, or (nil, nil).
|
||||
func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
|
||||
var c model.ApiConfig
|
||||
func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.APIConfig, error) {
|
||||
var c model.APIConfig
|
||||
err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
@@ -32,27 +32,40 @@ func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider strin
|
||||
}
|
||||
|
||||
// List returns all API configs.
|
||||
func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) {
|
||||
var rows []model.ApiConfig
|
||||
func (r *ApiConfigRepository) List(ctx context.Context) ([]model.APIConfig, error) {
|
||||
var rows []model.APIConfig
|
||||
err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// Upsert creates or updates an API config.
|
||||
func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.ApiConfig) error {
|
||||
return r.db.WithContext(ctx).Where("provider = ?", c.Provider).
|
||||
Assign(model.ApiConfig{
|
||||
Base: model.Base{UpdatedAt: time.Now()},
|
||||
APIKey: c.APIKey,
|
||||
BaseURL: c.BaseURL,
|
||||
Extra: c.Extra,
|
||||
Enabled: c.Enabled,
|
||||
}).FirstOrCreate(c).Error
|
||||
// 显式 map 更新:Assign(struct) 会跳过零值字段,导致 Enabled=false、
|
||||
// 清空 BaseURL/Extra 等撤销操作静默失效。
|
||||
func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.APIConfig) error {
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var existing model.APIConfig
|
||||
err := tx.Where("provider = ?", c.Provider).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return tx.Create(c).Error
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.ID = existing.ID
|
||||
c.CreatedAt = existing.CreatedAt
|
||||
return tx.Model(&existing).Updates(map[string]any{
|
||||
"api_key": c.APIKey,
|
||||
"base_url": c.BaseURL,
|
||||
"extra": c.Extra,
|
||||
"enabled": c.Enabled,
|
||||
"updated_at": time.Now(),
|
||||
}).Error
|
||||
})
|
||||
}
|
||||
|
||||
// Update updates an API config.
|
||||
func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) error {
|
||||
return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
|
||||
func (r *ApiConfigRepository) Update(ctx context.Context, c *model.APIConfig) error {
|
||||
return r.db.WithContext(ctx).Model(&model.APIConfig{}).
|
||||
Where("provider = ?", c.Provider).Updates(map[string]any{
|
||||
"api_key": c.APIKey,
|
||||
"base_url": c.BaseURL,
|
||||
@@ -64,13 +77,13 @@ func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) er
|
||||
|
||||
// Delete 物理删除 API 配置。
|
||||
func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error {
|
||||
return r.db.WithContext(ctx).Unscoped().Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
|
||||
return r.db.WithContext(ctx).Unscoped().Where("provider = ?", provider).Delete(&model.APIConfig{}).Error
|
||||
}
|
||||
|
||||
// UpdateTestResult 更新测试结果。
|
||||
func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error {
|
||||
now := time.Now()
|
||||
return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
|
||||
return r.db.WithContext(ctx).Model(&model.APIConfig{}).
|
||||
Where("provider = ?", provider).Updates(map[string]any{
|
||||
"test_result": result,
|
||||
"last_tested_at": &now,
|
||||
|
||||
@@ -125,3 +125,16 @@ func (r *EmbyMountRepository) DeleteByAccountID(ctx context.Context, accountID s
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
|
||||
// DeleteOrphans 清理账号已不存在的挂载(老版本删除账号未级联的历史残留)。
|
||||
func (r *EmbyMountRepository) DeleteOrphans(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).
|
||||
Where("account_id NOT IN (SELECT id FROM strm_accounts)").
|
||||
Delete(&model.EmbyMount{})
|
||||
count = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
|
||||
@@ -2,9 +2,9 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
@@ -13,25 +13,24 @@ import (
|
||||
// upserts on (UserID, MediaID) so resume always reads the latest position.
|
||||
type HistoryRepository struct{ db *gorm.DB }
|
||||
|
||||
// Upsert atomically inserts/updates the resume position.
|
||||
// Upsert atomically inserts/updates the resume position in a single statement,
|
||||
// relying on the uniq_user_history composite unique index. Concurrent progress
|
||||
// reports for the same (user, media) can no longer double-insert.
|
||||
func (r *HistoryRepository) Upsert(ctx context.Context, h *model.PlaybackHistory) error {
|
||||
var existing model.PlaybackHistory
|
||||
err := r.db.WithContext(ctx).
|
||||
Where("user_id = ? AND media_id = ?", h.UserID, h.MediaID).
|
||||
First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return r.db.WithContext(ctx).Create(h).Error
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
existing.PositionMs = h.PositionMs
|
||||
if h.DurationMs > 0 {
|
||||
existing.DurationMs = h.DurationMs
|
||||
}
|
||||
existing.WatchedAt = h.WatchedAt
|
||||
existing.Completed = h.Completed
|
||||
return r.db.WithContext(ctx).Save(&existing).Error
|
||||
return r.db.WithContext(ctx).Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "user_id"}, {Name: "media_id"}},
|
||||
DoUpdates: clause.Assignments(map[string]any{
|
||||
"position_ms": h.PositionMs,
|
||||
// 沿用旧语义:未知时长(0)不覆盖已记录的时长。
|
||||
"duration_ms": gorm.Expr(
|
||||
"CASE WHEN ? > 0 THEN ? ELSE playback_histories.duration_ms END",
|
||||
h.DurationMs, h.DurationMs,
|
||||
),
|
||||
"watched_at": h.WatchedAt,
|
||||
"completed": h.Completed,
|
||||
"deleted_at": nil,
|
||||
}),
|
||||
}).Create(h).Error
|
||||
}
|
||||
|
||||
// ListByUser returns the most recent history rows for the user.
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/database"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
func TestHistoryUpsertSingleRowPerUserMedia(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := New(db)
|
||||
ctx := t.Context()
|
||||
watched := time.Now()
|
||||
|
||||
first := &model.PlaybackHistory{UserID: "u-1", MediaID: "m-1", PositionMs: 30_000, DurationMs: 0, WatchedAt: watched, Completed: false}
|
||||
if err := repos.History.Upsert(ctx, first); err != nil {
|
||||
t.Fatalf("first upsert: %v", err)
|
||||
}
|
||||
second := &model.PlaybackHistory{UserID: "u-1", MediaID: "m-1", PositionMs: 90_000, DurationMs: 120_000, WatchedAt: watched.Add(time.Minute), Completed: true}
|
||||
if err := repos.History.Upsert(ctx, second); err != nil {
|
||||
t.Fatalf("second upsert: %v", err)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.Model(&model.PlaybackHistory{}).Where("user_id = ? AND media_id = ?", "u-1", "m-1").Count(&count).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Fatalf("expected 1 history row after upserts, got %d", count)
|
||||
}
|
||||
var got model.PlaybackHistory
|
||||
if err := db.Where("user_id = ? AND media_id = ?", "u-1", "m-1").First(&got).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.PositionMs != 90_000 || !got.Completed {
|
||||
t.Fatalf("position/completion not updated: %#v", got)
|
||||
}
|
||||
if got.DurationMs != 120_000 {
|
||||
t.Fatalf("duration should update when known, got %d", got.DurationMs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHistoryUpsertKeepsDurationWhenUnknown(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := New(db)
|
||||
ctx := t.Context()
|
||||
watched := time.Now()
|
||||
|
||||
if err := repos.History.Upsert(ctx, &model.PlaybackHistory{UserID: "u-1", MediaID: "m-2", PositionMs: 10, DurationMs: 600_000, WatchedAt: watched}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := repos.History.Upsert(ctx, &model.PlaybackHistory{UserID: "u-1", MediaID: "m-2", PositionMs: 20, DurationMs: 0, WatchedAt: watched.Add(time.Second)}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var got model.PlaybackHistory
|
||||
if err := db.Where("user_id = ? AND media_id = ?", "u-1", "m-2").First(&got).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.DurationMs != 600_000 {
|
||||
t.Fatalf("duration_ms=0 upsert must not clear stored duration, got %d", got.DurationMs)
|
||||
}
|
||||
}
|
||||
@@ -22,45 +22,93 @@ import (
|
||||
// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
|
||||
// 永远捞不到数据。
|
||||
func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.upsertWithDB(ctx, r.db, m)
|
||||
var indexIDs []string
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
id, uerr := r.upsertWithDB(ctx, r.db, m)
|
||||
if uerr != nil {
|
||||
return uerr
|
||||
}
|
||||
indexIDs = append(indexIDs[:0], id)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.indexByIDBestEffort(ctx, indexIDs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpsertBatch 在单个事务里逐条执行 Upsert:扫描一批只提交(fsync)一次,
|
||||
// 而不是每条一个隐式事务。任一条目落库失败不影响批内已成功的条目——
|
||||
// 事务回滚后由调用方退回逐条 Upsert 兜底。
|
||||
//
|
||||
// OpenSearch 索引同步(HTTP,4s 超时)必须在事务提交之后统一执行:放在
|
||||
// 事务内会把 SQLite 写锁挂起在网络 IO 上,且批内用非事务连接回读只能
|
||||
// 拿到提交前的旧版本数据,把陈旧内容写进索引。
|
||||
func (r *MediaRepository) UpsertBatch(ctx context.Context, items []*model.Media) error {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
indexIDs := make([]string, 0, len(items))
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
indexIDs = indexIDs[:0]
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
for _, m := range items {
|
||||
if m == nil {
|
||||
continue
|
||||
}
|
||||
if err := r.upsertWithDB(ctx, tx, m); err != nil {
|
||||
id, err := r.upsertWithDB(ctx, tx, m)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if id != "" {
|
||||
indexIDs = append(indexIDs, id)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func (r *MediaRepository) upsertWithDB(ctx context.Context, db *gorm.DB, m *model.Media) error {
|
||||
existing, created, err := r.findOrCreateMediaByPath(ctx, db, m)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.indexByIDBestEffort(ctx, indexIDs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// indexByIDBestEffort 在事务提交后按 ID 回读最新行并同步搜索索引。
|
||||
func (r *MediaRepository) indexByIDBestEffort(ctx context.Context, ids []string) {
|
||||
for _, id := range ids {
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
if fresh, err := r.FindByID(ctx, id); err == nil && fresh != nil {
|
||||
r.indexMediaBestEffort(ctx, *fresh)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// upsertWithDB 落库(新建或更新),返回需要重建索引的媒体 ID(无则空串)。
|
||||
func (r *MediaRepository) upsertWithDB(ctx context.Context, db *gorm.DB, m *model.Media) (string, error) {
|
||||
existing, created, err := r.findOrCreateMediaByPath(ctx, db, m)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if created {
|
||||
r.indexMediaBestEffort(ctx, *m)
|
||||
return nil
|
||||
return m.ID, nil
|
||||
}
|
||||
|
||||
updates := mediaUpsertUpdates(existing, *m)
|
||||
return r.applyMediaUpsertUpdates(ctx, db, m, existing, updates)
|
||||
if len(updates) == 0 {
|
||||
*m = existing
|
||||
return "", nil
|
||||
}
|
||||
if err := db.WithContext(ctx).Unscoped().Model(&model.Media{}).
|
||||
Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 回写 ID / 不可变字段,让 caller 拿到完整的现有行。
|
||||
*m = existing
|
||||
return existing.ID, nil
|
||||
}
|
||||
|
||||
func (r *MediaRepository) findOrCreateMediaByPath(ctx context.Context, db *gorm.DB, m *model.Media) (model.Media, bool, error) {
|
||||
@@ -75,6 +123,9 @@ func (r *MediaRepository) findOrCreateMediaByPath(ctx context.Context, db *gorm.
|
||||
return *m, true, nil
|
||||
} else if retryErr := db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil {
|
||||
return model.Media{}, false, createErr
|
||||
} else {
|
||||
// 并发插入竞态:重查已命中既有行,直接走更新分支。
|
||||
return existing, false, nil
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
@@ -259,24 +310,6 @@ func setNonEmptyMediaString(updates map[string]any, key, current, next string) {
|
||||
}
|
||||
}
|
||||
|
||||
func (r *MediaRepository) applyMediaUpsertUpdates(ctx context.Context, db *gorm.DB, m *model.Media, existing model.Media, updates map[string]any) error {
|
||||
if len(updates) == 0 {
|
||||
*m = existing
|
||||
return nil
|
||||
}
|
||||
if err := db.WithContext(ctx).Unscoped().Model(&model.Media{}).
|
||||
Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// 回写 ID / 不可变字段,让 caller 拿到完整的现有行。
|
||||
*m = existing
|
||||
if fresh, err := r.FindByID(ctx, existing.ID); err == nil && fresh != nil {
|
||||
*m = *fresh
|
||||
r.indexMediaBestEffort(ctx, *fresh)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
|
||||
if current != next {
|
||||
updates[key] = next
|
||||
|
||||
@@ -105,11 +105,17 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string,
|
||||
var total int64
|
||||
q := r.db.WithContext(ctx).Model(&model.Media{})
|
||||
q = applyMediaQueryFilter(q, filter)
|
||||
// SQLite 的 LIKE 对 ASCII 不区分大小写;Postgres 的 LIKE 区分大小写,
|
||||
// 需用 ILIKE 保持两端搜索行为一致。
|
||||
likeOp := "LIKE"
|
||||
if r.db.Dialector != nil && r.db.Dialector.Name() == "postgres" {
|
||||
likeOp = "ILIKE"
|
||||
}
|
||||
terms := mediaSearchTerms(query)
|
||||
for _, term := range terms {
|
||||
like := "%" + escapeLike(term) + "%"
|
||||
q = q.Where(
|
||||
"(title LIKE ? ESCAPE '\\' OR original_name LIKE ? ESCAPE '\\' OR path LIKE ? ESCAPE '\\' OR genres LIKE ? ESCAPE '\\')",
|
||||
"(title "+likeOp+" ? ESCAPE '\\' OR original_name "+likeOp+" ? ESCAPE '\\' OR path "+likeOp+" ? ESCAPE '\\' OR genres "+likeOp+" ? ESCAPE '\\')",
|
||||
like, like, like, like,
|
||||
)
|
||||
}
|
||||
@@ -120,7 +126,7 @@ func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string,
|
||||
prefix := escapeLike(query) + "%"
|
||||
exact := query
|
||||
q = q.Order(gorm.Expr(
|
||||
"CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title LIKE ? ESCAPE '\\' THEN 2 WHEN original_name LIKE ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc",
|
||||
"CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title "+likeOp+" ? ESCAPE '\\' THEN 2 WHEN original_name "+likeOp+" ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc",
|
||||
exact, exact, prefix, prefix,
|
||||
))
|
||||
} else {
|
||||
@@ -259,7 +265,9 @@ func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool {
|
||||
}
|
||||
r.searchIndexOnce.Do(func() {
|
||||
var count int64
|
||||
err := r.db.WithContext(ctx).
|
||||
// 用 Background 探测:sync.Once 只执行一次,若借用调用方的
|
||||
// ctx 且恰好被取消,FTS 会被永久误判为不可用。
|
||||
err := r.db.WithContext(context.Background()).
|
||||
Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`).
|
||||
Scan(&count).Error
|
||||
r.searchIndexAvailable = err == nil && count > 0
|
||||
|
||||
@@ -3,6 +3,7 @@ package repository
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
@@ -44,10 +45,43 @@ func (r *PermissionRepository) Update(ctx context.Context, userID string, update
|
||||
}
|
||||
|
||||
// Upsert creates or updates a permission record.
|
||||
// 显式 map 更新:Assign(struct) 会被 GORM 跳过零值字段,导致权限
|
||||
// "撤销"(false)保存后静默失效且无法重置。
|
||||
func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
|
||||
Assign(*p).FirstOrCreate(p).Error
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var existing model.UserPermission
|
||||
err := tx.Where("user_id = ?", p.UserID).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return tx.Create(p).Error
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.ID = existing.ID
|
||||
p.CreatedAt = existing.CreatedAt
|
||||
return tx.Model(&existing).Updates(map[string]any{
|
||||
"can_view_dashboard": p.CanViewDashboard,
|
||||
"can_play_media": p.CanPlayMedia,
|
||||
"can_cast": p.CanCast,
|
||||
"can_external_player": p.CanExternalPlayer,
|
||||
"can_favorite": p.CanFavorite,
|
||||
"can_view_history": p.CanViewHistory,
|
||||
"can_edit_media": p.CanEditMedia,
|
||||
"can_rescrape": p.CanRescrape,
|
||||
"can_use_ai": p.CanUseAI,
|
||||
"can_capture_frames": p.CanCaptureFrames,
|
||||
"can_manage_downloads": p.CanManageDownloads,
|
||||
"can_manage_subscriptions": p.CanManageSubscriptions,
|
||||
"can_manage_sites": p.CanManageSites,
|
||||
"can_use_ai_assistant": p.CanUseAIAssistant,
|
||||
"can_manage_users": p.CanManageUsers,
|
||||
"can_manage_files": p.CanManageFiles,
|
||||
"can_manage_strm": p.CanManageStrm,
|
||||
"can_access_settings": p.CanAccessSettings,
|
||||
"updated_at": time.Now(),
|
||||
}).Error
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -52,6 +52,40 @@ func (r *ScrapeTaskRepository) FindActiveByMediaID(ctx context.Context, mediaID
|
||||
return &t, err
|
||||
}
|
||||
|
||||
// FindActiveByMediaIDs 批量查询仍处于 pending/running 的任务媒体 ID 集合,
|
||||
// 供整库入队时去重(防止同一媒体被重复入队并被并发双刮)。
|
||||
func (r *ScrapeTaskRepository) FindActiveByMediaIDs(ctx context.Context, mediaIDs []string) (map[string]bool, error) {
|
||||
out := make(map[string]bool, len(mediaIDs))
|
||||
if len(mediaIDs) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var rows []model.ScrapeTask
|
||||
err := r.db.WithContext(ctx).
|
||||
Select("media_id").
|
||||
Where("media_id IN ? AND status IN ?", mediaIDs, []string{model.ScrapeTaskPending, model.ScrapeTaskRunning}).
|
||||
Find(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, r := range rows {
|
||||
out[r.MediaID] = true
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务重置为 pending,
|
||||
// 否则任务永久卡死(ClaimPending 只认 pending,重试按钮也拒绝 running)。
|
||||
func (r *ScrapeTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) {
|
||||
res := r.db.WithContext(ctx).Model(&model.ScrapeTask{}).
|
||||
Where("status = ?", model.ScrapeTaskRunning).
|
||||
Updates(map[string]any{
|
||||
"status": model.ScrapeTaskPending,
|
||||
"error": "服务重启,任务已重置",
|
||||
"started_at": nil,
|
||||
})
|
||||
return res.RowsAffected, res.Error
|
||||
}
|
||||
|
||||
func (r *ScrapeTaskRepository) List(ctx context.Context, status string, page, pageSize int) ([]model.ScrapeTask, int64, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
|
||||
@@ -302,6 +302,39 @@ func (r *StrmDownloadTaskRepository) Update(ctx context.Context, t *model.StrmDo
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateIfRunning 仅当任务在 DB 中仍为 running 时写入给定字段。
|
||||
// 返回 false 表示任务已被外部改变状态(如用户取消),收尾不得覆盖。
|
||||
func (r *StrmDownloadTaskRepository) UpdateIfRunning(ctx context.Context, id string, updates map[string]any) (bool, error) {
|
||||
var ok bool
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
updates["updated_at"] = time.Now()
|
||||
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
|
||||
Where("id = ? AND status = ?", id, model.StrmTaskRunning).Updates(updates)
|
||||
ok = res.RowsAffected > 0
|
||||
return res.Error
|
||||
})
|
||||
return ok, err
|
||||
}
|
||||
|
||||
// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务全部重置为
|
||||
// pending(清空退避时间以便立即可被认领),否则任务永久卡死且会阻塞
|
||||
// 该文件的重复下载。
|
||||
func (r *StrmDownloadTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) {
|
||||
var n int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
|
||||
Where("status = ?", model.StrmTaskRunning).
|
||||
Updates(map[string]any{
|
||||
"status": model.StrmTaskPending,
|
||||
"error": "服务重启,任务已重置",
|
||||
"started_at": nil,
|
||||
})
|
||||
n = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error
|
||||
@@ -620,6 +653,36 @@ func (r *StrmUploadTaskRepository) Update(ctx context.Context, t *model.StrmUplo
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateIfRunning 仅当任务在 DB 中仍为 running 时写入给定字段。
|
||||
func (r *StrmUploadTaskRepository) UpdateIfRunning(ctx context.Context, id string, updates map[string]any) (bool, error) {
|
||||
var ok bool
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
updates["updated_at"] = time.Now()
|
||||
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
|
||||
Where("id = ? AND status = ?", id, model.StrmTaskRunning).Updates(updates)
|
||||
ok = res.RowsAffected > 0
|
||||
return res.Error
|
||||
})
|
||||
return ok, err
|
||||
}
|
||||
|
||||
// ResetRunningToPending 启动自愈:进程中断遗留的 running 任务全部重置为 pending。
|
||||
func (r *StrmUploadTaskRepository) ResetRunningToPending(ctx context.Context) (int64, error) {
|
||||
var n int64
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
res := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).
|
||||
Where("status = ?", model.StrmTaskRunning).
|
||||
Updates(map[string]any{
|
||||
"status": model.StrmTaskPending,
|
||||
"error": "服务重启,任务已重置",
|
||||
"started_at": nil,
|
||||
})
|
||||
n = res.RowsAffected
|
||||
return res.Error
|
||||
})
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) error {
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Unscoped().Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error
|
||||
|
||||
@@ -20,7 +20,7 @@ func TestAdultProviderRouting(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("failed to open sqlite: %v", err)
|
||||
}
|
||||
_ = db.AutoMigrate(&model.Setting{}, &model.ApiConfig{})
|
||||
_ = db.AutoMigrate(&model.Setting{}, &model.APIConfig{})
|
||||
|
||||
repos := repository.New(db)
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ func (s *ApiConfigService) TestConnection(ctx context.Context, provider string)
|
||||
}
|
||||
|
||||
// testTMDb 测试 TMDb API 连接。
|
||||
func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testTMDb(cfg *model.APIConfig) (string, error) {
|
||||
if cfg.APIKey == "" {
|
||||
return "error", errors.New("API key is required")
|
||||
}
|
||||
@@ -74,7 +74,7 @@ func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) {
|
||||
}
|
||||
|
||||
// testOpenAI 测试 OpenAI API 连接。
|
||||
func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testOpenAI(cfg *model.APIConfig) (string, error) {
|
||||
if cfg.APIKey == "" {
|
||||
return "error", errors.New("API key is required")
|
||||
}
|
||||
@@ -108,7 +108,7 @@ func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) {
|
||||
}
|
||||
|
||||
// testDeepSeek 测试 DeepSeek API 连接。
|
||||
func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testDeepSeek(cfg *model.APIConfig) (string, error) {
|
||||
if cfg.APIKey == "" {
|
||||
return "error", errors.New("API key is required")
|
||||
}
|
||||
@@ -142,7 +142,7 @@ func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) {
|
||||
}
|
||||
|
||||
// testSiliconFlow 测试 SiliconFlow API 连接。
|
||||
func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testSiliconFlow(cfg *model.APIConfig) (string, error) {
|
||||
if cfg.APIKey == "" {
|
||||
return "error", errors.New("API key is required")
|
||||
}
|
||||
@@ -176,7 +176,7 @@ func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error)
|
||||
}
|
||||
|
||||
// testAdult 测试 Adult (JavDB/JavBus) 刮削数据源连接与年龄验证。
|
||||
func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.APIConfig) (string, error) {
|
||||
bases := []string{}
|
||||
if cfg.BaseURL != "" {
|
||||
bases = append(bases, adultConfiguredBases(cfg.BaseURL)...)
|
||||
@@ -245,7 +245,7 @@ func (s *ApiConfigService) testAdult(ctx context.Context, cfg *model.ApiConfig)
|
||||
}
|
||||
|
||||
// testMetaTube 测试 MetaTube Server 连接与 Token。
|
||||
func (s *ApiConfigService) testMetaTube(ctx context.Context, cfg *model.ApiConfig) (string, error) {
|
||||
func (s *ApiConfigService) testMetaTube(ctx context.Context, cfg *model.APIConfig) (string, error) {
|
||||
serverURL := strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
|
||||
if serverURL == "" {
|
||||
serverURL = "http://127.0.0.1:7700"
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
)
|
||||
|
||||
// GetEffectiveConfig 获取生效的 API 配置(数据库配置优先于配置文件)。
|
||||
func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider string) (*model.ApiConfig, error) {
|
||||
func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider string) (*model.APIConfig, error) {
|
||||
// 首先尝试从数据库获取
|
||||
cfg, err := s.GetByProvider(ctx, provider)
|
||||
if err == nil && cfg != nil {
|
||||
@@ -21,7 +21,7 @@ func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider stri
|
||||
}
|
||||
|
||||
// getConfigFromFile 从配置文件获取 API 配置。
|
||||
func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig, error) {
|
||||
func (s *ApiConfigService) getConfigFromFile(provider string) (*model.APIConfig, error) {
|
||||
var apiKey string
|
||||
var hasKey bool
|
||||
|
||||
@@ -44,7 +44,7 @@ func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig,
|
||||
return nil, ErrApiConfigNotFound
|
||||
}
|
||||
|
||||
return &model.ApiConfig{
|
||||
return &model.APIConfig{
|
||||
Provider: provider,
|
||||
APIKey: apiKey,
|
||||
Enabled: true,
|
||||
|
||||
@@ -33,7 +33,7 @@ var (
|
||||
)
|
||||
|
||||
// GetByProvider 获取指定提供者的 API 配置。
|
||||
func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
|
||||
func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (*model.APIConfig, error) {
|
||||
cfg, err := s.repo.ApiConfig.FindByProvider(ctx, provider)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -49,7 +49,7 @@ func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (
|
||||
}
|
||||
|
||||
// List 返回所有 API 配置。
|
||||
func (s *ApiConfigService) List(ctx context.Context) ([]model.ApiConfig, error) {
|
||||
func (s *ApiConfigService) List(ctx context.Context) ([]model.APIConfig, error) {
|
||||
configs, err := s.repo.ApiConfig.List(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -69,7 +69,7 @@ func (s *ApiConfigService) GetProviders() []model.ApiProvider {
|
||||
}
|
||||
|
||||
// Upsert 创建或更新 API 配置,自动加密敏感字段。
|
||||
func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) (*model.ApiConfig, error) {
|
||||
func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) (*model.APIConfig, error) {
|
||||
// 验证提供者是否有效
|
||||
if !s.isValidProvider(provider) {
|
||||
return nil, ErrInvalidProvider
|
||||
@@ -81,7 +81,7 @@ func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey,
|
||||
encryptedKey = s.crypto.Encrypt(apiKey)
|
||||
}
|
||||
|
||||
cfg := &model.ApiConfig{
|
||||
cfg := &model.APIConfig{
|
||||
Provider: provider,
|
||||
APIKey: encryptedKey,
|
||||
BaseURL: baseURL,
|
||||
@@ -112,7 +112,7 @@ func (s *ApiConfigService) Update(ctx context.Context, provider string, apiKey,
|
||||
encryptedKey = s.crypto.Encrypt(apiKey)
|
||||
}
|
||||
|
||||
cfg := &model.ApiConfig{
|
||||
cfg := &model.APIConfig{
|
||||
Provider: provider,
|
||||
APIKey: encryptedKey,
|
||||
BaseURL: baseURL,
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
@@ -50,6 +51,8 @@ func (a *AuditService) RecordBestEffort(userID, action, target, ip, detail strin
|
||||
go func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
a.Record(ctx, userID, action, target, ip, detail)
|
||||
helper.Run(a.log, "audit.record", func() {
|
||||
a.Record(ctx, userID, action, target, ip, detail)
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
@@ -176,9 +177,11 @@ func (s *AuthService) touchLoginBestEffort(userID string) {
|
||||
go func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
if err := s.repo.User.TouchLogin(ctx, userID); err != nil && s.log != nil {
|
||||
s.log.Debug("touch login delayed", zap.String("user_id", userID), zap.Error(err))
|
||||
}
|
||||
helper.Run(s.log, "auth.touchLogin", func() {
|
||||
if err := s.repo.User.TouchLogin(ctx, userID); err != nil && s.log != nil {
|
||||
s.log.Debug("touch login delayed", zap.String("user_id", userID), zap.Error(err))
|
||||
}
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,8 @@ import (
|
||||
"net/url"
|
||||
"path"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
|
||||
@@ -22,11 +24,18 @@ type cloudDrive2Provider struct {
|
||||
base *url.URL
|
||||
username string
|
||||
password string
|
||||
token string
|
||||
token string // 配置的静态令牌(构造后只读)
|
||||
ua string
|
||||
apiBase *url.URL
|
||||
client *http.Client
|
||||
proxy bool
|
||||
|
||||
// tokenMu / loginToken / loginTokenSeen 保护 OpenList 用户名密码登录的
|
||||
// token 缓存:多 worker 并发时单飞登录,缓存有效期内直接复用,
|
||||
// 401 时清缓存重登(见 clouddrive2_openlist.go 的 openListAPIToken)。
|
||||
tokenMu sync.Mutex
|
||||
loginToken string
|
||||
loginTokenSeen time.Time
|
||||
}
|
||||
|
||||
func newCloudDrive2(cfg map[string]any, client *http.Client) *cloudDrive2Provider {
|
||||
|
||||
@@ -36,9 +36,10 @@ func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, p.decorateDAVStatusError(resp, target)
|
||||
}
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
|
||||
// 流式解码:超大目录(如上万条目的网盘目录)响应可能远超旧 4MB 截断上限,
|
||||
// 直接 xml.Unmarshal 会截断报错;这里用 LimitReader(64MB) + Decoder 边读边解
|
||||
var multi cloudDAVMultiStatus
|
||||
if err := xml.Unmarshal(body, &multi); err != nil {
|
||||
if err := xml.NewDecoder(io.LimitReader(resp.Body, 64<<20)).Decode(&multi); err != nil {
|
||||
return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err)
|
||||
}
|
||||
basePath := strings.TrimRight(p.base.EscapedPath(), "/")
|
||||
|
||||
@@ -131,10 +131,13 @@ func (p *cloudDrive2Provider) openListAPIMove(ctx context.Context, source, targe
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath string, payload any, action string) error {
|
||||
token, err := p.openListAPIToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := doWithOpenListAPIToken(ctx, p, func(token string) (struct{}, error) {
|
||||
return struct{}{}, p.openListAPIPostWithToken(ctx, apiPath, payload, action, token)
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) openListAPIPostWithToken(ctx context.Context, apiPath string, payload any, action, token string) error {
|
||||
body, _ := json.Marshal(payload)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body))
|
||||
if err != nil {
|
||||
@@ -151,6 +154,9 @@ func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath strin
|
||||
return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return errOpenListAPITokenExpired
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("%s: api %s returned http %d", p.name, action, resp.StatusCode)
|
||||
}
|
||||
@@ -200,27 +206,43 @@ func (p *cloudDrive2Provider) PutFile(ctx context.Context, remotePath string, r
|
||||
}
|
||||
|
||||
// openListAPIPutFile 通过 OpenList /api/fs/form 上传(QMediaSync 同款契约:
|
||||
// PUT + multipart + File-Path 头)。
|
||||
// PUT + multipart + File-Path 头)。使用 io.Pipe + multipart.Writer 边写边发,
|
||||
// 避免把整个文件读进内存。
|
||||
func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath string, r io.Reader) error {
|
||||
token, err := p.openListAPIToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
encodedPath := openListPathEscape(remotePath)
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
formFile, err := writer.CreateFormFile("file", path.Base(remotePath))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := io.Copy(formFile, r); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writer.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPut, p.openListAPIURL("/api/fs/form"), body)
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
writer := multipart.NewWriter(pw)
|
||||
go func() {
|
||||
var writeErr error
|
||||
defer func() {
|
||||
// 读源失败必须传给 pipe 写端,让 HTTP 请求以失败收场而不是静默截断
|
||||
if writeErr != nil {
|
||||
_ = pw.CloseWithError(writeErr)
|
||||
return
|
||||
}
|
||||
_ = pw.Close()
|
||||
}()
|
||||
formFile, err := writer.CreateFormFile("file", path.Base(remotePath))
|
||||
if err != nil {
|
||||
writeErr = err
|
||||
return
|
||||
}
|
||||
if _, err := io.Copy(formFile, r); err != nil {
|
||||
writeErr = err
|
||||
return
|
||||
}
|
||||
writeErr = writer.Close()
|
||||
}()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPut, p.openListAPIURL("/api/fs/form"), pr)
|
||||
if err != nil {
|
||||
// 关闭读端以释放仍在等待写入的后台 goroutine(其 Write 会立即失败返回)
|
||||
_ = pr.Close()
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Authorization", token)
|
||||
@@ -230,9 +252,15 @@ func (p *cloudDrive2Provider) openListAPIPutFile(ctx context.Context, remotePath
|
||||
req.Header.Set("Overwrite", "true")
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
// 传输层失败(含提前断开)时 net/http 会关闭请求 body,解除后台 goroutine 阻塞
|
||||
return decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/form"), err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
// 流式 body 无法重放,不能自动重试:清除登录 token 缓存让下次上传重新登录,
|
||||
// 本次返回明确错误交由调用方重试
|
||||
p.invalidateOpenListAPIToken()
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return p.openListAPIStatusError("upload", remotePath, resp.StatusCode)
|
||||
}
|
||||
|
||||
@@ -4,18 +4,50 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) {
|
||||
// errOpenListAPITokenExpired 标记 OpenList 返回 401(登录 token 已失效):
|
||||
// 调用方收到后应清缓存重登一次再重试原请求。
|
||||
var errOpenListAPITokenExpired = errors.New("openlist api token expired")
|
||||
|
||||
// openListAPITokenCacheTTL 登录 token 缓存有效期(OpenList 默认签发 48h JWT,
|
||||
// 这里保守取 30 分钟,过期自动重新登录)。
|
||||
const openListAPITokenCacheTTL = 30 * time.Minute
|
||||
|
||||
// doWithOpenListAPIToken 获取 OpenList API token 后执行 fn;若请求命中 401
|
||||
// (登录 token 失效)则清缓存重登一次并重试,避免一次 token 轮换导致整批请求失败。
|
||||
func doWithOpenListAPIToken[T any](ctx context.Context, p *cloudDrive2Provider, fn func(token string) (T, error)) (T, error) {
|
||||
var zero T
|
||||
token, err := p.openListAPIToken(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return zero, err
|
||||
}
|
||||
result, err := fn(token)
|
||||
if err == nil || !errors.Is(err, errOpenListAPITokenExpired) {
|
||||
return result, err
|
||||
}
|
||||
p.invalidateOpenListAPIToken()
|
||||
token, err = p.openListAPIToken(ctx)
|
||||
if err != nil {
|
||||
return zero, err
|
||||
}
|
||||
return fn(token)
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) {
|
||||
return doWithOpenListAPIToken(ctx, p, func(token string) ([]FileEntry, error) {
|
||||
return p.listOpenListAPIWithToken(ctx, dir, token)
|
||||
})
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) listOpenListAPIWithToken(ctx context.Context, dir, token string) ([]FileEntry, error) {
|
||||
const pageSize = 500
|
||||
target := normalizeCloudDAVPath(dir)
|
||||
out := make([]FileEntry, 0, pageSize)
|
||||
@@ -45,6 +77,9 @@ func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) (
|
||||
var decoded openListListResponse
|
||||
decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded)
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return nil, errOpenListAPITokenExpired
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, p.openListAPIStatusError("list", target, resp.StatusCode)
|
||||
}
|
||||
@@ -85,10 +120,12 @@ func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) (
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
|
||||
token, err := p.openListAPIToken(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return doWithOpenListAPIToken(ctx, p, func(token string) (*DirectLink, error) {
|
||||
return p.resolveOpenListAPIDirectWithToken(ctx, fileRef, token)
|
||||
})
|
||||
}
|
||||
|
||||
func (p *cloudDrive2Provider) resolveOpenListAPIDirectWithToken(ctx context.Context, fileRef, token string) (*DirectLink, error) {
|
||||
payload, _ := json.Marshal(map[string]string{"path": normalizeCloudDAVPath(fileRef), "password": ""})
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/get"), bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
@@ -105,6 +142,9 @@ func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, file
|
||||
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return nil, errOpenListAPITokenExpired
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, p.openListAPIStatusError("get", fileRef, resp.StatusCode)
|
||||
}
|
||||
@@ -163,6 +203,10 @@ func (p *cloudDrive2Provider) hasOpenListAPICredentials() bool {
|
||||
return strings.TrimSpace(p.token) != "" || (strings.TrimSpace(p.username) != "" && p.password != "")
|
||||
}
|
||||
|
||||
// openListAPIToken 返回 OpenList API 访问令牌:
|
||||
// - 配置了静态 token 时直接使用(构造后只读,无并发问题);
|
||||
// - 否则用用户名密码登录,并在缓存有效期内单飞复用——8 个同步 worker 并发时
|
||||
// 只会有一个 goroutine 真正执行登录,避免登录风暴;登录 token 的写入受 tokenMu 保护。
|
||||
func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) {
|
||||
if token := strings.TrimSpace(p.token); token != "" {
|
||||
return token, nil
|
||||
@@ -170,6 +214,30 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err
|
||||
if strings.TrimSpace(p.username) == "" || p.password == "" {
|
||||
return "", nil
|
||||
}
|
||||
p.tokenMu.Lock()
|
||||
defer p.tokenMu.Unlock()
|
||||
if p.loginToken != "" && time.Since(p.loginTokenSeen) < openListAPITokenCacheTTL {
|
||||
return p.loginToken, nil
|
||||
}
|
||||
token, err := p.openListAPILogin(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
p.loginToken = token
|
||||
p.loginTokenSeen = time.Now()
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// invalidateOpenListAPIToken 清除登录 token 缓存(收到 401 时调用,下次请求重新登录)。
|
||||
func (p *cloudDrive2Provider) invalidateOpenListAPIToken() {
|
||||
p.tokenMu.Lock()
|
||||
p.loginToken = ""
|
||||
p.loginTokenSeen = time.Time{}
|
||||
p.tokenMu.Unlock()
|
||||
}
|
||||
|
||||
// openListAPILogin 调用 OpenList /api/auth/login 换取登录 token。
|
||||
func (p *cloudDrive2Provider) openListAPILogin(ctx context.Context) (string, error) {
|
||||
payload, _ := json.Marshal(map[string]string{
|
||||
"username": p.username,
|
||||
"password": p.password,
|
||||
@@ -204,7 +272,6 @@ func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, err
|
||||
if token == "" {
|
||||
return "", fmt.Errorf("%s: api login returned empty token", p.name)
|
||||
}
|
||||
p.token = token
|
||||
return token, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -48,7 +48,7 @@ func (p *openAPI115Provider) Ping(ctx context.Context) error {
|
||||
if strings.TrimSpace(p.c.AppID) == "" {
|
||||
return fmt.Errorf("115: 缺少开放平台应用 ID,请重新授权")
|
||||
}
|
||||
if strings.TrimSpace(p.c.AccessToken) == "" {
|
||||
if strings.TrimSpace(p.c.CurrentAccessToken()) == "" {
|
||||
return fmt.Errorf("115: 缺少访问令牌,请重新授权")
|
||||
}
|
||||
_, _, err := p.c.GetFsList(ctx, "0", 0, 1)
|
||||
@@ -70,7 +70,7 @@ func (p *openAPI115Provider) List(ctx context.Context, dirID string) ([]FileEntr
|
||||
Name: f.FileName,
|
||||
IsDir: f.Category == cloud115.TypeDir,
|
||||
Size: f.FileSize,
|
||||
MTime: f.Utime,
|
||||
MTime: f.ModifiedAt(),
|
||||
PickCode: f.PickCode,
|
||||
})
|
||||
}
|
||||
@@ -127,12 +127,15 @@ func (p *openAPI115Provider) PutFileNamed(ctx context.Context, parentCID, fileNa
|
||||
if err := tmp.Close(); err != nil {
|
||||
return fmt.Errorf("115: 关闭临时文件失败:%w", err)
|
||||
}
|
||||
// 重命名为目标文件名,保证上传到 115 后保留原始文件名
|
||||
// 重命名为目标文件名,保证上传到 115 后保留原始文件名。
|
||||
// 重命名失败必须 fail fast:静默用随机临时名上传会导致 115 上的文件名
|
||||
// 变成 mebox-upload-xxx,破坏元数据文件名契约。
|
||||
if fileName != "" && fileName != filepath.Base(tmpPath) {
|
||||
namedPath := filepath.Join(filepath.Dir(tmpPath), fileName)
|
||||
if err := os.Rename(tmpPath, namedPath); err == nil {
|
||||
tmpPath = namedPath
|
||||
if err := os.Rename(tmpPath, namedPath); err != nil {
|
||||
return fmt.Errorf("115: 重命名临时文件为 %s 失败:%w", fileName, err)
|
||||
}
|
||||
tmpPath = namedPath
|
||||
}
|
||||
_, err = p.c.Upload(ctx, tmpPath, parentCID, "", "")
|
||||
if err != nil {
|
||||
|
||||
@@ -23,9 +23,14 @@ type OpenClient struct {
|
||||
RefreshTokenStr string
|
||||
executor *QueueExecutor
|
||||
|
||||
// tokenMu 保护令牌刷新:业务请求中途 access_token 失效时自动刷新重试,
|
||||
// 多 goroutine(同步列表 + 下载队列)并发下只允许一次刷新进行。
|
||||
tokenMu sync.Mutex
|
||||
// OnTokenRefreshed 在 access_token 刷新成功后回调(参数为新令牌对),
|
||||
// 供上层持久化新令牌使用;nil 安全,且在 tokenMu 释放后调用以避免死锁。
|
||||
OnTokenRefreshed func(accessToken, refreshToken string)
|
||||
|
||||
// tokenMu 保护 AccessToken / RefreshTokenStr 的并发读写:业务请求中途
|
||||
// access_token 失效时自动刷新重试,多 goroutine(同步列表 + 下载队列)
|
||||
// 并发下只允许一次刷新进行。
|
||||
tokenMu sync.RWMutex
|
||||
}
|
||||
|
||||
// default115HTTPClient 创建带有防 405 重定向保护的 http.Client。
|
||||
@@ -57,12 +62,40 @@ func NewOpenClient(appID, accessToken, refreshToken string) *OpenClient {
|
||||
}
|
||||
}
|
||||
|
||||
// SetAuthToken 更新认证令牌。
|
||||
// SetAuthToken 更新认证令牌(并发安全)。
|
||||
func (c *OpenClient) SetAuthToken(accessToken, refreshToken string) {
|
||||
c.tokenMu.Lock()
|
||||
c.setAuthTokenLocked(accessToken, refreshToken)
|
||||
c.tokenMu.Unlock()
|
||||
}
|
||||
|
||||
// setAuthTokenLocked 无锁更新令牌,调用方必须已持有 tokenMu 写锁
|
||||
// (tryRefreshTokenLocked 等已持锁流程内部使用,避免重入死锁)。
|
||||
func (c *OpenClient) setAuthTokenLocked(accessToken, refreshToken string) {
|
||||
c.AccessToken = accessToken
|
||||
c.RefreshTokenStr = refreshToken
|
||||
}
|
||||
|
||||
// currentAccessToken 返回当前 access_token(并发安全)。
|
||||
func (c *OpenClient) currentAccessToken() string {
|
||||
c.tokenMu.RLock()
|
||||
defer c.tokenMu.RUnlock()
|
||||
return c.AccessToken
|
||||
}
|
||||
|
||||
// currentRefreshToken 返回当前 refresh_token(并发安全)。
|
||||
func (c *OpenClient) currentRefreshToken() string {
|
||||
c.tokenMu.RLock()
|
||||
defer c.tokenMu.RUnlock()
|
||||
return c.RefreshTokenStr
|
||||
}
|
||||
|
||||
// CurrentAccessToken 返回当前 access_token 快照(并发安全),
|
||||
// 供上层在无锁环境下安全读取(如 Ping 时探测令牌是否存在)。
|
||||
func (c *OpenClient) CurrentAccessToken() string {
|
||||
return c.currentAccessToken()
|
||||
}
|
||||
|
||||
// RespState 兼容 115 不同端点返回的 state 类型(proapi 返回布尔、passport 返回数字)。
|
||||
type RespState bool
|
||||
|
||||
@@ -191,6 +224,14 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map
|
||||
if access {
|
||||
// 刷新失败(或已刷新仍失败)时返回明确错误
|
||||
lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115: access_token 校验失败且刷新未成功")
|
||||
} else {
|
||||
// 未携带令牌的请求(登录/刷新流程)命中 token 类错误码:
|
||||
// 必须返回显式 error,避免调用方把 (resp, nil) 当作成功处理
|
||||
msg := base.Message
|
||||
if msg == "" {
|
||||
msg = base.Error
|
||||
}
|
||||
lastErr = fmt.Errorf("115: 认证失败(code=%d): %s", base.Code, msg)
|
||||
}
|
||||
return &base, lastErr
|
||||
}
|
||||
@@ -242,8 +283,11 @@ func (c *OpenClient) buildRequestWithUA(ctx context.Context, method, rawURL stri
|
||||
if method == http.MethodPost && len(form) > 0 {
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
}
|
||||
if access && c.AccessToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.AccessToken)
|
||||
if access {
|
||||
// RLock 读取令牌,避免与刷新流程的写入产生数据竞争
|
||||
if accessToken := c.currentAccessToken(); accessToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
}
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
@@ -261,33 +305,57 @@ func (c *OpenClient) doAuthJSONWithUA(ctx context.Context, method, rawURL string
|
||||
// tryRefreshTokenLocked 并发安全地刷新 access_token;成功返回 true(调用方
|
||||
// 应使用内存中的新 token 重试原请求)。
|
||||
//
|
||||
// 拿到写锁后在锁内读取 oldAccess,与持锁期间的当前值对比:若已被其他
|
||||
// goroutine 刷新过则直接复用新 token,避免并发请求连环轮转消耗 115 的
|
||||
// 一次性 refresh_token。全程持写锁读写 token 字段,无 TOCTOU 窗口。
|
||||
//
|
||||
// 对"refresh_token 本身已失效/被吊销"(IsRefreshTokenDead,如 40140114/116/119/120)
|
||||
// 这类不可恢复的错误直接放弃并清空内存 token(提示需重新授权)。
|
||||
// 对其它失败(网络瞬时抖动、刷新接口可重试错误码等)做指数退避重试几次再放弃,
|
||||
// 避免同步长任务中途 token 到期时恰好撞上一个短暂的刷新失败就整体失败。
|
||||
func (c *OpenClient) tryRefreshTokenLocked(ctx context.Context) bool {
|
||||
c.tokenMu.Lock()
|
||||
defer c.tokenMu.Unlock()
|
||||
// 在已持有写锁内读取当前 token 作为"刷新前快照",消除双重加锁窗口:
|
||||
// 若在拿锁期间已有其他 goroutine 完成刷新,refreshTokenWhileLocked
|
||||
// 内的 c.AccessToken != oldAccess 判断会立即命中并返回复用。
|
||||
oldAccess := c.AccessToken
|
||||
newToken, ok := c.refreshTokenWhileLocked(ctx, oldAccess)
|
||||
c.tokenMu.Unlock()
|
||||
// 回调必须在 tokenMu 释放后调用,避免上层在回调内访问客户端时死锁
|
||||
if ok && newToken != nil && c.OnTokenRefreshed != nil {
|
||||
c.OnTokenRefreshed(newToken.AccessToken, newToken.RefreshToken)
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
// refreshTokenWhileLocked 在已持有 tokenMu 写锁的前提下执行刷新。
|
||||
// 返回 (新令牌, 是否成功);命中"他人已刷新"捷径时新令牌为 nil。
|
||||
func (c *OpenClient) refreshTokenWhileLocked(ctx context.Context, oldAccess string) (*TokenData, bool) {
|
||||
if c.AccessToken != oldAccess {
|
||||
// 其他 goroutine 刚刷新过:直接复用内存中的新 token 重试原请求
|
||||
return nil, true
|
||||
}
|
||||
refreshToken := c.RefreshTokenStr
|
||||
for attempt := 0; attempt < refreshAttempts; attempt++ {
|
||||
token, err := c.RefreshToken(c.RefreshTokenStr)
|
||||
token, err := c.doRefreshToken(refreshToken)
|
||||
if err == nil {
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return true
|
||||
c.setAuthTokenLocked(token.AccessToken, token.RefreshToken)
|
||||
return token, true
|
||||
}
|
||||
if IsRefreshTokenDead(err) {
|
||||
c.SetAuthToken("", "")
|
||||
return false
|
||||
c.setAuthTokenLocked("", "")
|
||||
return nil, false
|
||||
}
|
||||
// 可恢复失败:退避后重试。ctx 取消时立即放弃。
|
||||
if attempt < refreshAttempts-1 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
return nil, false
|
||||
case <-time.After(refreshBackoff(attempt)):
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// refreshAttempts 是刷新 access_token 失败时的最大尝试次数(含首次)。
|
||||
@@ -312,7 +380,14 @@ func isTokenCode(code int) bool {
|
||||
}
|
||||
|
||||
// openList 解析 data 为对象或数组(StructOrArray 语义)。
|
||||
// 115 部分接口在鉴权/业务异常时会返回 data:null 或 data:{},此时若直接
|
||||
// 反序列化会得到零值元素 + nil error,调用方会把空数据当成功处理;
|
||||
// 这里对 null/空对象显式报错。
|
||||
func openList[T any](raw json.RawMessage) ([]T, error) {
|
||||
trimmed := bytes.TrimSpace(raw)
|
||||
if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) || bytes.Equal(trimmed, []byte("{}")) {
|
||||
return nil, fmt.Errorf("115: data 为空(%s)", string(trimmed))
|
||||
}
|
||||
var single T
|
||||
if err := json.Unmarshal(raw, &single); err == nil {
|
||||
return []T{single}, nil
|
||||
@@ -324,12 +399,16 @@ func openList[T any](raw json.RawMessage) ([]T, error) {
|
||||
return nil, fmt.Errorf("115: data 既不是对象也不是数组")
|
||||
}
|
||||
|
||||
// openFirstList 取 data 的第一个元素。
|
||||
// openFirstList 取 data 的第一个元素;data 为空(null/空数组)时返回显式错误,
|
||||
// 避免调用方拿到 (nil, nil) 后解引用空指针。
|
||||
func openFirstList[T any](raw json.RawMessage) (*T, error) {
|
||||
items, err := openList[T](raw)
|
||||
if err != nil || len(items) == 0 {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return nil, fmt.Errorf("115: data 为空数组")
|
||||
}
|
||||
return &items[0], nil
|
||||
}
|
||||
|
||||
|
||||
@@ -317,13 +317,18 @@ func appendCallbackParams(rawURL string, params url.Values) (string, error) {
|
||||
return callbackURL.String(), nil
|
||||
}
|
||||
|
||||
// oauthHTTPClient 是 OAuth 授权服务专用 HTTP 客户端。http.DefaultClient 无超时,
|
||||
// 授权服务无响应时会永久阻塞授权/轮询协程,这里统一 30s 超时(ctx 仍经
|
||||
// NewRequestWithContext 传导,可提前取消)。
|
||||
var oauthHTTPClient = &http.Client{Timeout: 30 * time.Second}
|
||||
|
||||
func httpGetJSON(ctx context.Context, endpoint string) (map[string]any, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("User-Agent", DefaultUA)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
resp, err := oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -300,6 +300,11 @@ func (c *OpenClient) GetQrCode() (*QrCodeDataReturn, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 关键字段缺失时显式报错:空 uid/sign 会导致后续扫码轮询必然失败,
|
||||
// 不能把残缺响应当成功返回给界面。
|
||||
if code.Uid == "" || code.Sign == "" {
|
||||
return nil, fmt.Errorf("115: 设备码响应缺少 uid/sign,无法发起扫码授权")
|
||||
}
|
||||
return &QrCodeDataReturn{QrCodeData: *code, CodeVerifier: codeVerifier}, nil
|
||||
}
|
||||
|
||||
@@ -352,6 +357,10 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 空凭证绝不能 SetAuthToken 后当成功返回:界面会显示"授权成功"但账号不可用
|
||||
if token.AccessToken == "" || token.RefreshToken == "" {
|
||||
return nil, fmt.Errorf("115: 设备码换 token 返回空凭证(access_token/refresh_token 缺失)")
|
||||
}
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return token, nil
|
||||
}
|
||||
@@ -359,11 +368,30 @@ func (c *OpenClient) GetToken(qrCode *QrCodeDataReturn) (*TokenData, error) {
|
||||
// RefreshToken 刷新访问令牌。
|
||||
func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) {
|
||||
if refreshToken == "" {
|
||||
refreshToken = c.RefreshTokenStr
|
||||
refreshToken = c.currentRefreshToken()
|
||||
}
|
||||
if refreshToken == "" {
|
||||
return nil, fmt.Errorf("没有可用的 refresh_token")
|
||||
}
|
||||
token, err := c.doRefreshToken(refreshToken)
|
||||
if err != nil {
|
||||
// refresh_token 已失效时清空内存令牌(提示需重新授权)
|
||||
if IsRefreshTokenDead(err) {
|
||||
c.SetAuthToken("", "")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if token.AccessToken == "" || token.RefreshToken == "" {
|
||||
return nil, fmt.Errorf("115: 刷新返回空凭证(access_token/refresh_token 缺失)")
|
||||
}
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// doRefreshToken 调用 115 刷新接口换取新令牌,不修改客户端内存状态;
|
||||
// 拆出无状态方法供 tryRefreshTokenLocked(已持 tokenMu 写锁)复用,
|
||||
// 避免在持锁期间重入 SetAuthToken 造成死锁。
|
||||
func (c *OpenClient) doRefreshToken(refreshToken string) (*TokenData, error) {
|
||||
params := map[string]string{"refresh_token": refreshToken}
|
||||
resp, err := c.doJSON(context.Background(), "POST", PassportAPIBase+"/open/refreshToken", params, false, 0)
|
||||
if err != nil && resp == nil {
|
||||
@@ -373,18 +401,9 @@ func (c *OpenClient) RefreshToken(refreshToken string) (*TokenData, error) {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.State {
|
||||
apiErr := NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败")
|
||||
if IsRefreshTokenDead(apiErr) {
|
||||
c.SetAuthToken("", "")
|
||||
}
|
||||
return nil, apiErr
|
||||
return nil, NewOpenAPIResponseError(resp.Code, resp.Errno, resp.Message, resp.Error, "115 开放平台刷新访问凭证失败")
|
||||
}
|
||||
token, err := openFirstList[TokenData](resp.Data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.SetAuthToken(token.AccessToken, token.RefreshToken)
|
||||
return token, nil
|
||||
return openFirstList[TokenData](resp.Data)
|
||||
}
|
||||
|
||||
// ─── 用户信息 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"os"
|
||||
"sort"
|
||||
|
||||
@@ -106,14 +107,23 @@ func (u *OSSMultipartUploader) UploadFile(ctx context.Context, input OSSMultipar
|
||||
return result.CallbackResult, nil
|
||||
}
|
||||
|
||||
// UploadedPart 是 OSS 已上传分片的定位信息(断点续传时复用 ETag 用)。
|
||||
type UploadedPart struct {
|
||||
PartNumber int32
|
||||
Size int64
|
||||
ETag string
|
||||
}
|
||||
|
||||
// UploadFileWithResult 上传文件并返回 multipart 结果。
|
||||
func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input OSSMultipartUploadInput) (OSSMultipartUploadResult, error) {
|
||||
// 任一失败路径(分片上传失败 / callback 校验失败 / Complete 失败 / 文件打开失败等)
|
||||
// 都会经 defer 统一 AbortMultipartUpload 丢弃本次 Initiate 出的 multipart
|
||||
// (abort 失败仅记日志),避免 OSS 分片永久泄漏;成功路径不 Abort。
|
||||
func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input OSSMultipartUploadInput) (result OSSMultipartUploadResult, err error) {
|
||||
if input.PartRetryMax <= 0 {
|
||||
input.PartRetryMax = 3
|
||||
}
|
||||
partSize := input.PartSize
|
||||
totalParts := 0
|
||||
var err error
|
||||
if partSize <= 0 {
|
||||
partSize, totalParts, err = CalculateMultipartPartSize(input.FileSize)
|
||||
if err != nil {
|
||||
@@ -124,28 +134,45 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
|
||||
}
|
||||
|
||||
uploadId := input.UploadId
|
||||
if uploadId == "" {
|
||||
initResult, err := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{
|
||||
// ownUploadId 标记 uploadId 是否为本调用 Initiate 出来的:仅自建的
|
||||
// multipart 在失败时由本函数 Abort;调用方显式传入的 uploadId(断点续传)
|
||||
// 失败后保留现场,由调用方决定重试或清理。
|
||||
ownUploadId := uploadId == ""
|
||||
if ownUploadId {
|
||||
initResult, initErr := u.client.InitiateMultipartUpload(ctx, &oss.InitiateMultipartUploadRequest{
|
||||
Bucket: oss.Ptr(input.Bucket),
|
||||
Key: oss.Ptr(input.Object),
|
||||
RequestCommon: oss.RequestCommon{
|
||||
Parameters: map[string]string{"sequential": "1"},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", err)
|
||||
if initErr != nil {
|
||||
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 失败:%w", initErr)
|
||||
}
|
||||
if initResult.UploadId == nil || *initResult.UploadId == "" {
|
||||
return OSSMultipartUploadResult{}, fmt.Errorf("初始化 OSS multipart 返回空 upload_id")
|
||||
}
|
||||
uploadId = *initResult.UploadId
|
||||
}
|
||||
defer func() {
|
||||
if err == nil || !ownUploadId || uploadId == "" {
|
||||
return
|
||||
}
|
||||
// 失败路径统一 Abort 丢弃已上传分片;ctx 可能已取消,脱离其取消信号尽力清理
|
||||
abortCtx := context.WithoutCancel(ctx)
|
||||
if _, abortErr := u.client.AbortMultipartUpload(abortCtx, &oss.AbortMultipartUploadRequest{
|
||||
Bucket: oss.Ptr(input.Bucket),
|
||||
Key: oss.Ptr(input.Object),
|
||||
UploadId: oss.Ptr(uploadId),
|
||||
}); abortErr != nil {
|
||||
log.Printf("115: 中止 OSS multipart 失败(upload_id=%s,可能残留分片):%v", uploadId, abortErr)
|
||||
}
|
||||
}()
|
||||
|
||||
existingPartMap := make(map[int32]int64)
|
||||
existingParts, err := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId)
|
||||
if err == nil {
|
||||
existingPartMap := make(map[int32]UploadedPart)
|
||||
if existingParts, listErr := u.ListUploadedParts(ctx, input.Bucket, input.Object, uploadId); listErr == nil {
|
||||
for _, part := range existingParts {
|
||||
existingPartMap[part.PartNumber] = part.Size
|
||||
existingPartMap[part.PartNumber] = part
|
||||
}
|
||||
}
|
||||
|
||||
@@ -164,13 +191,20 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
|
||||
if length < 0 {
|
||||
length = 0
|
||||
}
|
||||
if existingSize, ok := existingPartMap[int32(partNumber)]; ok && existingSize == length {
|
||||
// 断点续传:分片已完整上传(大小一致即代表分片大小未变)时直接复用
|
||||
// ListParts 返回的 ETag,跳过重传,也不再重复累加统计
|
||||
if existing, ok := existingPartMap[int32(partNumber)]; ok && existing.Size == length && existing.ETag != "" {
|
||||
uploadedBytes += length
|
||||
uploadedParts++
|
||||
completeParts = append(completeParts, oss.UploadPart{
|
||||
PartNumber: int32(partNumber),
|
||||
ETag: oss.Ptr(existing.ETag),
|
||||
})
|
||||
continue
|
||||
}
|
||||
etag, err := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length)
|
||||
if err != nil {
|
||||
return OSSMultipartUploadResult{}, err
|
||||
etag, uploadErr := u.uploadPartWithRetry(ctx, input, uploadId, int32(partNumber), file, offset, length)
|
||||
if uploadErr != nil {
|
||||
return OSSMultipartUploadResult{}, uploadErr
|
||||
}
|
||||
uploadedBytes += length
|
||||
uploadedParts++
|
||||
@@ -224,29 +258,34 @@ func (u *OSSMultipartUploader) UploadFileWithResult(ctx context.Context, input O
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListUploadedParts 查询 OSS 已上传分片。
|
||||
func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]struct {
|
||||
PartNumber int32
|
||||
Size int64
|
||||
}, error) {
|
||||
parts := []struct {
|
||||
PartNumber int32
|
||||
Size int64
|
||||
}{}
|
||||
result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{
|
||||
Bucket: oss.Ptr(bucket),
|
||||
Key: oss.Ptr(object),
|
||||
UploadId: oss.Ptr(uploadId),
|
||||
MaxParts: 1000,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err)
|
||||
}
|
||||
for _, part := range result.Parts {
|
||||
parts = append(parts, struct {
|
||||
PartNumber int32
|
||||
Size int64
|
||||
}{PartNumber: part.PartNumber, Size: part.Size})
|
||||
// ListUploadedParts 查询 OSS 已上传分片(MaxParts 上限 1000,超过时按
|
||||
// NextPartNumberMarker 自动翻页取全量,否则断点续传只能看到前 1000 片)。
|
||||
func (u *OSSMultipartUploader) ListUploadedParts(ctx context.Context, bucket, object, uploadId string) ([]UploadedPart, error) {
|
||||
parts := []UploadedPart{}
|
||||
var marker int32
|
||||
for {
|
||||
result, err := u.client.ListParts(ctx, &oss.ListPartsRequest{
|
||||
Bucket: oss.Ptr(bucket),
|
||||
Key: oss.Ptr(object),
|
||||
UploadId: oss.Ptr(uploadId),
|
||||
MaxParts: 1000,
|
||||
PartNumberMarker: marker,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("查询 OSS 已上传分片失败:%w", err)
|
||||
}
|
||||
for _, part := range result.Parts {
|
||||
etag := ""
|
||||
if part.ETag != nil {
|
||||
etag = *part.ETag
|
||||
}
|
||||
parts = append(parts, UploadedPart{PartNumber: part.PartNumber, Size: part.Size, ETag: etag})
|
||||
}
|
||||
if !result.IsTruncated || result.NextPartNumberMarker <= marker {
|
||||
// 防御:marker 不前进时终止循环,避免异常响应导致死循环
|
||||
break
|
||||
}
|
||||
marker = result.NextPartNumberMarker
|
||||
}
|
||||
return parts, nil
|
||||
}
|
||||
|
||||
@@ -262,7 +262,10 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s
|
||||
}
|
||||
switch status {
|
||||
case UploadInitStatusRapidUploaded:
|
||||
// 秒传成功
|
||||
// 秒传成功:必须带远端文件定位信息,否则视为异常响应
|
||||
if initResult.FileId == "" || initResult.PickCode == "" {
|
||||
return nil, fmt.Errorf("115: 秒传成功但缺少 file_id/pick_code(status=%d)", status)
|
||||
}
|
||||
return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil
|
||||
case UploadInitStatusSignFailed:
|
||||
return nil, fmt.Errorf("115: 签名验证后失败")
|
||||
@@ -271,7 +274,8 @@ func (c *OpenClient) Upload(ctx context.Context, filePath, parentCID, signKey, s
|
||||
case UploadInitStatusNeedUpload:
|
||||
// 真实上传:OSS multipart
|
||||
default:
|
||||
return &UploadCompleteResult{FileId: initResult.FileId, PickCode: initResult.PickCode}, nil
|
||||
// 未知状态不能当成功返回(会静默丢文件),显式报错便于排查
|
||||
return nil, fmt.Errorf("115: 未知的 upload/init 状态 %d", status)
|
||||
}
|
||||
|
||||
if initResult.Bucket == "" || initResult.Object == "" {
|
||||
|
||||
@@ -1,14 +1,26 @@
|
||||
package cloud115
|
||||
|
||||
import "math/rand"
|
||||
import (
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
"math/big"
|
||||
)
|
||||
|
||||
const randCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
|
||||
|
||||
// RandomString 生成指定长度的随机字符串(PKCE code_verifier 等)。
|
||||
// RandomString 生成指定长度的密码学安全随机字符串(PKCE code_verifier、
|
||||
// OAuth state 等安全敏感场景)。必须使用 crypto/rand:math/rand 未播种时
|
||||
// 序列可预测,会造成 PKCE 防御失效。
|
||||
func RandomString(length int) string {
|
||||
b := make([]byte, length)
|
||||
max := big.NewInt(int64(len(randCharset)))
|
||||
for i := range b {
|
||||
b[i] = randCharset[rand.Intn(len(randCharset))]
|
||||
n, err := rand.Int(rand.Reader, max)
|
||||
if err != nil {
|
||||
// 仅在系统熵源不可用时发生;静默降级为弱随机不可接受,直接暴露
|
||||
panic(fmt.Errorf("115: 生成安全随机字符串失败:%w", err))
|
||||
}
|
||||
b[i] = randCharset[n.Int64()]
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
@@ -36,6 +36,10 @@ type DLNAService struct {
|
||||
cachedAt time.Time
|
||||
}
|
||||
|
||||
// dlnaHTTPClient 是 DLNA 专用 HTTP 客户端:SSDP 描述拉取与 SOAP 投递
|
||||
// 都应快速失败,不占用全局 DefaultClient,也不无限悬挂。
|
||||
var dlnaHTTPClient = &http.Client{Timeout: 15 * time.Second}
|
||||
|
||||
// NewDLNAService is the constructor.
|
||||
func NewDLNAService(log *zap.Logger) *DLNAService {
|
||||
return &DLNAService{log: log}
|
||||
@@ -153,7 +157,7 @@ func (d *DLNAService) fetchDescription(ctx context.Context, location string) (*D
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
resp, err := dlnaHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -267,7 +271,7 @@ func (d *DLNAService) soap(ctx context.Context, controlURL, action, envelope str
|
||||
req.Header.Set("Content-Type", `text/xml; charset="utf-8"`)
|
||||
req.Header.Set("SOAPAction",
|
||||
fmt.Sprintf(`"urn:schemas-upnp-org:service:AVTransport:1#%s"`, action))
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
resp, err := dlnaHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
@@ -267,12 +268,19 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type remoteResult struct {
|
||||
items []any
|
||||
type remoteReply struct {
|
||||
acct *model.StrmAccount
|
||||
envelope map[string]any
|
||||
}
|
||||
mounts, aerr := e.remote.ListMounts(ctx)
|
||||
results := make([]remoteResult, 0, len(mounts))
|
||||
replies := make([]*remoteReply, 0, len(mounts))
|
||||
if aerr == nil {
|
||||
type mountSearchJob struct {
|
||||
idx int
|
||||
mount *model.EmbyMount
|
||||
acct *model.StrmAccount
|
||||
}
|
||||
jobs := make([]*mountSearchJob, 0, len(mounts))
|
||||
for i := range mounts {
|
||||
m := mounts[i]
|
||||
if !m.Enabled {
|
||||
@@ -285,40 +293,76 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[
|
||||
if acct == nil {
|
||||
continue
|
||||
}
|
||||
// 按挂载逐个搜索:搜索结果归属明确(伪装 ID 正确),也天然只搜已
|
||||
// 挂载的媒体库。
|
||||
searchParams := p
|
||||
searchParams.ParentID = "" // RemoteSearchMount 内部设 ParentId
|
||||
remote, rerr := e.remote.RemoteSearchMount(ctx, &m, acct, p)
|
||||
if rerr != nil {
|
||||
if e.log != nil {
|
||||
e.log.Warn("remote emby search failed",
|
||||
zap.String("account", acct.Name), zap.Error(rerr))
|
||||
// idx 使用 jobs 内的序号(而非 mounts 下标):fetched 按
|
||||
// len(jobs) 分配,必须与 jobs 下标对齐,否则越界 panic。
|
||||
jobs = append(jobs, &mountSearchJob{idx: len(jobs), mount: &mounts[i], acct: acct})
|
||||
}
|
||||
// 并发搜索各挂载(限并发 + 单挂载超时):串行时每挂载最多
|
||||
// 15s×线路数,多挂载下首屏延迟被成倍放大。结果按挂载顺序合并。
|
||||
sem := make(chan struct{}, 4)
|
||||
var wg sync.WaitGroup
|
||||
fetched := make([]*remoteReply, len(jobs))
|
||||
for _, job := range jobs {
|
||||
wg.Add(1)
|
||||
go func(job *mountSearchJob) {
|
||||
defer wg.Done()
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
sctx, cancel := context.WithTimeout(ctx, 8*time.Second)
|
||||
defer cancel()
|
||||
sp := p
|
||||
// 远程只取首页:此前每个远程各自按 StartIndex 分页,拼接后
|
||||
// 又被 sliceSearchItems 再切一次——分页被二次偏移,远程结果
|
||||
// 首屏不可见、翻页错位。合并后由 sliceSearchItems 单点分页。
|
||||
sp.StartIndex = 0
|
||||
sp.ParentID = "" // RemoteSearchMount 内部设 ParentId
|
||||
remote, rerr := e.remote.RemoteSearchMount(sctx, job.mount, job.acct, sp)
|
||||
if rerr != nil {
|
||||
if e.log != nil {
|
||||
e.log.Warn("remote emby search failed",
|
||||
zap.String("account", job.acct.Name), zap.Error(rerr))
|
||||
}
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := e.mergeRemoteUserData(ctx, p.UserID, remote); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw, ok := remote["Items"].([]any); ok {
|
||||
results = append(results, remoteResult{items: raw})
|
||||
} else if rawMap, ok := remote["Items"].([]map[string]any); ok {
|
||||
converted := make([]any, 0, len(rawMap))
|
||||
for _, m := range rawMap {
|
||||
converted = append(converted, any(m))
|
||||
}
|
||||
results = append(results, remoteResult{items: converted})
|
||||
fetched[job.idx] = &remoteReply{acct: job.acct, envelope: remote}
|
||||
}(job)
|
||||
}
|
||||
wg.Wait()
|
||||
for _, r := range fetched {
|
||||
if r != nil {
|
||||
replies = append(replies, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
items := make([]any, 0, len(localItemsAsAny(local))+len(results)*p.Limit)
|
||||
items := make([]any, 0, len(localItemsAsAny(local))+len(replies)*p.Limit)
|
||||
items = append(items, localItemsAsAny(local)...)
|
||||
for _, res := range results {
|
||||
items = append(items, res.items...)
|
||||
for _, reply := range replies {
|
||||
if err := e.mergeRemoteUserData(ctx, p.UserID, reply.envelope); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, remoteItemsAsAny(reply.envelope)...)
|
||||
}
|
||||
return sliceSearchItems(items, p), nil
|
||||
}
|
||||
|
||||
// remoteItemsAsAny 提取远程载荷的 Items 列表(兼容 []any 与 []map 形态)。
|
||||
func remoteItemsAsAny(envelope map[string]any) []any {
|
||||
if envelope == nil {
|
||||
return nil
|
||||
}
|
||||
if raw, ok := envelope["Items"].([]any); ok {
|
||||
return raw
|
||||
}
|
||||
if rawMap, ok := envelope["Items"].([]map[string]any); ok {
|
||||
converted := make([]any, 0, len(rawMap))
|
||||
for _, m := range rawMap {
|
||||
converted = append(converted, any(m))
|
||||
}
|
||||
return converted
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func localItemsAsAny(envelope map[string]any) []any {
|
||||
if envelope == nil {
|
||||
return nil
|
||||
|
||||
@@ -54,7 +54,7 @@ func (e *EmbyService) countVisibleSeries(ctx context.Context, userID string) (in
|
||||
for i := range rows {
|
||||
key := strings.TrimSpace(rows[i].SeriesID)
|
||||
if key == "" {
|
||||
key = stableEmbyID(embyVirtualSeriesPrefix, rows[i].LibraryID, e.seriesNameForMedia(&rows[i]))
|
||||
key = stableEmbyID(embyVirtualSeriesPrefix, rows[i].LibraryID, e.seriesNameForMedia(ctx, &rows[i]))
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
}
|
||||
|
||||
@@ -98,7 +98,8 @@ func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[str
|
||||
pos = h.PositionMs
|
||||
}
|
||||
}
|
||||
return e.itemPayload(ctx, m, fav, pos), nil
|
||||
// 单条目 payload 内部对库类型/series 标题有多次查找,挂请求级缓存合并。
|
||||
return e.itemPayload(e.withPayloadCache(ctx), m, fav, pos), nil
|
||||
}
|
||||
|
||||
// LatestItems 最近添加,全库或指定库。远程媒体库(parentID 带前缀)直接透传远程。
|
||||
@@ -177,7 +178,7 @@ func (e *EmbyService) latestSeriesItemsForLibrary(ctx context.Context, userID, l
|
||||
if err := q.Order(mediaReleaseOrderSQL(true)).Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups := e.seriesGroupsFromMedia(rows)
|
||||
groups := e.seriesGroupsFromMedia(ctx, rows)
|
||||
sortSeriesGroups(groups, ItemsParams{SortBy: "premieredate", SortOrder: "Descending"})
|
||||
if len(groups) > limit {
|
||||
groups = groups[:limit]
|
||||
@@ -337,10 +338,12 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st
|
||||
return map[string]any{"Items": []any{}, "TotalRecordCount": int64(0), "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
|
||||
// 历史记录限行:此前无上限全量加载,远程条目多时既拖慢 SQL 也放大
|
||||
// 下面的远程详情请求量。
|
||||
var hist []model.PlaybackHistory
|
||||
if err := e.repo.DB.WithContext(ctx).
|
||||
Where("user_id = ? AND completed = ? AND position_ms > 0", p.UserID, false).
|
||||
Order("watched_at desc").Find(&hist).Error; err != nil {
|
||||
Order("watched_at desc").Limit(200).Find(&hist).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(hist) == 0 {
|
||||
@@ -366,18 +369,31 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st
|
||||
}
|
||||
}
|
||||
|
||||
items := make([]map[string]any, 0, len(hist))
|
||||
// 分页前置:凑满 StartIndex+Limit 条即停,不再为「总数」逐条发远程
|
||||
// 详情 GET(此前每条远程记录一次串行 GET,远程慢时请求挂起数分钟)。
|
||||
// 总数用候选行数(本地过滤后 + 远程候选),对继续观看行的翻页语义
|
||||
// 足够准确。
|
||||
needed := p.StartIndex + p.Limit
|
||||
items := make([]map[string]any, 0, p.Limit)
|
||||
localTotal, remoteTotal := 0, 0
|
||||
for _, h := range hist {
|
||||
if m, ok := byID[h.MediaID]; ok {
|
||||
if p.ParentID != "" && m.LibraryID != p.ParentID && m.SeriesID != p.ParentID {
|
||||
continue
|
||||
}
|
||||
items = append(items, e.itemPayload(ctx, m, false, h.PositionMs))
|
||||
localTotal++
|
||||
if produced := len(items); produced < needed {
|
||||
items = append(items, e.itemPayload(ctx, m, false, h.PositionMs))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if e.remote == nil || !IsEmbyRemoteID(h.MediaID) {
|
||||
continue
|
||||
}
|
||||
remoteTotal++
|
||||
if len(items) >= needed {
|
||||
continue
|
||||
}
|
||||
mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID)
|
||||
mount, acct, err := e.remote.ResolveMount(ctx, mountID)
|
||||
if err != nil || mount == nil || acct == nil {
|
||||
@@ -398,7 +414,7 @@ func (e *EmbyService) resumableItems(ctx context.Context, p ItemsParams) (map[st
|
||||
items = append(items, item)
|
||||
}
|
||||
|
||||
total := int64(len(items))
|
||||
total := int64(localTotal + remoteTotal)
|
||||
if p.StartIndex >= len(items) {
|
||||
return map[string]any{"Items": []map[string]any{}, "TotalRecordCount": total, "StartIndex": p.StartIndex}, nil
|
||||
}
|
||||
@@ -415,9 +431,9 @@ func (e *EmbyService) itemPayload(ctx context.Context, m *model.Media, fav bool,
|
||||
seasonID := ""
|
||||
if e.mediaShouldBeEpisode(ctx, m) {
|
||||
itemType = "Episode"
|
||||
seriesID = e.seriesIDForMedia(m)
|
||||
seriesName = e.seriesNameForMedia(m)
|
||||
seasonID = e.seasonIDForMedia(m)
|
||||
seriesID = e.seriesIDForMedia(ctx, m)
|
||||
seriesName = e.seriesNameForMedia(ctx, m)
|
||||
seasonID = e.seasonIDForMedia(ctx, m)
|
||||
parentID = seasonID
|
||||
episodeTitle := strings.TrimSpace(m.EpisodeTitle)
|
||||
if episodeTitle != "" {
|
||||
|
||||
@@ -63,7 +63,7 @@ func primarySupportedEmbySort(sortBy string, resumeFilter bool) string {
|
||||
for _, part := range strings.Split(sortBy, ",") {
|
||||
key := strings.ToLower(strings.TrimSpace(part))
|
||||
switch key {
|
||||
case "sortname", "name", "premieredate", "productionyear", "datecreated", "communityrating":
|
||||
case "sortname", "name", "premieredate", "productionyear", "datecreated", "datelastmediaadded", "datelastcontentadded", "communityrating":
|
||||
return key
|
||||
case "dateplayed":
|
||||
if resumeFilter {
|
||||
|
||||
@@ -75,9 +75,9 @@ func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string
|
||||
orderIncludesDirection = false
|
||||
case "premieredate", "productionyear":
|
||||
order = mediaReleaseOrderSQL(desc)
|
||||
case "datecreated":
|
||||
order = "media.created_at"
|
||||
orderIncludesDirection = false
|
||||
case "datecreated", "datelastmediaadded", "datelastcontentadded":
|
||||
order = "media.created_at"
|
||||
orderIncludesDirection = false
|
||||
case "dateplayed":
|
||||
order = "resume.watched_at"
|
||||
orderIncludesDirection = false
|
||||
@@ -149,6 +149,9 @@ func (e *EmbyService) episodeItems(ctx context.Context, rows []model.Media, p It
|
||||
}
|
||||
|
||||
func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media, userID string) ([]map[string]any, error) {
|
||||
// 请求级缓存:库类型与 series 标题整页只查一次,消除逐条目 N+1。
|
||||
ctx = e.withPayloadCache(ctx)
|
||||
e.prefetchPayloadCache(ctx, rows)
|
||||
rows = e.collapseMediaVersionRows(ctx, rows)
|
||||
userFavs := map[string]bool{}
|
||||
userPos := map[string]int64{}
|
||||
@@ -240,7 +243,7 @@ func (e *EmbyService) seriesItemsForLibrary(ctx context.Context, libraryID strin
|
||||
if err := q.Order(mediaReleaseOrderSQL(true)).Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups := e.seriesGroupsFromMedia(rows)
|
||||
groups := e.seriesGroupsFromMedia(ctx, rows)
|
||||
sortSeriesGroups(groups, p)
|
||||
total := len(groups)
|
||||
items := make([]map[string]any, 0, minInt(p.Limit, len(groups)))
|
||||
|
||||
@@ -65,7 +65,7 @@ func (e *EmbyService) movieLibraryItems(ctx context.Context, p ItemsParams) (map
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
seriesGroups := e.seriesGroupsFromMedia(episodicRows)
|
||||
seriesGroups := e.seriesGroupsFromMedia(ctx, episodicRows)
|
||||
|
||||
// 真正的电影 -> Movie 项(剔除剧集结构行)。
|
||||
movieQ := apply(e.repo.DB.WithContext(ctx).Model(&model.Media{}))
|
||||
@@ -135,10 +135,11 @@ func (e *EmbyService) libraryIsEpisodic(ctx context.Context, libraryID string) (
|
||||
if strings.TrimSpace(libraryID) == "" {
|
||||
return false, nil
|
||||
}
|
||||
if lib, err := e.repo.Library.FindByID(ctx, libraryID); err != nil {
|
||||
// 走请求级缓存(若有),避免同一请求内对同一库重复查表。
|
||||
if typ, ok, err := e.payloadLibraryType(ctx, libraryID); err != nil {
|
||||
return false, err
|
||||
} else if lib != nil {
|
||||
return embyLibraryTypeIsEpisodic(lib.Type), nil
|
||||
} else if ok {
|
||||
return embyLibraryTypeIsEpisodic(typ), nil
|
||||
}
|
||||
var count int64
|
||||
err := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
|
||||
@@ -151,11 +152,11 @@ func (e *EmbyService) mediaBelongsToEpisodicLibrary(ctx context.Context, m *mode
|
||||
if e == nil || m == nil || strings.TrimSpace(m.LibraryID) == "" {
|
||||
return false
|
||||
}
|
||||
lib, err := e.repo.Library.FindByID(ctx, m.LibraryID)
|
||||
if err != nil || lib == nil {
|
||||
typ, ok, err := e.payloadLibraryType(ctx, m.LibraryID)
|
||||
if err != nil || !ok {
|
||||
return false
|
||||
}
|
||||
return embyLibraryTypeIsEpisodic(lib.Type)
|
||||
return embyLibraryTypeIsEpisodic(typ)
|
||||
}
|
||||
|
||||
func (e *EmbyService) mediaShouldBeEpisode(ctx context.Context, m *model.Media) bool {
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
// 请求级 payload 构建缓存:/Items 列表为每行构建 payload 时,
|
||||
// mediaShouldBeEpisode 需要库类型、剧集 payload 需要 series 标题。
|
||||
// 一次页面请求内这些值高度重复(同一库、同一部剧),挂在 ctx 上的
|
||||
// 小缓存可以把每条目 2-3 次 DB 查询降为整个请求各 1 次预取。
|
||||
|
||||
type embyPayloadCacheKey struct{}
|
||||
|
||||
type embyLibraryTypeEntry struct {
|
||||
typ string
|
||||
found bool // 库不存在时 found=false,调用方可退回计数启发式
|
||||
}
|
||||
|
||||
type embyPayloadCache struct {
|
||||
mu sync.Mutex
|
||||
libTypes map[string]embyLibraryTypeEntry
|
||||
series map[string]string // series_id -> title("" 表示不存在/无标题)
|
||||
}
|
||||
|
||||
func (c *embyPayloadCache) libraryType(id string) (embyLibraryTypeEntry, bool) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
entry, ok := c.libTypes[id]
|
||||
return entry, ok
|
||||
}
|
||||
|
||||
func (c *embyPayloadCache) setLibraryType(id string, entry embyLibraryTypeEntry) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.libTypes[id] = entry
|
||||
}
|
||||
|
||||
func (c *embyPayloadCache) seriesTitle(id string) (string, bool) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
title, ok := c.series[id]
|
||||
return title, ok
|
||||
}
|
||||
|
||||
func (c *embyPayloadCache) setSeriesTitle(id, title string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.series[id] = title
|
||||
}
|
||||
|
||||
// withPayloadCache attaches a fresh request-scoped cache if none exists yet.
|
||||
func (e *EmbyService) withPayloadCache(ctx context.Context) context.Context {
|
||||
if e == nil || e.repo == nil {
|
||||
return ctx
|
||||
}
|
||||
if ctx.Value(embyPayloadCacheKey{}) != nil {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, embyPayloadCacheKey{}, &embyPayloadCache{
|
||||
libTypes: map[string]embyLibraryTypeEntry{},
|
||||
series: map[string]string{},
|
||||
})
|
||||
}
|
||||
|
||||
// prefetchPayloadCache warms the cache for the given media rows with two bulk
|
||||
// queries (library types, series titles) instead of per-item lookups.
|
||||
func (e *EmbyService) prefetchPayloadCache(ctx context.Context, rows []model.Media) {
|
||||
cache, ok := ctx.Value(embyPayloadCacheKey{}).(*embyPayloadCache)
|
||||
if !ok || len(rows) == 0 {
|
||||
return
|
||||
}
|
||||
libIDs := make([]string, 0, 8)
|
||||
seriesIDs := make([]string, 0, 8)
|
||||
seenLib := map[string]struct{}{}
|
||||
seenSeries := map[string]struct{}{}
|
||||
for i := range rows {
|
||||
row := &rows[i]
|
||||
if id := strings.TrimSpace(row.LibraryID); id != "" {
|
||||
if _, done := seenLib[id]; !done {
|
||||
// 已在缓存中的库不必再查。
|
||||
if _, hit := cache.libraryType(id); !hit {
|
||||
seenLib[id] = struct{}{}
|
||||
libIDs = append(libIDs, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
if id := strings.TrimSpace(row.SeriesID); id != "" {
|
||||
if _, done := seenSeries[id]; !done {
|
||||
if _, hit := cache.seriesTitle(id); !hit {
|
||||
seenSeries[id] = struct{}{}
|
||||
seriesIDs = append(seriesIDs, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(libIDs) > 0 {
|
||||
var libs []model.Library
|
||||
if err := e.repo.DB.WithContext(ctx).Select("id, type").Where("id IN ?", libIDs).Find(&libs).Error; err == nil {
|
||||
found := map[string]string{}
|
||||
for _, lib := range libs {
|
||||
found[lib.ID] = lib.Type
|
||||
}
|
||||
for _, id := range libIDs {
|
||||
typ, ok := found[id]
|
||||
cache.setLibraryType(id, embyLibraryTypeEntry{typ: typ, found: ok})
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(seriesIDs) > 0 {
|
||||
var series []model.Series
|
||||
if err := e.repo.DB.WithContext(ctx).Select("id, title").Where("id IN ?", seriesIDs).Find(&series).Error; err == nil {
|
||||
for _, s := range series {
|
||||
cache.setSeriesTitle(s.ID, s.Title)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// payloadLibraryType resolves a library type through the request cache,
|
||||
// falling back to a direct lookup when no cache is attached. found=false
|
||||
// means the library row does not exist (soft-deleted or orphaned id).
|
||||
func (e *EmbyService) payloadLibraryType(ctx context.Context, libraryID string) (typ string, found bool, err error) {
|
||||
if cache, ok := ctx.Value(embyPayloadCacheKey{}).(*embyPayloadCache); ok {
|
||||
if entry, hit := cache.libraryType(libraryID); hit {
|
||||
return entry.typ, entry.found, nil
|
||||
}
|
||||
var lib model.Library
|
||||
if dbErr := e.repo.DB.WithContext(ctx).Select("id, type").Where("id = ?", libraryID).First(&lib).Error; dbErr != nil {
|
||||
cache.setLibraryType(libraryID, embyLibraryTypeEntry{})
|
||||
return "", false, nil
|
||||
}
|
||||
cache.setLibraryType(lib.ID, embyLibraryTypeEntry{typ: lib.Type, found: true})
|
||||
return lib.Type, true, nil
|
||||
}
|
||||
var lib model.Library
|
||||
if err = e.repo.DB.WithContext(ctx).Select("id, type").Where("id = ?", libraryID).First(&lib).Error; err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return "", false, nil
|
||||
}
|
||||
return "", false, err
|
||||
}
|
||||
return lib.Type, true, nil
|
||||
}
|
||||
|
||||
// payloadSeriesTitle resolves a series title through the request cache,
|
||||
// falling back to a direct lookup when no cache is attached.
|
||||
func (e *EmbyService) payloadSeriesTitle(ctx context.Context, seriesID string) (string, bool, error) {
|
||||
if cache, ok := ctx.Value(embyPayloadCacheKey{}).(*embyPayloadCache); ok {
|
||||
if title, hit := cache.seriesTitle(seriesID); hit {
|
||||
return title, true, nil
|
||||
}
|
||||
var s model.Series
|
||||
if err := e.repo.DB.WithContext(ctx).Select("id, title").Where("id = ?", seriesID).First(&s).Error; err != nil {
|
||||
cache.setSeriesTitle(seriesID, "")
|
||||
return "", true, nil
|
||||
}
|
||||
cache.setSeriesTitle(s.ID, s.Title)
|
||||
return s.Title, true, nil
|
||||
}
|
||||
series, err := e.repo.Series.FindByID(ctx, seriesID)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
if series == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
return series.Title, true, nil
|
||||
}
|
||||
+124
-24
@@ -26,6 +26,7 @@ import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
@@ -73,6 +74,7 @@ type EmbyRemoteService struct {
|
||||
repo *repository.Container
|
||||
crypto *CryptoService
|
||||
http *http.Client
|
||||
stream *http.Client // 流式代理专用(视频/字幕),无整体 Timeout
|
||||
cache *RuntimeCacheService
|
||||
}
|
||||
|
||||
@@ -87,6 +89,12 @@ func NewEmbyRemoteService(cfg *config.Config, log *zap.Logger, repo *repository.
|
||||
Timeout: embyRemoteHTTPTimeout,
|
||||
Transport: &embyRemoteTransport{base: http.DefaultTransport},
|
||||
},
|
||||
// 流式代理必须用无整体 Timeout 的 client:http.Client.Timeout
|
||||
// 覆盖整个响应体读取过程,15s 的常规超时会让代理播放播到
|
||||
// 15 秒整被掐断。生命周期由请求 ctx 控制。
|
||||
stream: &http.Client{
|
||||
Transport: &embyRemoteTransport{base: http.DefaultTransport},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -316,6 +324,65 @@ func (r *EmbyRemoteService) AutoSeedMounts(ctx context.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
// remoteConfigWithToken 解密账号配置并确保已有可用凭据(首次请求自动认证并
|
||||
// 回写 token 与 remote_user_id,等价于管理端「测试连接」),保证后续构造的
|
||||
// /Users/{userId} 路径使用远程真实用户 GUID,而不是未认证兜底的 "0"。
|
||||
func (r *EmbyRemoteService) remoteConfigWithToken(ctx context.Context, acct *model.StrmAccount) (*EmbyRemoteConfig, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.ensureToken(ctx, acct, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// api_key 直连(未配用户名/密码)的账号认证步骤不会回填用户 ID;此时用
|
||||
// api_key 拉一次用户列表取真实用户 ID 并回写,避免 /Users/{uid} 请求路径
|
||||
// 落回兜底 "0" 被远程 Emby 拒绝(Unrecognized Guid format)。
|
||||
if strings.TrimSpace(cfg.RemoteUserID) == "" {
|
||||
r.resolveRemoteUserID(ctx, acct, cfg)
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// resolveRemoteUserID 用已有 api_key 拉远程用户列表,把首个用户 ID 回写账号
|
||||
// 配置(取不到时静默跳过,保持兜底行为不变)。
|
||||
func (r *EmbyRemoteService) resolveRemoteUserID(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) {
|
||||
if acct == nil || cfg == nil || strings.TrimSpace(cfg.Token) == "" || strings.TrimSpace(cfg.RemoteUserID) != "" {
|
||||
return
|
||||
}
|
||||
q := url.Values{"api_key": {cfg.Token}}
|
||||
var users []map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users", q, &users); err != nil || len(users) == 0 {
|
||||
return
|
||||
}
|
||||
uid := strings.TrimSpace(remoteItemString(users[0], "Id"))
|
||||
if uid == "" {
|
||||
return
|
||||
}
|
||||
cfg.RemoteUserID = uid
|
||||
_ = r.updateAccountConfig(ctx, acct, func(raw map[string]string) {
|
||||
raw["remote_user_id"] = uid
|
||||
})
|
||||
}
|
||||
|
||||
// CleanupOrphanMounts 清理账号已删除的残留挂载(老版本删除账号未级联),
|
||||
// 避免挂载计数/列表出现永远清不掉的孤儿数据。
|
||||
func (r *EmbyRemoteService) CleanupOrphanMounts(ctx context.Context) {
|
||||
n, err := r.repo.EmbyMount.DeleteOrphans(ctx)
|
||||
if err != nil {
|
||||
if r.log != nil {
|
||||
r.log.Warn("cleanup orphan emby mounts failed", zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
if n > 0 {
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
if r.log != nil {
|
||||
r.log.Info("cleaned up orphan emby mounts", zap.Int64("mounts", n))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// configOf 解密账号配置。
|
||||
func (r *EmbyRemoteService) configOf(acct *model.StrmAccount) (*EmbyRemoteConfig, error) {
|
||||
raw := map[string]string{}
|
||||
@@ -430,19 +497,28 @@ func (r *EmbyRemoteService) ensureTokenOnLine(ctx context.Context, acct *model.S
|
||||
return nil
|
||||
}
|
||||
|
||||
// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。
|
||||
func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error {
|
||||
if acct == nil {
|
||||
// acctCfgMu 序列化对账号 Config 的读-改-写。并发请求若各自基于请求开始
|
||||
// 时的快照做整包覆盖,会互相丢失更新(刚持久化的 token / active_line 被
|
||||
// 旧快照覆盖回去)。
|
||||
var acctCfgMu sync.Mutex
|
||||
|
||||
// updateAccountConfig 在互斥下重读账号最新 Config,应用 mutate 后写回,
|
||||
// 并同步调用方持有的 acct 快照。
|
||||
func (r *EmbyRemoteService) updateAccountConfig(ctx context.Context, acct *model.StrmAccount, mutate func(raw map[string]string)) error {
|
||||
if acct == nil || r.repo == nil {
|
||||
return nil
|
||||
}
|
||||
acctCfgMu.Lock()
|
||||
defer acctCfgMu.Unlock()
|
||||
raw := map[string]string{}
|
||||
if fresh, err := r.repo.StrmAccount.FindByID(ctx, acct.ID); err == nil && fresh != nil {
|
||||
acct.Config = fresh.Config // 以 DB 最新值为基线,避免覆盖并发写入
|
||||
}
|
||||
if strings.TrimSpace(acct.Config) != "" {
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
}
|
||||
raw["api_key"] = r.crypto.Encrypt(cfg.Token)
|
||||
raw["remote_user_id"] = cfg.RemoteUserID
|
||||
if strings.TrimSpace(raw["username"]) == "" {
|
||||
raw["username"] = cfg.Username
|
||||
if mutate != nil {
|
||||
mutate(raw)
|
||||
}
|
||||
data, err := json.Marshal(raw)
|
||||
if err != nil {
|
||||
@@ -452,6 +528,20 @@ func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAc
|
||||
return r.repo.StrmAccount.Update(ctx, acct)
|
||||
}
|
||||
|
||||
// persistToken 把认证得到的 token / user id 加密写回账号配置(下次请求免登录)。
|
||||
func (r *EmbyRemoteService) persistToken(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig) error {
|
||||
if acct == nil {
|
||||
return nil
|
||||
}
|
||||
return r.updateAccountConfig(ctx, acct, func(raw map[string]string) {
|
||||
raw["api_key"] = r.crypto.Encrypt(cfg.Token)
|
||||
raw["remote_user_id"] = cfg.RemoteUserID
|
||||
if strings.TrimSpace(raw["username"]) == "" {
|
||||
raw["username"] = cfg.Username
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// doGet 向远程 Emby 发起带 api_key 的 GET,把响应 JSON 解码到 out。
|
||||
// 401 时自动重认证一次再重试(凭据过期场景)。连接失败时按线路优先级自动切换。
|
||||
func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, path string, q url.Values, out any) error {
|
||||
@@ -469,6 +559,10 @@ func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount,
|
||||
}
|
||||
}
|
||||
if lastErr != nil {
|
||||
if r.log != nil && acct != nil {
|
||||
r.log.Warn("remote emby request failed",
|
||||
zap.String("account", acct.Name), zap.String("path", path), zap.Error(lastErr))
|
||||
}
|
||||
return lastErr
|
||||
}
|
||||
return errors.New("远程 Emby 请求失败")
|
||||
@@ -498,22 +592,28 @@ func (r *EmbyRemoteService) doGetOnLine(ctx context.Context, acct *model.StrmAcc
|
||||
if err != nil {
|
||||
return fmt.Errorf("请求远程 Emby 失败: %w", err)
|
||||
}
|
||||
data, readErr := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
||||
// 读 8MB+1 以区分"刚好 8MB"与"被截断":截断的 JSON 会让
|
||||
// Unmarshal 报 unexpected end,难以定位;这里显式报错。
|
||||
data, readErr := io.ReadAll(io.LimitReader(resp.Body, (8<<20)+1))
|
||||
resp.Body.Close()
|
||||
if readErr != nil {
|
||||
return readErr
|
||||
}
|
||||
if len(data) > 8<<20 {
|
||||
return fmt.Errorf("远程 Emby 响应超过 8MB 上限(路径 %s):请减小分页或 Fields 字段", path)
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnauthorized && attempt == 0 {
|
||||
// 401:只清当前线路的内存 token 并立即重认证;不在此时删除
|
||||
// DB 里的 api_key——①外层还会按线路故障转移(其他线路可能
|
||||
// 存有自己的 token);②纯 api_key 账号删除后无法再认证,一次
|
||||
// 线路误报就会把账号“砖化”。重认证成功后 persistToken 会用
|
||||
// 新 token 覆盖 api_key。
|
||||
cfg.Token = ""
|
||||
master.Token = ""
|
||||
if acct != nil {
|
||||
raw := map[string]string{}
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
delete(raw, "api_key")
|
||||
enc, _ := json.Marshal(raw)
|
||||
acct.Config = string(enc)
|
||||
_ = r.repo.StrmAccount.Update(ctx, acct)
|
||||
if err := r.ensureTokenOnLine(ctx, acct, cfg); err != nil {
|
||||
return fmt.Errorf("认证重试失败: %w", err)
|
||||
}
|
||||
master.Token = cfg.Token
|
||||
master.RemoteUserID = cfg.RemoteUserID
|
||||
continue
|
||||
}
|
||||
if resp.StatusCode >= 300 {
|
||||
@@ -553,7 +653,7 @@ func (r *EmbyRemoteService) ProxyPlayOf(acct *model.StrmAccount) (bool, error) {
|
||||
|
||||
// RemoteViews 拉取远程媒体库(View)列表,返回远程原始 view map(未重写)。
|
||||
func (r *EmbyRemoteService) RemoteViews(ctx context.Context, acct *model.StrmAccount) ([]map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -588,7 +688,7 @@ func (r *EmbyRemoteService) remoteUserID(cfg *EmbyRemoteConfig) string {
|
||||
// RemoteItems 向远程 Emby 转发 /Items 浏览/搜索请求,返回重写后的响应载荷。
|
||||
// p 的分页/排序/过滤参数原样转发,分页语义完全由远程承接。
|
||||
func (r *EmbyRemoteService) RemoteItems(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, p ItemsParams) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -633,7 +733,7 @@ func (r *EmbyRemoteService) RemoteItems(ctx context.Context, mount *model.EmbyMo
|
||||
// RemoteSearchMount 对单个挂载的媒体库执行全局搜索(ParentId=挂载的远程库,
|
||||
// Recursive 返回库内全部命中),结果归属明确可直接伪装。
|
||||
func (r *EmbyRemoteService) RemoteSearchMount(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, p ItemsParams) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -666,7 +766,7 @@ func (r *EmbyRemoteService) RemoteSearchMount(ctx context.Context, mount *model.
|
||||
|
||||
// RemoteItem 拉取远程单条目详情(含响应的重写)。
|
||||
func (r *EmbyRemoteService) RemoteItem(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -681,7 +781,7 @@ func (r *EmbyRemoteService) RemoteItem(ctx context.Context, mount *model.EmbyMou
|
||||
|
||||
// RemoteLatest 拉取远程「最近添加」(用于 /Items/Latest 聚合)。
|
||||
func (r *EmbyRemoteService) RemoteLatest(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, parentID string, limit int) ([]map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -702,7 +802,7 @@ func (r *EmbyRemoteService) RemoteLatest(ctx context.Context, mount *model.EmbyM
|
||||
// 播放 URL:不代理=指向远程绝对地址(播放字节不过 MeBox);代理=指向 MeBox
|
||||
// 本地 /Videos/{encodedID} 端点(由 ProxyVideoStream 反代)。
|
||||
func (r *EmbyRemoteService) RemotePlaybackInfo(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID, userID string) (map[string]any, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -888,7 +988,7 @@ func (r *EmbyRemoteService) proxyVideoStreamOnLine(ctx context.Context, w http.R
|
||||
if rangeHeader := req.Header.Get("Range"); rangeHeader != "" {
|
||||
upstream.Header.Set("Range", rangeHeader)
|
||||
}
|
||||
resp, err := r.http.Do(upstream)
|
||||
resp, err := r.stream.Do(upstream)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接远程 Emby 视频流失败: %w", err)
|
||||
}
|
||||
@@ -959,7 +1059,7 @@ func (r *EmbyRemoteService) proxySubtitleOnLine(ctx context.Context, w http.Resp
|
||||
return err
|
||||
}
|
||||
upstream.Header.Set("X-Emby-Token", cfg.Token)
|
||||
resp, err := r.http.Do(upstream)
|
||||
resp, err := r.stream.Do(upstream)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接远程 Emby 字幕流失败: %w", err)
|
||||
}
|
||||
|
||||
@@ -124,10 +124,11 @@ func isEmbyLineFailoverError(err error) bool {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
// 注意:认证类错误(重认证失败 / 缺少凭据)不在此排除——401 后清空
|
||||
// 内存 token 重认证失败时应继续按线路故障转移,其他线路可能存有
|
||||
// 自己的 token。仅“登录失败”(密码错误)是账号级问题,无需换线。
|
||||
if strings.Contains(msg, "登录失败") ||
|
||||
strings.Contains(msg, "未返回 accesstoken") ||
|
||||
strings.Contains(msg, "缺少 emby 凭据") ||
|
||||
strings.Contains(msg, "认证重试失败") {
|
||||
strings.Contains(msg, "未返回 accesstoken") {
|
||||
return false
|
||||
}
|
||||
var urlErr *url.Error
|
||||
@@ -147,20 +148,13 @@ func (r *EmbyRemoteService) persistActiveLine(ctx context.Context, acct *model.S
|
||||
if acct == nil || cfg == nil || lineIndex < 0 || lineIndex >= len(cfg.Lines) {
|
||||
return nil
|
||||
}
|
||||
raw := map[string]string{}
|
||||
if strings.TrimSpace(acct.Config) != "" {
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
}
|
||||
raw["active_line"] = strconv.Itoa(lineIndex)
|
||||
raw["url"] = cfg.Lines[lineIndex].URL
|
||||
data, err := json.Marshal(raw)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
acct.Config = string(data)
|
||||
err := r.updateAccountConfig(ctx, acct, func(raw map[string]string) {
|
||||
raw["active_line"] = strconv.Itoa(lineIndex)
|
||||
raw["url"] = cfg.Lines[lineIndex].URL
|
||||
})
|
||||
cfg.ActiveLine = lineIndex
|
||||
cfg.BaseURL = normalizeEmbyRemoteURL(cfg.Lines[lineIndex].URL)
|
||||
return r.repo.StrmAccount.Update(ctx, acct)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) adoptWorkingLine(ctx context.Context, acct *model.StrmAccount, cfg *EmbyRemoteConfig, lineIndex int) {
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
@@ -57,7 +58,7 @@ func (r *EmbyRemoteService) RemoteLibraries(ctx context.Context) ([]RemoteLibrar
|
||||
acctData[m.AccountID] = nil
|
||||
continue
|
||||
}
|
||||
cfg, cfgErr := r.configOf(acct)
|
||||
cfg, cfgErr := r.remoteConfigWithToken(ctx, acct)
|
||||
if cfgErr != nil {
|
||||
acctData[m.AccountID] = nil
|
||||
continue
|
||||
@@ -198,6 +199,9 @@ func (r *EmbyRemoteService) MapRemoteItemToMedia(ctx context.Context, mount *mod
|
||||
media.CreatedAt = date
|
||||
media.UpdatedAt = date
|
||||
}
|
||||
if date, ok := parseEmbyRemoteDate(remoteItemString(item, "DateLastMediaAdded")); ok {
|
||||
media.UpdatedAt = date
|
||||
}
|
||||
// 只有远程明确存在图片标签才下发图片 URL。
|
||||
if remoteItemHasImageTag(item, "Primary") {
|
||||
media.PosterURL = r.remoteItemImageURL(cfg, remoteID, "Primary")
|
||||
@@ -301,28 +305,28 @@ func (r *EmbyRemoteService) MapRemoteItemToMedia(ctx context.Context, mount *mod
|
||||
media.SeasonNum = 0
|
||||
media.EpisodeNum = 0
|
||||
}
|
||||
if mount != nil && strings.TrimSpace(mount.RemoteViewID) != "" {
|
||||
libID := EncodeEmbyRemoteID(mount.ID, mount.RemoteViewID)
|
||||
media.DisplayLibraryID = libID
|
||||
media.LibraryID = libID
|
||||
libName := strings.TrimSpace(mount.Name)
|
||||
if libName == "" {
|
||||
libName = strings.TrimSpace(mount.RemoteViewName)
|
||||
}
|
||||
if libName == "" && acct != nil {
|
||||
libName = acct.Name
|
||||
} else if acct != nil && acct.Name != "" && !strings.Contains(libName, acct.Name) {
|
||||
libName = acct.Name + " · " + libName
|
||||
}
|
||||
media.LibraryName = libName
|
||||
media.DisplayLibraryName = libName
|
||||
if mount != nil && strings.TrimSpace(mount.RemoteViewID) != "" {
|
||||
libID := EncodeEmbyRemoteID(mount.ID, mount.RemoteViewID)
|
||||
media.DisplayLibraryID = libID
|
||||
media.LibraryID = libID
|
||||
libName := strings.TrimSpace(mount.Name)
|
||||
if libName == "" {
|
||||
libName = strings.TrimSpace(mount.RemoteViewName)
|
||||
}
|
||||
if libName == "" && acct != nil {
|
||||
libName = acct.Name
|
||||
} else if acct != nil && acct.Name != "" && !strings.Contains(libName, acct.Name) {
|
||||
libName = acct.Name + " · " + libName
|
||||
}
|
||||
media.LibraryName = libName
|
||||
media.DisplayLibraryName = libName
|
||||
}
|
||||
return media
|
||||
}
|
||||
|
||||
// RemoteLibraryMedia 拉远程库直属条目(电影库=Movie,剧集库=Series),映射分页。
|
||||
func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string, itemTypes string, offset, limit int) ([]model.Media, int64, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
@@ -367,37 +371,52 @@ func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model
|
||||
|
||||
// RemoteMediaDetail 拉远程单条目映射为 Media(网页详情页)。
|
||||
func (r *EmbyRemoteService) RemoteMediaDetail(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
m, _, err := r.remoteMediaDetailRaw(ctx, mount, acct, remoteID)
|
||||
return m, err
|
||||
}
|
||||
|
||||
// remoteMediaDetailRaw 拉取远程条目详情,同时返回原始载荷(ID 已伪装),
|
||||
// 供调用方免二次请求读取 Type / SeriesId 等字段。
|
||||
func (r *EmbyRemoteService) remoteMediaDetailRaw(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) (*model.Media, map[string]any, error) {
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID)
|
||||
path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating"
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
RewriteEmbyRemoteIDs(out, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, out)
|
||||
return &m, nil
|
||||
return &m, out, nil
|
||||
}
|
||||
|
||||
// RemoteEpisodes 拉远程条目下的集列表(Series/Season/Folder→子集;Episode→同系列;
|
||||
// Movie→自身单条),按季/集排序,与本地 ListMediaEpisodes 行为一致。
|
||||
func (r *EmbyRemoteService) RemoteEpisodes(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteID string) ([]model.Media, error) {
|
||||
detail, err := r.RemoteMediaDetail(ctx, mount, acct, remoteID)
|
||||
detail, rawDetail, err := r.remoteMediaDetailRaw(ctx, mount, acct, remoteID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 用远程详情载荷精判类型(Episode→同系列;Series/Season/Folder→子集;Movie→单条)。
|
||||
itemType := r.remoteItemType(ctx, acct, remoteID)
|
||||
// Type/SeriesId 都在详情载荷里现成可用,不再为判定类型/系列额外发起
|
||||
// 两次重复的远程全量 GET(远程慢时页面延迟直接×3)。
|
||||
itemType := remoteItemString(rawDetail, "Type")
|
||||
if itemType == "" {
|
||||
itemType = remoteItemTypeOf(detail)
|
||||
}
|
||||
var parentID string
|
||||
switch itemType {
|
||||
case "Episode":
|
||||
parentID = r.remoteItemSeriesID(ctx, acct, remoteID)
|
||||
parentID = remoteItemString(rawDetail, "SeriesId")
|
||||
if _, rid, ok := DecodeEmbyRemoteID(parentID); ok {
|
||||
parentID = rid // 载荷 ID 已伪装,远程查询需要原始 ID
|
||||
}
|
||||
if parentID == "" {
|
||||
parentID = r.remoteItemSeriesID(ctx, acct, remoteID)
|
||||
}
|
||||
if parentID == "" {
|
||||
parentID = remoteID
|
||||
}
|
||||
@@ -423,7 +442,7 @@ func (r *EmbyRemoteService) RemoteEpisodes(ctx context.Context, mount *model.Emb
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, parentID string) ([]model.Media, int64, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
@@ -431,28 +450,47 @@ func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.E
|
||||
q.Set("ParentId", parentID)
|
||||
q.Set("IncludeItemTypes", "Episode")
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("StartIndex", "0")
|
||||
q.Set("Limit", "500")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
items := make([]model.Media, 0, 64)
|
||||
total := int64(0)
|
||||
// 每页 200 循环拉全:MediaStreams/MediaSources 重字段下单页 500 条
|
||||
// 已贴近 8MB 截断上限;单次大页超限会静默解析失败。
|
||||
const episodePageSize = 200
|
||||
for startIndex := 0; ; startIndex += episodePageSize {
|
||||
q.Set("StartIndex", strconv.Itoa(startIndex))
|
||||
q.Set("Limit", strconv.Itoa(episodePageSize))
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
total = body.TotalRecordCount
|
||||
if len(body.Items) == 0 {
|
||||
break
|
||||
}
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
items = append(items, m)
|
||||
}
|
||||
if len(body.Items) < episodePageSize {
|
||||
break
|
||||
}
|
||||
}
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
items := make([]model.Media, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
items = append(items, m)
|
||||
}
|
||||
return items, body.TotalRecordCount, nil
|
||||
return items, total, nil
|
||||
}
|
||||
|
||||
// RemoteSeriesCards 远程剧集库的系列卡片(ChildCount 作为集数)。
|
||||
//
|
||||
// 远程 Emby 的 Series DTO 不会返回 DateLastMediaAdded 字段(即使请求 Fields
|
||||
// 也缺失),但其服务端排序支持 SortBy=DateLastContentAdded——即客户端"上次
|
||||
// 添加集日期"排序。因此这里直接按该键倒序分页拉全量,返回的卡片顺序与对方
|
||||
// Emby 客户端选择"上次添加集日期"完全一致;LastAddedAt 在远程提供字段时
|
||||
// 才填充,否则保持 nil(前端对无该值的卡片维持服务器顺序,不再回退加入日期)。
|
||||
func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string) ([]SeriesCard, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -465,28 +503,46 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.
|
||||
q.Set("ParentId", remoteViewID)
|
||||
q.Set("IncludeItemTypes", "Series")
|
||||
q.Set("Recursive", "false")
|
||||
q.Set("StartIndex", "0")
|
||||
q.Set("Limit", "1000")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,RecursiveItemCount,SeriesPrimaryImage,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
q.Set("SortBy", "DateLastContentAdded")
|
||||
q.Set("SortOrder", "Descending")
|
||||
// 每页 200:Fields 带全量重字段(Overview/MediaStreams 等)时单页 1000
|
||||
// 条的载荷会超过 doGet 的 8MB 截断上限,JSON 被静默截断直接解析失败。
|
||||
q.Set("Limit", "200")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,RecursiveItemCount,SeriesPrimaryImage,DateCreated,DateLastMediaAdded,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
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, err
|
||||
}
|
||||
cards := make([]SeriesCard, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
// 集数优先用递归条目数(ChildCount 只算直属 Season 文件夹数)。
|
||||
count := remoteItemInt(it, "RecursiveItemCount")
|
||||
if count == 0 {
|
||||
count = remoteItemInt(it, "ChildCount")
|
||||
cards := make([]SeriesCard, 0)
|
||||
for startIndex := 0; ; startIndex += 200 {
|
||||
q.Set("StartIndex", strconv.Itoa(startIndex))
|
||||
body.Items = nil
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Items", q, &body); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count == 0 {
|
||||
count = 1
|
||||
if len(body.Items) == 0 {
|
||||
break
|
||||
}
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, mount.ID)
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
// 集数优先用递归条目数(ChildCount 只算直属 Season 文件夹数)。
|
||||
count := remoteItemInt(it, "RecursiveItemCount")
|
||||
if count == 0 {
|
||||
count = remoteItemInt(it, "ChildCount")
|
||||
}
|
||||
if count == 0 {
|
||||
count = 1
|
||||
}
|
||||
var lastAdded *time.Time
|
||||
if date, ok := parseEmbyRemoteDate(remoteItemString(it, "DateLastMediaAdded")); ok {
|
||||
lastAdded = &date
|
||||
}
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: count, LastAddedAt: lastAdded})
|
||||
}
|
||||
if int64(len(cards)) >= body.TotalRecordCount || len(body.Items) < 1000 {
|
||||
break
|
||||
}
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: count})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
@@ -496,7 +552,7 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.
|
||||
|
||||
// RemoteLatestCards 远程库最新条目(首页预览卡片),映射 SeriesCard。
|
||||
func (r *EmbyRemoteService) RemoteLatestCards(ctx context.Context, mount *model.EmbyMount, acct *model.StrmAccount, remoteViewID string, limit int) ([]SeriesCard, error) {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -512,13 +568,21 @@ func (r *EmbyRemoteService) RemoteLatestCards(ctx context.Context, mount *model.
|
||||
cards := make([]SeriesCard, 0, len(items))
|
||||
for _, it := range items {
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: 0})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
var lastAdded *time.Time
|
||||
if !m.UpdatedAt.IsZero() {
|
||||
t := m.UpdatedAt
|
||||
lastAdded = &t
|
||||
} else if !m.CreatedAt.IsZero() {
|
||||
t := m.CreatedAt
|
||||
lastAdded = &t
|
||||
}
|
||||
return cards, nil
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: 0, LastAddedAt: lastAdded})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
// RemoteSearchMedia 在全部启用的挂载库中并发搜索影视条目(Movie,Series),
|
||||
// 并将远程结果映射为 model.Media。遵循当前用户的 MediaVisibility 权限规则。
|
||||
@@ -579,7 +643,7 @@ func (r *EmbyRemoteService) RemoteSearchMedia(ctx context.Context, query string,
|
||||
if acct == nil {
|
||||
continue
|
||||
}
|
||||
cfg, cfgErr := r.configOf(acct)
|
||||
cfg, cfgErr := r.remoteConfigWithToken(ctx, acct)
|
||||
if cfgErr != nil {
|
||||
continue
|
||||
}
|
||||
@@ -609,32 +673,34 @@ func (r *EmbyRemoteService) RemoteSearchMedia(ctx context.Context, query string,
|
||||
return
|
||||
}
|
||||
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", target.mount.RemoteViewID)
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("SearchTerm", query)
|
||||
q.Set("IncludeItemTypes", "Movie,Series")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
q.Set("Limit", strconv.Itoa(limit))
|
||||
q.Set("StartIndex", "0")
|
||||
helper.Run(r.log, "emby.remoteSearch", func() {
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", target.mount.RemoteViewID)
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("SearchTerm", query)
|
||||
q.Set("IncludeItemTypes", "Movie,Series")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources,DateCreated,DateLastMediaAdded,PremiereDate,ProductionYear,CommunityRating,CriticRating")
|
||||
q.Set("Limit", strconv.Itoa(limit))
|
||||
q.Set("StartIndex", "0")
|
||||
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
}
|
||||
if err := r.doGet(searchCtx, target.acct, target.cfg, "/Users/"+url.PathEscape(r.remoteUserID(target.cfg))+"/Items", q, &body); err != nil {
|
||||
if r.log != nil {
|
||||
r.log.Warn("remote search failed",
|
||||
zap.String("mount", target.mount.RemoteViewName), zap.Error(err))
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
}
|
||||
return
|
||||
}
|
||||
medias := make([]model.Media, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, target.mount.ID)
|
||||
m := r.MapRemoteItemToMedia(searchCtx, &target.mount, target.acct, target.cfg, it)
|
||||
medias = append(medias, m)
|
||||
}
|
||||
results[idx] = searchResult{items: medias}
|
||||
if err := r.doGet(searchCtx, target.acct, target.cfg, "/Users/"+url.PathEscape(r.remoteUserID(target.cfg))+"/Items", q, &body); err != nil {
|
||||
if r.log != nil {
|
||||
r.log.Warn("remote search failed",
|
||||
zap.String("mount", target.mount.RemoteViewName), zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
medias := make([]model.Media, 0, len(body.Items))
|
||||
for _, it := range body.Items {
|
||||
RewriteEmbyRemoteIDs(it, target.mount.ID)
|
||||
m := r.MapRemoteItemToMedia(searchCtx, &target.mount, target.acct, target.cfg, it)
|
||||
medias = append(medias, m)
|
||||
}
|
||||
results[idx] = searchResult{items: medias}
|
||||
})
|
||||
}(i, t)
|
||||
}
|
||||
wg.Wait()
|
||||
@@ -661,7 +727,7 @@ func (r *EmbyRemoteService) WebStreamURL(ctx context.Context, acct *model.StrmAc
|
||||
|
||||
// remoteItemType 轻量查询远程条目 Type(避免依赖映射载荷)。
|
||||
func (r *EmbyRemoteService) remoteItemType(ctx context.Context, acct *model.StrmAccount, remoteID string) string {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
@@ -674,7 +740,7 @@ func (r *EmbyRemoteService) remoteItemType(ctx context.Context, acct *model.Strm
|
||||
|
||||
// remoteItemSeriesID 轻量查询 Episode 的 SeriesId。
|
||||
func (r *EmbyRemoteService) remoteItemSeriesID(ctx context.Context, acct *model.StrmAccount, remoteID string) string {
|
||||
cfg, err := r.configOf(acct)
|
||||
cfg, err := r.remoteConfigWithToken(ctx, acct)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -68,10 +70,227 @@ func TestMapRemoteItemToMediaCriticRatingFallback(t *testing.T) {
|
||||
if media.Rating != 9.2 {
|
||||
t.Fatalf("Rating = %f, want 9.2 from CriticRating", media.Rating)
|
||||
}
|
||||
if media.Year != 2022 {
|
||||
t.Fatalf("Year = %d, want 2022 from PremiereDate", media.Year)
|
||||
}
|
||||
if media.Year != 2022 {
|
||||
t.Fatalf("Year = %d, want 2022 from PremiereDate", media.Year)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapRemoteItemToMediaDateLastMediaAdded(t *testing.T) {
|
||||
svc := &EmbyRemoteService{}
|
||||
mount := &model.EmbyMount{Base: model.Base{ID: "mount-1"}}
|
||||
acct := &model.StrmAccount{Base: model.Base{ID: "acct-1"}}
|
||||
cfg := &EmbyRemoteConfig{BaseURL: "http://localhost:8096"}
|
||||
|
||||
item := map[string]any{
|
||||
"Id": "series-1",
|
||||
"Name": "测试剧集",
|
||||
"DateCreated": "2023-01-01T00:00:00.0000000Z",
|
||||
"DateLastMediaAdded": "2024-05-20T10:00:00.0000000Z",
|
||||
}
|
||||
|
||||
media := svc.MapRemoteItemToMedia(context.Background(), mount, acct, cfg, item)
|
||||
expectedCreated, _ := time.Parse(time.RFC3339, "2023-01-01T00:00:00Z")
|
||||
expectedLastAdded, _ := time.Parse(time.RFC3339, "2024-05-20T10:00:00Z")
|
||||
if !media.CreatedAt.Equal(expectedCreated) {
|
||||
t.Fatalf("CreatedAt = %v, want %v", media.CreatedAt, expectedCreated)
|
||||
}
|
||||
if !media.UpdatedAt.Equal(expectedLastAdded) {
|
||||
t.Fatalf("UpdatedAt = %v, want %v", media.UpdatedAt, expectedLastAdded)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteSeriesCardsAutoAuthOnFirstBrowse(t *testing.T) {
|
||||
// 模拟远程 Emby:未认证兜底用户 ID "0" 被拒绝(与真实服务器一致),
|
||||
// 只有认证拿到的真实用户 GUID 才能浏览。
|
||||
var zeroUserHits atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/AuthenticateByName") {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"AccessToken": "real-token",
|
||||
"User": map[string]any{"Id": "real-user-guid"},
|
||||
})
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/emby/Users/0/Items" {
|
||||
zeroUserHits.Add(1)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_, _ = w.Write([]byte("Unrecognized Guid format."))
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/emby/Users/real-user-guid/Items" {
|
||||
q := r.URL.Query()
|
||||
if q.Get("ParentId") != "view-1" || q.Get("IncludeItemTypes") != "Series" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"TotalRecordCount": 1,
|
||||
"Items": []map[string]any{
|
||||
{
|
||||
"Id": "series-100",
|
||||
"Name": "测试剧",
|
||||
"Type": "Series",
|
||||
"ProductionYear": 2024,
|
||||
"RecursiveItemCount": 12,
|
||||
"ChildCount": 2,
|
||||
},
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
db := newServiceTestDB(t, &model.StrmAccount{}, &model.EmbyMount{})
|
||||
repos := repository.New(db)
|
||||
svc := NewEmbyRemoteService(&config.Config{}, zap.NewNop(), repos, NewCryptoService("", zap.NewNop()))
|
||||
|
||||
// 账号只配置了用户名/密码,从未「测试连接」:无 api_key、无 remote_user_id。
|
||||
rawConfig, _ := json.Marshal(map[string]string{
|
||||
"url": server.URL,
|
||||
"username": "user",
|
||||
"password": "pass",
|
||||
})
|
||||
acct := &model.StrmAccount{
|
||||
Base: model.Base{ID: "acct-1"},
|
||||
Name: "test-emby",
|
||||
Provider: model.StrmProviderEmbyRemote,
|
||||
Config: string(rawConfig),
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.StrmAccount.Create(t.Context(), acct); err != nil {
|
||||
t.Fatalf("create account: %v", err)
|
||||
}
|
||||
mount := &model.EmbyMount{
|
||||
Base: model.Base{ID: "mount-1"},
|
||||
AccountID: acct.ID,
|
||||
RemoteViewID: "view-1",
|
||||
RemoteViewName: "剧集库",
|
||||
CollectionType: "tvshows",
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.EmbyMount.Create(t.Context(), mount); err != nil {
|
||||
t.Fatalf("create mount: %v", err)
|
||||
}
|
||||
|
||||
cards, err := svc.RemoteSeriesCards(t.Context(), mount, acct, "view-1")
|
||||
if err != nil {
|
||||
t.Fatalf("RemoteSeriesCards on first browse failed: %v", err)
|
||||
}
|
||||
if len(cards) != 1 {
|
||||
t.Fatalf("cards = %d, want 1", len(cards))
|
||||
}
|
||||
if cards[0].Rep.Title != "测试剧" {
|
||||
t.Fatalf("title = %q, want 测试剧", cards[0].Rep.Title)
|
||||
}
|
||||
if cards[0].Count != 12 {
|
||||
t.Fatalf("count = %d, want 12 (RecursiveItemCount)", cards[0].Count)
|
||||
}
|
||||
if n := zeroUserHits.Load(); n != 0 {
|
||||
t.Fatalf("request hit /Users/0/Items %d time(s), want 0 (must use real user id)", n)
|
||||
}
|
||||
|
||||
// 首次浏览自动认证应把 token 与 remote_user_id 回写账号配置(等价于测试连接)。
|
||||
stored := map[string]string{}
|
||||
if err := json.Unmarshal([]byte(acct.Config), &stored); err != nil {
|
||||
t.Fatalf("decode account config: %v", err)
|
||||
}
|
||||
if stored["api_key"] == "" {
|
||||
t.Fatalf("account config missing api_key after first browse: %v", stored)
|
||||
}
|
||||
if stored["remote_user_id"] != "real-user-guid" {
|
||||
t.Fatalf("remote_user_id = %q, want real-user-guid (config %v)", stored["remote_user_id"], stored)
|
||||
}
|
||||
|
||||
// 第二次浏览不再需要认证步骤,直接命中真实用户 ID。
|
||||
if _, err := svc.RemoteSeriesCards(t.Context(), mount, acct, "view-1"); err != nil {
|
||||
t.Fatalf("RemoteSeriesCards second browse failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteSeriesCardsResolveUserIDFromAPIKey(t *testing.T) {
|
||||
// api_key 直连场景:账号只填了 token(无用户名/密码),从未回写过
|
||||
// remote_user_id。首次浏览应通过 /Users 列表解析出真实用户 GUID。
|
||||
var zeroUserHits atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/emby/Users" {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode([]map[string]any{
|
||||
{"Id": "real-user-guid", "Name": "admin"},
|
||||
})
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/emby/Users/0/Items" {
|
||||
zeroUserHits.Add(1)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_, _ = w.Write([]byte("Unrecognized Guid format."))
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/emby/Users/real-user-guid/Items" {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"TotalRecordCount": 1,
|
||||
"Items": []map[string]any{
|
||||
{"Id": "series-200", "Name": "API剧", "Type": "Series", "RecursiveItemCount": 8},
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
db := newServiceTestDB(t, &model.StrmAccount{}, &model.EmbyMount{})
|
||||
repos := repository.New(db)
|
||||
svc := NewEmbyRemoteService(&config.Config{}, zap.NewNop(), repos, NewCryptoService("", zap.NewNop()))
|
||||
|
||||
rawConfig, _ := json.Marshal(map[string]string{
|
||||
"url": server.URL,
|
||||
"token": "api-key-only",
|
||||
})
|
||||
acct := &model.StrmAccount{
|
||||
Base: model.Base{ID: "acct-2"},
|
||||
Name: "api-key-emby",
|
||||
Provider: model.StrmProviderEmbyRemote,
|
||||
Config: string(rawConfig),
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.StrmAccount.Create(t.Context(), acct); err != nil {
|
||||
t.Fatalf("create account: %v", err)
|
||||
}
|
||||
mount := &model.EmbyMount{
|
||||
Base: model.Base{ID: "mount-2"},
|
||||
AccountID: acct.ID,
|
||||
RemoteViewID: "view-2",
|
||||
RemoteViewName: "剧集库",
|
||||
CollectionType: "tvshows",
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.EmbyMount.Create(t.Context(), mount); err != nil {
|
||||
t.Fatalf("create mount: %v", err)
|
||||
}
|
||||
|
||||
cards, err := svc.RemoteSeriesCards(t.Context(), mount, acct, "view-2")
|
||||
if err != nil {
|
||||
t.Fatalf("RemoteSeriesCards with api_key only failed: %v", err)
|
||||
}
|
||||
if len(cards) != 1 || cards[0].Rep.Title != "API剧" {
|
||||
t.Fatalf("cards = %#v, want 1 card 测试剧", cards)
|
||||
}
|
||||
if n := zeroUserHits.Load(); n != 0 {
|
||||
t.Fatalf("request hit /Users/0/Items %d time(s), want 0", n)
|
||||
}
|
||||
stored := map[string]string{}
|
||||
if err := json.Unmarshal([]byte(acct.Config), &stored); err != nil {
|
||||
t.Fatalf("decode account config: %v", err)
|
||||
}
|
||||
if stored["remote_user_id"] != "real-user-guid" {
|
||||
t.Fatalf("remote_user_id = %q, want real-user-guid (config %v)", stored["remote_user_id"], stored)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteSearchMedia(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
+124
-46
@@ -10,19 +10,20 @@ import (
|
||||
)
|
||||
|
||||
type embySeriesGroup struct {
|
||||
ID string
|
||||
LibraryID string
|
||||
Name string
|
||||
PosterURL string
|
||||
BackdropURL string
|
||||
Overview string
|
||||
Rating float32
|
||||
Year int
|
||||
ReleaseDate string
|
||||
TMDbID int
|
||||
BangumiID int
|
||||
CreatedAt time.Time
|
||||
Episodes []model.Media
|
||||
ID string
|
||||
LibraryID string
|
||||
Name string
|
||||
PosterURL string
|
||||
BackdropURL string
|
||||
Overview string
|
||||
Rating float32
|
||||
Year int
|
||||
ReleaseDate string
|
||||
TMDbID int
|
||||
BangumiID int
|
||||
CreatedAt time.Time
|
||||
DateLastMediaAdded time.Time
|
||||
Episodes []model.Media
|
||||
}
|
||||
|
||||
type embySeasonGroup struct {
|
||||
@@ -49,11 +50,15 @@ func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (e
|
||||
q = e.applyUserMediaVisibility(ctx, q, userID)
|
||||
if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
|
||||
q = q.Where("series_id = ?", id)
|
||||
} else {
|
||||
// 虚拟 series ID 只可能来自 series_id 为空的媒体:
|
||||
// 有 series_id 时分组 key 就是 series_id 本身(UUID,不带虚拟前缀)。
|
||||
q = q.Where("series_id IS NULL OR series_id = ''")
|
||||
}
|
||||
if err := q.Order("media.season_num asc, media.episode_num asc, media.created_at asc").Limit(embySeriesGroupingLimit).Find(&rows).Error; err != nil {
|
||||
return embySeriesGroup{}, false, err
|
||||
}
|
||||
for _, group := range e.seriesGroupsFromMedia(rows) {
|
||||
for _, group := range e.seriesGroupsFromMedia(ctx, rows) {
|
||||
if group.ID == id {
|
||||
e.rememberSeriesGroup(group)
|
||||
return group, true, nil
|
||||
@@ -63,19 +68,20 @@ func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (e
|
||||
if series, err := e.repo.Series.FindByID(ctx, id); err != nil {
|
||||
return embySeriesGroup{}, false, err
|
||||
} else if series != nil {
|
||||
return embySeriesGroup{
|
||||
ID: series.ID,
|
||||
LibraryID: series.LibraryID,
|
||||
Name: series.Title,
|
||||
PosterURL: series.PosterURL,
|
||||
BackdropURL: series.BackdropURL,
|
||||
Overview: series.Overview,
|
||||
Rating: series.Rating,
|
||||
Year: series.Year,
|
||||
TMDbID: series.TMDbID,
|
||||
BangumiID: series.BangumiID,
|
||||
CreatedAt: series.CreatedAt,
|
||||
}, true, nil
|
||||
return embySeriesGroup{
|
||||
ID: series.ID,
|
||||
LibraryID: series.LibraryID,
|
||||
Name: series.Title,
|
||||
PosterURL: series.PosterURL,
|
||||
BackdropURL: series.BackdropURL,
|
||||
Overview: series.Overview,
|
||||
Rating: series.Rating,
|
||||
Year: series.Year,
|
||||
TMDbID: series.TMDbID,
|
||||
BangumiID: series.BangumiID,
|
||||
CreatedAt: series.CreatedAt,
|
||||
DateLastMediaAdded: series.CreatedAt,
|
||||
}, true, nil
|
||||
}
|
||||
}
|
||||
return embySeriesGroup{}, false, nil
|
||||
@@ -88,9 +94,18 @@ func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (e
|
||||
if season, ok := e.cachedSeasonGroup(id); ok {
|
||||
return season, true, nil
|
||||
}
|
||||
// 虚拟 Season ID 是 hash(seriesKey, seasonNum),无法反解出 series。
|
||||
// 常见情况(已刮削、series_id 非空)先用一条小型 DISTINCT 查询枚举候选对,
|
||||
// 在内存中算哈希匹配,命中后只加载该一部剧的剧集行,避免整库扫描。
|
||||
if season, ok, err := e.findSeasonGroupBySeriesCandidates(ctx, id, userID); err != nil {
|
||||
return embySeasonGroup{}, false, err
|
||||
} else if ok {
|
||||
return season, true, nil
|
||||
}
|
||||
// 回退:未刮削(series_id 为空,虚拟 key 由库名+名称派生)的媒体只能全量分组。
|
||||
var rows []model.Media
|
||||
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
|
||||
Where("season_num > 0 OR episode_num > 0")
|
||||
Where("(series_id IS NULL OR series_id = '') AND (season_num > 0 OR episode_num > 0)")
|
||||
q = e.applyUserMediaVisibility(ctx, q, userID)
|
||||
if err := q.
|
||||
Order("media.season_num asc, media.episode_num asc, media.created_at asc").
|
||||
@@ -98,7 +113,7 @@ func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (e
|
||||
Find(&rows).Error; err != nil {
|
||||
return embySeasonGroup{}, false, err
|
||||
}
|
||||
for _, series := range e.seriesGroupsFromMedia(rows) {
|
||||
for _, series := range e.seriesGroupsFromMedia(ctx, rows) {
|
||||
for _, season := range e.seasonsForSeries(series) {
|
||||
if season.ID == id {
|
||||
e.rememberSeriesGroup(series)
|
||||
@@ -109,30 +124,93 @@ func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (e
|
||||
return embySeasonGroup{}, false, nil
|
||||
}
|
||||
|
||||
func (e *EmbyService) seriesGroupsFromMedia(rows []model.Media) []embySeriesGroup {
|
||||
// findSeasonGroupBySeriesCandidates resolves virtual season IDs for media that
|
||||
// carry a real series_id: enumerate distinct (series_id, season_num) pairs via
|
||||
// SQL, hash each candidate to find the matching season, then load only that
|
||||
// one series' episodes.
|
||||
func (e *EmbyService) findSeasonGroupBySeriesCandidates(ctx context.Context, id, userID string) (embySeasonGroup, bool, error) {
|
||||
type seasonCandidate struct {
|
||||
SeriesID string
|
||||
SeasonNum int
|
||||
}
|
||||
var candidates []seasonCandidate
|
||||
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
|
||||
Select("DISTINCT series_id, season_num").
|
||||
Where("series_id <> '' AND (season_num > 0 OR episode_num > 0)")
|
||||
q = e.applyUserMediaVisibility(ctx, q, userID)
|
||||
if err := q.Find(&candidates).Error; err != nil {
|
||||
return embySeasonGroup{}, false, err
|
||||
}
|
||||
matched := make([]string, 0, 1)
|
||||
for _, cand := range candidates {
|
||||
if seasonID(cand.SeriesID, cand.SeasonNum) == id {
|
||||
matched = append(matched, cand.SeriesID)
|
||||
}
|
||||
}
|
||||
for _, matchedSeries := range matched {
|
||||
season, ok, err := e.seasonGroupForSeries(ctx, id, matchedSeries, userID)
|
||||
if err != nil || ok {
|
||||
return season, ok, err
|
||||
}
|
||||
}
|
||||
return embySeasonGroup{}, false, nil
|
||||
}
|
||||
|
||||
// seasonGroupForSeries rebuilds the season groups of one series (small row
|
||||
// set) and returns the one matching the virtual season id.
|
||||
func (e *EmbyService) seasonGroupForSeries(ctx context.Context, id, seriesID, userID string) (embySeasonGroup, bool, error) {
|
||||
var rows []model.Media
|
||||
rq := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
|
||||
Where("series_id = ? AND (season_num > 0 OR episode_num > 0)", seriesID)
|
||||
rq = e.applyUserMediaVisibility(ctx, rq, userID)
|
||||
if err := rq.
|
||||
Order("media.season_num asc, media.episode_num asc, media.created_at asc").
|
||||
Limit(embySeriesGroupingLimit).
|
||||
Find(&rows).Error; err != nil {
|
||||
return embySeasonGroup{}, false, err
|
||||
}
|
||||
for _, series := range e.seriesGroupsFromMedia(ctx, rows) {
|
||||
if series.ID != seriesID {
|
||||
continue
|
||||
}
|
||||
for _, season := range e.seasonsForSeries(series) {
|
||||
if season.ID == id {
|
||||
e.rememberSeriesGroup(series)
|
||||
return season, true, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return embySeasonGroup{}, false, nil
|
||||
}
|
||||
|
||||
func (e *EmbyService) seriesGroupsFromMedia(ctx context.Context, rows []model.Media) []embySeriesGroup {
|
||||
byID := map[string]*embySeriesGroup{}
|
||||
order := []string{}
|
||||
for _, row := range rows {
|
||||
row := row
|
||||
seriesID := e.seriesIDForMedia(&row)
|
||||
seriesID := e.seriesIDForMedia(ctx, &row)
|
||||
group, ok := byID[seriesID]
|
||||
if !ok {
|
||||
group = &embySeriesGroup{
|
||||
ID: seriesID,
|
||||
LibraryID: row.LibraryID,
|
||||
Name: e.seriesNameForMedia(&row),
|
||||
Year: row.Year,
|
||||
ReleaseDate: row.ReleaseDate,
|
||||
TMDbID: row.TMDbID,
|
||||
BangumiID: row.BangumiID,
|
||||
CreatedAt: row.CreatedAt,
|
||||
group = &embySeriesGroup{
|
||||
ID: seriesID,
|
||||
LibraryID: row.LibraryID,
|
||||
Name: e.seriesNameForMedia(ctx, &row),
|
||||
Year: row.Year,
|
||||
ReleaseDate: row.ReleaseDate,
|
||||
TMDbID: row.TMDbID,
|
||||
BangumiID: row.BangumiID,
|
||||
CreatedAt: row.CreatedAt,
|
||||
DateLastMediaAdded: row.CreatedAt,
|
||||
}
|
||||
byID[seriesID] = group
|
||||
order = append(order, seriesID)
|
||||
}
|
||||
if row.CreatedAt.Before(group.CreatedAt) || group.CreatedAt.IsZero() {
|
||||
group.CreatedAt = row.CreatedAt
|
||||
}
|
||||
if row.CreatedAt.After(group.DateLastMediaAdded) {
|
||||
group.DateLastMediaAdded = row.CreatedAt
|
||||
}
|
||||
byID[seriesID] = group
|
||||
order = append(order, seriesID)
|
||||
}
|
||||
if row.CreatedAt.After(group.CreatedAt) {
|
||||
group.CreatedAt = row.CreatedAt
|
||||
}
|
||||
if strings.TrimSpace(row.ReleaseDate) != "" && mediaReleaseSortTime(row).After(embySeriesReleaseSortTime(*group)) {
|
||||
group.ReleaseDate = row.ReleaseDate
|
||||
if row.Year > 0 {
|
||||
|
||||
@@ -429,10 +429,87 @@ func TestInferSeriesNameFromPath(t *testing.T) {
|
||||
want: "间谍过家家",
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
got := inferSeriesNameFromPath(tc.path)
|
||||
if got != tc.want {
|
||||
t.Errorf("inferSeriesNameFromPath(%q) = %q, want %q", tc.path, got, tc.want)
|
||||
for _, tc := range tests {
|
||||
got := inferSeriesNameFromPath(tc.path)
|
||||
if got != tc.want {
|
||||
t.Errorf("inferSeriesNameFromPath(%q) = %q, want %q", tc.path, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbySeriesSortByDateLastMediaAdded(t *testing.T) {
|
||||
svc := newTestEmbyService(t)
|
||||
lib := model.Library{Name: "测试剧库", Path: `/media/tv`, Type: "tv", Enabled: true}
|
||||
if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
|
||||
t.Fatalf("create library: %v", err)
|
||||
}
|
||||
t0 := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC)
|
||||
tOld := time.Date(2025, 6, 1, 0, 0, 0, 0, time.UTC)
|
||||
tNew := time.Date(2026, 8, 1, 0, 0, 0, 0, time.UTC)
|
||||
|
||||
// Series A: 较早创建,但最近添加了新一集 (Last episode at tNew)
|
||||
// Series B: 较晚创建,但最后一集在 tOld
|
||||
rows := []model.Media{
|
||||
{
|
||||
Base: model.Base{ID: "showA-ep01", CreatedAt: t0, UpdatedAt: t0},
|
||||
LibraryID: lib.ID,
|
||||
Title: "剧集A",
|
||||
Path: `/media/tv/剧集A/Season 01/剧集A.S01E01.mkv`,
|
||||
SeasonNum: 1,
|
||||
EpisodeNum: 1,
|
||||
},
|
||||
{
|
||||
Base: model.Base{ID: "showA-ep02", CreatedAt: tNew, UpdatedAt: tNew},
|
||||
LibraryID: lib.ID,
|
||||
Title: "剧集A",
|
||||
Path: `/media/tv/剧集A/Season 01/剧集A.S01E02.mkv`,
|
||||
SeasonNum: 1,
|
||||
EpisodeNum: 2,
|
||||
},
|
||||
{
|
||||
Base: model.Base{ID: "showB-ep01", CreatedAt: tOld.Add(-24 * time.Hour), UpdatedAt: tOld.Add(-24 * time.Hour)},
|
||||
LibraryID: lib.ID,
|
||||
Title: "剧集B",
|
||||
Path: `/media/tv/剧集B/Season 01/剧集B.S01E01.mkv`,
|
||||
SeasonNum: 1,
|
||||
EpisodeNum: 1,
|
||||
},
|
||||
{
|
||||
Base: model.Base{ID: "showB-ep02", CreatedAt: tOld, UpdatedAt: tOld},
|
||||
LibraryID: lib.ID,
|
||||
Title: "剧集B",
|
||||
Path: `/media/tv/剧集B/Season 01/剧集B.S01E02.mkv`,
|
||||
SeasonNum: 1,
|
||||
EpisodeNum: 2,
|
||||
},
|
||||
}
|
||||
for _, m := range rows {
|
||||
if err := svc.repo.DB.Create(&m).Error; err != nil {
|
||||
t.Fatalf("create media: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 降序排序:剧集A最后一集在 tNew,剧集B最后一集在 tOld,剧集A应排在第一位
|
||||
res, err := svc.Items(t.Context(), ItemsParams{
|
||||
ParentID: lib.ID,
|
||||
SortBy: "DateLastMediaAdded",
|
||||
SortOrder: "Descending",
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("items DateLastMediaAdded: %v", err)
|
||||
}
|
||||
items := res["Items"].([]map[string]any)
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("items count = %d, want 2", len(items))
|
||||
}
|
||||
if items[0]["Name"] != "剧集A" {
|
||||
t.Fatalf("first item = %v, want 剧集A (last episode at tNew)", items[0]["Name"])
|
||||
}
|
||||
if items[1]["Name"] != "剧集B" {
|
||||
t.Fatalf("second item = %v, want 剧集B", items[1]["Name"])
|
||||
}
|
||||
if items[0]["DateLastMediaAdded"] != tNew {
|
||||
t.Fatalf("DateLastMediaAdded = %v, want %v", items[0]["DateLastMediaAdded"], tNew)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,21 +14,22 @@ import (
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
func (e *EmbyService) seriesIDForMedia(m *model.Media) string {
|
||||
func (e *EmbyService) seriesIDForMedia(ctx context.Context, m *model.Media) string {
|
||||
if strings.TrimSpace(m.SeriesID) != "" {
|
||||
return m.SeriesID
|
||||
}
|
||||
return stableEmbyID(embyVirtualSeriesPrefix, m.LibraryID, e.seriesNameForMedia(m))
|
||||
return stableEmbyID(embyVirtualSeriesPrefix, m.LibraryID, e.seriesNameForMedia(ctx, m))
|
||||
}
|
||||
|
||||
func (e *EmbyService) seasonIDForMedia(m *model.Media) string {
|
||||
return seasonID(e.seriesIDForMedia(m), m.SeasonNum)
|
||||
func (e *EmbyService) seasonIDForMedia(ctx context.Context, m *model.Media) string {
|
||||
return seasonID(e.seriesIDForMedia(ctx, m), m.SeasonNum)
|
||||
}
|
||||
|
||||
func (e *EmbyService) seriesNameForMedia(m *model.Media) string {
|
||||
func (e *EmbyService) seriesNameForMedia(ctx context.Context, m *model.Media) string {
|
||||
if strings.TrimSpace(m.SeriesID) != "" {
|
||||
if series, err := e.repo.Series.FindByID(context.Background(), m.SeriesID); err == nil && series != nil && strings.TrimSpace(series.Title) != "" {
|
||||
return series.Title
|
||||
// 走请求级缓存;无缓存 ctx 时退化为单次查询。
|
||||
if title, ok, err := e.payloadSeriesTitle(ctx, m.SeriesID); err == nil && ok && strings.TrimSpace(title) != "" {
|
||||
return title
|
||||
}
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(m.ScrapeStatus), "matched") && strings.TrimSpace(m.Title) != "" {
|
||||
@@ -125,13 +126,28 @@ func sortSeriesGroups(groups []embySeriesGroup, p ItemsParams) {
|
||||
}
|
||||
return groups[i].Name < groups[j].Name
|
||||
})
|
||||
case "datecreated":
|
||||
sort.SliceStable(groups, func(i, j int) bool {
|
||||
if strings.EqualFold(p.SortOrder, "Ascending") {
|
||||
return groups[i].CreatedAt.Before(groups[j].CreatedAt)
|
||||
}
|
||||
return groups[i].CreatedAt.After(groups[j].CreatedAt)
|
||||
})
|
||||
case "datecreated":
|
||||
sort.SliceStable(groups, func(i, j int) bool {
|
||||
if strings.EqualFold(p.SortOrder, "Ascending") {
|
||||
return groups[i].CreatedAt.Before(groups[j].CreatedAt)
|
||||
}
|
||||
return groups[i].CreatedAt.After(groups[j].CreatedAt)
|
||||
})
|
||||
case "datelastmediaadded", "datelastcontentadded":
|
||||
sort.SliceStable(groups, func(i, j int) bool {
|
||||
tI := groups[i].DateLastMediaAdded
|
||||
if tI.IsZero() {
|
||||
tI = groups[i].CreatedAt
|
||||
}
|
||||
tJ := groups[j].DateLastMediaAdded
|
||||
if tJ.IsZero() {
|
||||
tJ = groups[j].CreatedAt
|
||||
}
|
||||
if strings.EqualFold(p.SortOrder, "Ascending") {
|
||||
return tI.Before(tJ)
|
||||
}
|
||||
return tI.After(tJ)
|
||||
})
|
||||
default:
|
||||
sort.SliceStable(groups, func(i, j int) bool {
|
||||
if strings.EqualFold(p.SortOrder, "Ascending") {
|
||||
|
||||
@@ -10,22 +10,27 @@ func (e *EmbyService) seriesPayload(group embySeriesGroup) map[string]any {
|
||||
if group.BackdropURL != "" {
|
||||
backdropTags = append(backdropTags, group.ID+"-bd")
|
||||
}
|
||||
item := map[string]any{
|
||||
"Id": group.ID,
|
||||
"Name": group.Name,
|
||||
"ServerId": embyServerID,
|
||||
"Type": "Series",
|
||||
"MediaType": "Video",
|
||||
"IsFolder": true,
|
||||
"ParentId": group.LibraryID,
|
||||
"ProductionYear": group.Year,
|
||||
"Overview": group.Overview,
|
||||
"CommunityRating": group.Rating,
|
||||
"RecursiveItemCount": len(group.Episodes),
|
||||
"ChildCount": len(e.seasonsForSeries(group)),
|
||||
"DateCreated": group.CreatedAt,
|
||||
"ImageTags": imageTags,
|
||||
"BackdropImageTags": backdropTags,
|
||||
lastMediaAdded := group.DateLastMediaAdded
|
||||
if lastMediaAdded.IsZero() {
|
||||
lastMediaAdded = group.CreatedAt
|
||||
}
|
||||
item := map[string]any{
|
||||
"Id": group.ID,
|
||||
"Name": group.Name,
|
||||
"ServerId": embyServerID,
|
||||
"Type": "Series",
|
||||
"MediaType": "Video",
|
||||
"IsFolder": true,
|
||||
"ParentId": group.LibraryID,
|
||||
"ProductionYear": group.Year,
|
||||
"Overview": group.Overview,
|
||||
"CommunityRating": group.Rating,
|
||||
"RecursiveItemCount": len(group.Episodes),
|
||||
"ChildCount": len(e.seasonsForSeries(group)),
|
||||
"DateCreated": group.CreatedAt,
|
||||
"DateLastMediaAdded": lastMediaAdded,
|
||||
"ImageTags": imageTags,
|
||||
"BackdropImageTags": backdropTags,
|
||||
"ProviderIds": map[string]string{
|
||||
"Tmdb": intToStr(group.TMDbID),
|
||||
"Bangumi": intToStr(group.BangumiID),
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
@@ -218,7 +219,22 @@ func mergedRemoteUserData(raw any, history *model.PlaybackHistory) map[string]an
|
||||
return userData
|
||||
}
|
||||
|
||||
// embyInvalMu 节流全量缓存失效:播放期间客户端每 5-10s 上报一次进度,
|
||||
// 每次都 SCAN+DEL 全部 media:emby:* 缓存会把缓存命中率持续打穿(其他
|
||||
// 客户端每次翻页都回源 SQL)。条目载荷的 UserData 在请求时动态合并,
|
||||
// 进度类变更做 30s 节流即可,不影响正确性观感。
|
||||
var (
|
||||
embyInvalMu sync.Mutex
|
||||
embyInvalLast time.Time
|
||||
)
|
||||
|
||||
func (e *EmbyService) invalidateEmbyItemsCache(ctx context.Context) {
|
||||
embyInvalMu.Lock()
|
||||
defer embyInvalMu.Unlock()
|
||||
if !embyInvalLast.IsZero() && time.Since(embyInvalLast) < 30*time.Second {
|
||||
return
|
||||
}
|
||||
embyInvalLast = time.Now()
|
||||
if e.cache != nil {
|
||||
e.cache.DeletePrefix(ctx, "media:emby:")
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
|
||||
@@ -152,7 +153,7 @@ func (s *FFmpegToolsService) StartInstall(ctx context.Context) error {
|
||||
s.mu.Unlock()
|
||||
|
||||
s.setMessage("准备下载…")
|
||||
go s.runInstall()
|
||||
helper.Go(s.log, "ffmpeg.install", s.runInstall)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -13,9 +13,12 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
@@ -52,6 +55,33 @@ func NewImageProxy(cfg *config.Config, log *zap.Logger) *ImageProxy {
|
||||
// from image.tmdb.org via their HTTP proxy without extra config. On
|
||||
// Windows we also honor the current user's system proxy settings.
|
||||
transport := NewExternalTransport()
|
||||
if proxyConfiguredForImageFetch() {
|
||||
// 走本地代理(如 127.0.0.1:7890)时,拨号目标是代理本身,
|
||||
// 连接层 SSRF 校验会误杀本地回环代理;此时沿用 URL 级校验。
|
||||
log.Info("image proxy: outbound proxy detected, connection-level SSRF guard disabled")
|
||||
} else {
|
||||
// 仅 URL 解析层的 isPrivateHost 可被十进制/十六进制 IP、解析到
|
||||
// 私网的域名与 DNS rebinding 绕过;在拨号层对最终连接 IP 做二次
|
||||
// 校验(含重定向后的每条连接)堵住该旁路。
|
||||
dialer := &net.Dialer{
|
||||
Timeout: 15 * time.Second,
|
||||
Control: func(_, address string, _ syscall.RawConn) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
return errors.New("image proxy: refusing non-IP dial target")
|
||||
}
|
||||
if isPrivateIP(ip) {
|
||||
return errors.New("image proxy: requests to private/internal hosts are not allowed")
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
transport.DialContext = dialer.DialContext
|
||||
}
|
||||
return &ImageProxy{
|
||||
cfg: cfg,
|
||||
log: log,
|
||||
@@ -60,6 +90,16 @@ func NewImageProxy(cfg *config.Config, log *zap.Logger) *ImageProxy {
|
||||
}
|
||||
}
|
||||
|
||||
// proxyConfiguredForImageFetch 探测环境变量或系统代理是否会影响图片抓取。
|
||||
func proxyConfiguredForImageFetch() bool {
|
||||
req, err := http.NewRequest(http.MethodGet, "https://image.tmdb.org/", nil)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
proxy, err := ProxyFromEnvironmentOrSystem(req)
|
||||
return err == nil && proxy != nil
|
||||
}
|
||||
|
||||
// SetLibraryRootsProvider injects a callback that returns the current set of
|
||||
// media library root directories. Sidecar posters live under these roots
|
||||
// (which are arbitrary, user-defined, and not necessarily under the
|
||||
|
||||
@@ -35,13 +35,18 @@ func isPrivateHost(host string) bool {
|
||||
if host == "" {
|
||||
return true
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip != nil {
|
||||
return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified()
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return isPrivateIP(ip)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isPrivateIP 判定单个 IP 是否属于回环/私网/链路本地/未指定地址。
|
||||
func isPrivateIP(ip net.IP) bool {
|
||||
return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() ||
|
||||
ip.IsLinkLocalMulticast() || ip.IsUnspecified()
|
||||
}
|
||||
|
||||
// isAllowedLocalPath restricts local file reads to known-safe roots.
|
||||
func (p *ImageProxy) isAllowedLocalPath(abs string) bool {
|
||||
roots := []string{p.cfg.App.DataDir, p.cfg.Cache.CacheDir, p.cfg.Media.MoviesDir, p.cfg.Media.TVDir, p.cfg.Media.AnimeDir}
|
||||
|
||||
@@ -27,10 +27,11 @@ type libraryRowsCacheValue struct {
|
||||
}
|
||||
|
||||
type SeriesCard struct {
|
||||
Key string `json:"key"`
|
||||
Rep model.Media `json:"rep"`
|
||||
LinkMedia model.Media `json:"linkMedia"`
|
||||
Count int `json:"count"`
|
||||
Key string `json:"key"`
|
||||
Rep model.Media `json:"rep"`
|
||||
LinkMedia model.Media `json:"linkMedia"`
|
||||
Count int `json:"count"`
|
||||
LastAddedAt *time.Time `json:"last_added_at,omitempty"`
|
||||
}
|
||||
|
||||
type seriesCardGroup struct {
|
||||
@@ -189,19 +190,40 @@ func (s *MediaService) ListMediaEpisodes(ctx context.Context, mediaID string, vi
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没有聚合到多集,尝试同父目录匹配
|
||||
if len(out) <= 1 && target.Path != "" {
|
||||
targetDir := filepath.Dir(strings.ReplaceAll(target.Path, "\\", "/"))
|
||||
dirMatches := make([]model.Media, 0)
|
||||
for _, row := range rows {
|
||||
if row.Path != "" && filepath.Dir(strings.ReplaceAll(row.Path, "\\", "/")) == targetDir {
|
||||
dirMatches = append(dirMatches, row)
|
||||
// 如果没有聚合到多集,尝试同父目录匹配(排除合集目录和公共分类目录,且同目录文件不能是互不相同的独立电影)
|
||||
if len(out) <= 1 && target.Path != "" {
|
||||
targetDir := filepath.Dir(strings.ReplaceAll(target.Path, "\\", "/"))
|
||||
parentBase := filepath.Base(targetDir)
|
||||
if !mediaParentLooksLikeCollection(target.Path) && !seriesTitleIsGenericContainer(parentBase, *target) {
|
||||
targetTitleNorm := normalizeSeriesTitle(target.Title)
|
||||
targetDirNorm := normalizeSeriesTitle(parentBase)
|
||||
dirMatches := make([]model.Media, 0)
|
||||
for _, row := range rows {
|
||||
if row.Path == "" || filepath.Dir(strings.ReplaceAll(row.Path, "\\", "/")) != targetDir {
|
||||
continue
|
||||
}
|
||||
if row.ID == target.ID {
|
||||
dirMatches = append(dirMatches, row)
|
||||
continue
|
||||
}
|
||||
rowTitleNorm := normalizeSeriesTitle(row.Title)
|
||||
allowMatch := false
|
||||
if isGenericMovieTitle(rowTitleNorm) || isGenericMovieTitle(targetTitleNorm) {
|
||||
allowMatch = true
|
||||
} else if rowTitleNorm != "" && rowTitleNorm == targetTitleNorm {
|
||||
allowMatch = true
|
||||
} else if rowTitleNorm != "" && targetDirNorm != "" && rowTitleNorm == targetDirNorm {
|
||||
allowMatch = true
|
||||
}
|
||||
if allowMatch {
|
||||
dirMatches = append(dirMatches, row)
|
||||
}
|
||||
}
|
||||
if len(dirMatches) > 1 {
|
||||
out = dirMatches
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(dirMatches) > 1 {
|
||||
out = dirMatches
|
||||
}
|
||||
}
|
||||
|
||||
if len(out) == 0 {
|
||||
out = []model.Media{*target}
|
||||
@@ -256,12 +278,20 @@ func groupMediaSeriesCards(items []model.Media) []SeriesCard {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
itemAdded := item.CreatedAt
|
||||
if itemAdded.IsZero() {
|
||||
itemAdded = item.UpdatedAt
|
||||
}
|
||||
if idx, ok := byKey[key]; ok {
|
||||
group := &groups[idx]
|
||||
if latest := seriesMediaTime(item); latest.After(group.latest) {
|
||||
group.latest = latest
|
||||
}
|
||||
card := &group.card
|
||||
if card.LastAddedAt == nil || (!itemAdded.IsZero() && itemAdded.After(*card.LastAddedAt)) {
|
||||
t := itemAdded
|
||||
card.LastAddedAt = &t
|
||||
}
|
||||
// A shared external ID means duplicate encodes/locations for movies,
|
||||
// not multiple episodes. Keep a single movie card without presenting
|
||||
// its versions as an "N episodes" collection.
|
||||
@@ -284,9 +314,14 @@ func groupMediaSeriesCards(items []model.Media) []SeriesCard {
|
||||
}
|
||||
continue
|
||||
}
|
||||
var initialLastAdded *time.Time
|
||||
if !itemAdded.IsZero() {
|
||||
t := itemAdded
|
||||
initialLastAdded = &t
|
||||
}
|
||||
byKey[key] = len(groups)
|
||||
groups = append(groups, seriesCardGroup{
|
||||
card: SeriesCard{Key: key, Rep: item, LinkMedia: item, Count: 1},
|
||||
card: SeriesCard{Key: key, Rep: item, LinkMedia: item, Count: 1, LastAddedAt: initialLastAdded},
|
||||
latest: seriesMediaTime(item),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -11,6 +11,16 @@ import (
|
||||
|
||||
var episodicPathRE = regexp.MustCompile(`(?i)[\\/](?:电视剧|剧集|连续剧|短剧|国产剧|国剧|大陆剧|华语剧|国产电视剧|大陆电视剧|华语电视剧|欧美剧|欧美电视剧|美剧|英剧|日韩剧|日韩电视剧|日剧|韩剧|港剧|台剧|港台剧|泰剧|综艺|纪录片|儿童|动漫|番剧|国漫|日番|韩漫|美漫|欧美动漫|欧美动画|其他动漫|tv|series|shows?|season[\s._-]*\d|s\d{1,2}(?:[\s._-]|[\\/])|special[\s._-]*episodes?|specials?|sp|ovas?|oads?|extras?|bonus(?:es)?|omake|特别篇|特別篇|番外篇?|特典|外传|外傳|总集篇|總集篇)[\\/]`)
|
||||
|
||||
var genericMovieTitleRE = regexp.MustCompile(`(?i)^(?:cd\s*\d+|part\s*\d+|disc\s*\d+|disk\s*\d+|dvd\s*\d+|movie|film|video|main|feature|track\s*\d+|preview|sample|trailer|\d{3,4}p|4k|2160p|1080p|720p)$`)
|
||||
|
||||
func isGenericMovieTitle(title string) bool {
|
||||
title = strings.TrimSpace(title)
|
||||
if title == "" {
|
||||
return true
|
||||
}
|
||||
return genericMovieTitleRE.MatchString(title)
|
||||
}
|
||||
|
||||
func mediaSeriesKey(media model.Media) string {
|
||||
return compactSeriesKey(mediaSeriesRawKey(media))
|
||||
}
|
||||
@@ -54,7 +64,11 @@ func mediaSeriesRawKey(media model.Media) string {
|
||||
return seriesFingerprint("movie-external", fmt.Sprintf("bgm:%d", media.BangumiID))
|
||||
}
|
||||
if fromPath != "" && !mediaParentLooksLikeCollection(media.Path) {
|
||||
return seriesFingerprint("library-path", media.LibraryID, fromPath)
|
||||
titleNorm := normalizeSeriesTitle(media.Title)
|
||||
fromPathNorm := normalizeSeriesTitle(fromPath)
|
||||
if titleNorm == "" || isGenericMovieTitle(titleNorm) || titleNorm == fromPathNorm {
|
||||
return seriesFingerprint("library-path", media.LibraryID, fromPath)
|
||||
}
|
||||
}
|
||||
return seriesFingerprint("library-title", media.LibraryID, normalizeSeriesTitle(media.Title))
|
||||
}
|
||||
|
||||
@@ -47,10 +47,14 @@ func TestListRecentSeriesCardsCountsAllEpisodesInSeries(t *testing.T) {
|
||||
if len(cards) != 1 {
|
||||
t.Fatalf("recent cards = %#v, want one series card", cards)
|
||||
}
|
||||
if cards[0].Count != 40 {
|
||||
t.Fatalf("recent series count = %d, want full 40 episodes", cards[0].Count)
|
||||
if cards[0].Count != 40 {
|
||||
t.Fatalf("recent series count = %d, want full 40 episodes", cards[0].Count)
|
||||
}
|
||||
expectedLastAdded := now.Add(40 * time.Minute)
|
||||
if cards[0].LastAddedAt == nil || !cards[0].LastAddedAt.Equal(expectedLastAdded) {
|
||||
t.Fatalf("recent series LastAddedAt = %v, want %v", cards[0].LastAddedAt, expectedLastAdded)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMediaSeriesKeyCollapsesNestedSpecialFolders(t *testing.T) {
|
||||
main := model.Media{
|
||||
@@ -445,3 +449,67 @@ func TestGroupMediaSeriesCardsBridgesReleaseFoldersByMatchedSeriesTitle(t *testi
|
||||
t.Fatalf("cards=%#v, want one series bridged by matched title", cards)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGroupMediaSeriesCardsKeepsIndependentMoviesSeparateInSharedSubdirectory(t *testing.T) {
|
||||
// 同一分类子目录下存放多部不同标题的独立电影,不应被强制折叠成 1 部
|
||||
items := []model.Media{
|
||||
{LibraryID: "movies", Title: "老师2024偷窥篇", Path: `/media/小姐姐/国产/nana/老师2024偷窥篇.strm`},
|
||||
{LibraryID: "movies", Title: "紫光灯下的肉体诱惑", Path: `/media/小姐姐/国产/nana/紫光灯下的肉体诱惑.strm`},
|
||||
{LibraryID: "movies", Title: "修洗衣机", Path: `/media/小姐姐/国产/nana/修洗衣机.strm`},
|
||||
}
|
||||
cards := groupMediaSeriesCards(items)
|
||||
if len(cards) != 3 {
|
||||
t.Fatalf("got %d cards, want 3 independent movie cards", len(cards))
|
||||
}
|
||||
|
||||
// 但同一部电影的 CD1 和 CD2 仍应正确折叠为 1 部
|
||||
cdItems := []model.Media{
|
||||
{LibraryID: "movies", Title: "cd1", Path: `/media/电影/指环王 (2001)/cd1.mkv`},
|
||||
{LibraryID: "movies", Title: "cd2", Path: `/media/电影/指环王 (2001)/cd2.mkv`},
|
||||
}
|
||||
cdCards := groupMediaSeriesCards(cdItems)
|
||||
if len(cdCards) != 1 {
|
||||
t.Fatalf("got %d cards for cd1/cd2, want 1 folded movie card", len(cdCards))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListMediaEpisodesKeepsIndependentMoviesSeparate(t *testing.T) {
|
||||
db := newServiceTestDB(t, &model.Library{}, &model.Media{})
|
||||
repos := repository.New(db)
|
||||
lib := model.Library{Base: model.Base{ID: "lib-movies"}, Name: "电影", Type: "movies", Enabled: true}
|
||||
if err := repos.DB.Create(&lib).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
m1 := model.Media{
|
||||
Base: model.Base{ID: "m1"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "老师2024偷窥篇",
|
||||
Path: `/media/小姐姐/国产/nana/老师2024偷窥篇.strm`,
|
||||
}
|
||||
m2 := model.Media{
|
||||
Base: model.Base{ID: "m2"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "紫光灯下的肉体诱惑",
|
||||
Path: `/media/小姐姐/国产/nana/紫光灯下的肉体诱惑.strm`,
|
||||
}
|
||||
m3 := model.Media{
|
||||
Base: model.Base{ID: "m3"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "修洗衣机",
|
||||
Path: `/media/小姐姐/国产/nana/修洗衣机.strm`,
|
||||
}
|
||||
if err := repos.DB.Create(&[]model.Media{m1, m2, m3}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
|
||||
eps, err := svc.ListMediaEpisodes(t.Context(), "m1", MediaVisibility{IncludeNSFW: true})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(eps) != 1 || eps[0].ID != "m1" {
|
||||
t.Fatalf("ListMediaEpisodes got %#v, want exactly m1", eps)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
@@ -219,24 +220,50 @@ func (o *OrganizerService) replaceVersions(ctx context.Context, src string, exis
|
||||
o.log.Warn("organize replace sidecar artwork failed",
|
||||
zap.String("from", src), zap.String("to", dst), zap.Error(err))
|
||||
}
|
||||
// New file is safely staged; the transfer succeeded so it is now safe to
|
||||
// supersede the existing lower-res versions.
|
||||
for _, e := range existing {
|
||||
if nfo := nfoPath(e); nfo != "" {
|
||||
_ = os.Remove(nfo)
|
||||
// New file is safely staged. 先把现有 dst 改名为备份、rename stage→dst
|
||||
// 成功后,才删除旧版本:此前顺序是先删旧版本再 rename,一旦 rename
|
||||
// 失败(Windows 下 dst 被播放器/杀软占用很常见),cleanup 会删掉
|
||||
// stage——旧版本已删、move 模式下源已不在、新文件也删,数据彻底丢失。
|
||||
var backup string
|
||||
if _, err := os.Stat(dst); err == nil {
|
||||
backup = dst + ".replacing-" + randomSuffix()
|
||||
if err := os.Rename(dst, backup); err != nil {
|
||||
cleanup()
|
||||
return fmt.Errorf("备份现有文件失败(可能被其他程序占用):%w", err)
|
||||
}
|
||||
if err := os.Remove(e); err != nil && !os.IsNotExist(err) {
|
||||
o.log.Warn("organize replace remove existing failed",
|
||||
zap.String("path", e), zap.Error(err))
|
||||
}
|
||||
if err := os.Rename(stage, dst); err != nil {
|
||||
if backup != "" {
|
||||
if rbErr := os.Rename(backup, dst); rbErr != nil {
|
||||
o.log.Error("organize replace restore backup failed",
|
||||
zap.String("backup", backup), zap.Error(rbErr))
|
||||
}
|
||||
}
|
||||
cleanup()
|
||||
return err
|
||||
}
|
||||
// 新文件已就位,现在才删除被取代的旧版本。dst 路径此时已是新文件,
|
||||
// 文件级删除必须跳过(DB 行仍按原语义清理)。
|
||||
for _, e := range existing {
|
||||
if e != dst {
|
||||
if nfo := nfoPath(e); nfo != "" {
|
||||
_ = os.Remove(nfo)
|
||||
}
|
||||
if err := os.Remove(e); err != nil && !os.IsNotExist(err) {
|
||||
o.log.Warn("organize replace remove existing failed",
|
||||
zap.String("path", e), zap.Error(err))
|
||||
}
|
||||
}
|
||||
if o.repo != nil && o.repo.DB != nil {
|
||||
_ = o.repo.DB.WithContext(ctx).Unscoped().Where("path = ?", e).Delete(&model.Media{}).Error
|
||||
}
|
||||
}
|
||||
// Move staged file + sidecars into the final path.
|
||||
if err := os.Rename(stage, dst); err != nil {
|
||||
cleanup()
|
||||
return err
|
||||
if backup != "" {
|
||||
// 备份文件即被取代的旧 dst 内容,新文件已成功落位后移除。
|
||||
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
|
||||
o.log.Warn("organize replace remove backup failed",
|
||||
zap.String("path", backup), zap.Error(err))
|
||||
}
|
||||
}
|
||||
moveSidecarRename(nfoPath(stage), nfoPath(dst))
|
||||
moveStagedArtwork(stage, dst)
|
||||
|
||||
@@ -291,10 +291,29 @@ func (p *PlaybackService) GetPlaylist(ctx context.Context, playlistID string) (*
|
||||
return &PlaylistDetail{Playlist: pl, Items: ordered}, nil
|
||||
}
|
||||
|
||||
// ErrPlaylistForbidden 表示当前用户无权操作目标播放列表。
|
||||
var ErrPlaylistForbidden = errors.New("forbidden")
|
||||
|
||||
// EnsurePlaylistOwner 校验播放列表属主;admin 可操作任意列表。
|
||||
// 非存在的列表返回 gorm.ErrRecordNotFound。
|
||||
func (p *PlaybackService) EnsurePlaylistOwner(ctx context.Context, playlistID, userID string, isAdmin bool) error {
|
||||
var pl model.Playlist
|
||||
if err := p.repo.DB.WithContext(ctx).Select("user_id").Where("id = ?", playlistID).First(&pl).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if pl.UserID != userID && !isAdmin {
|
||||
return ErrPlaylistForbidden
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddToPlaylist appends a media item to the end of a playlist.
|
||||
func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID string) error {
|
||||
func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, userID, mediaID string, isAdmin bool) error {
|
||||
if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil {
|
||||
return err
|
||||
}
|
||||
var count int64
|
||||
if err := p.repo.DB.Model(&model.PlaylistItem{}).
|
||||
if err := p.repo.DB.WithContext(ctx).Model(&model.PlaylistItem{}).
|
||||
Where("playlist_id = ?", playlistID).Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -307,14 +326,20 @@ func (p *PlaybackService) AddToPlaylist(ctx context.Context, playlistID, mediaID
|
||||
}
|
||||
|
||||
// RemoveFromPlaylist 物理删除播放列表项(幂等)。
|
||||
func (p *PlaybackService) RemoveFromPlaylist(ctx context.Context, playlistID, mediaID string) error {
|
||||
func (p *PlaybackService) RemoveFromPlaylist(ctx context.Context, playlistID, userID, mediaID string, isAdmin bool) error {
|
||||
if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil {
|
||||
return err
|
||||
}
|
||||
return p.repo.DB.WithContext(ctx).Unscoped().
|
||||
Where("playlist_id = ? AND media_id = ?", playlistID, mediaID).
|
||||
Delete(&model.PlaylistItem{}).Error
|
||||
}
|
||||
|
||||
// DeletePlaylist 物理删除播放列表及其全部条目。
|
||||
func (p *PlaybackService) DeletePlaylist(ctx context.Context, playlistID string) error {
|
||||
func (p *PlaybackService) DeletePlaylist(ctx context.Context, playlistID, userID string, isAdmin bool) error {
|
||||
if err := p.EnsurePlaylistOwner(ctx, playlistID, userID, isAdmin); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := p.repo.DB.WithContext(ctx).Unscoped().Where("playlist_id = ?", playlistID).
|
||||
Delete(&model.PlaylistItem{}).Error; err != nil {
|
||||
return err
|
||||
|
||||
@@ -28,6 +28,8 @@ func ApplyRuntimeSettings(ctx context.Context, cfg *config.Config, repos *reposi
|
||||
}
|
||||
|
||||
func ApplyRuntimeSetting(cfg *config.Config, key, value string) {
|
||||
config.RuntimeMu.Lock()
|
||||
defer config.RuntimeMu.Unlock()
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
@@ -45,7 +46,7 @@ func (s *ScannerService) startLocalMediaProbeWorkers() {
|
||||
s.localMediaProbeOnce.Do(func() {
|
||||
workers := s.ffprobeWorkerCount()
|
||||
for i := 0; i < workers; i++ {
|
||||
go s.localMediaProbeWorker()
|
||||
helper.Go(s.log, "scanner.probeWorker", s.localMediaProbeWorker)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
)
|
||||
|
||||
func (s *ScannerService) invalidateMediaCache(ctx context.Context) {
|
||||
@@ -18,10 +20,12 @@ func (s *ScannerService) startAutoScrape(ctx context.Context, libraryID string)
|
||||
scrapeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Minute)
|
||||
go func() {
|
||||
defer cancel()
|
||||
_, err := s.scraper.EnrichLibraryDetailedWithOptions(scrapeCtx, libraryID, skipEpisodeArtworkOptions(false))
|
||||
if err != nil {
|
||||
s.log.Warn("scraper enrich failed", zap.Error(err))
|
||||
return
|
||||
}
|
||||
// 扫描触发的后台刮削与请求线程无关,panic 只记日志,不能带崩进程。
|
||||
helper.Run(s.log, "scanner.autoScrape", func() {
|
||||
_, err := s.scraper.EnrichLibraryDetailedWithOptions(scrapeCtx, libraryID, skipEpisodeArtworkOptions(false))
|
||||
if err != nil {
|
||||
s.log.Warn("scraper enrich failed", zap.Error(err))
|
||||
}
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
@@ -16,8 +16,33 @@ func (s *ScannerService) RemovePath(ctx context.Context, path string) (int64, er
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return 0, nil // still exists; nothing to remove
|
||||
}
|
||||
// 目录整体消失(删除/改名离开):连同其子树下的媒体行一并移除。
|
||||
// 此前只删 path 精确匹配的行——目录本身通常没有 media 行,导致目录
|
||||
// 改名后旧子树记录全部失联,只有全量扫描才能修复。
|
||||
prefix := filepath.Clean(path) + string(filepath.Separator)
|
||||
var rows []struct {
|
||||
ID string
|
||||
Path string
|
||||
}
|
||||
if err := s.repo.DB.WithContext(ctx).Model(&model.Media{}).
|
||||
Select("id, path").
|
||||
Where("path = ? OR path LIKE ?", path, prefix+"%").
|
||||
Find(&rows).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
ids := make([]string, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
// LIKE 里的 % _ 是通配符(候选集只会偏大),用 Go 前缀精确过滤,
|
||||
// 避免对含 % / _ 的路径误删。
|
||||
if row.Path == path || strings.HasPrefix(filepath.Clean(row.Path), prefix) {
|
||||
ids = append(ids, row.ID)
|
||||
}
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
res := s.repo.DB.WithContext(ctx).Unscoped().
|
||||
Where("path = ?", path).
|
||||
Where("id IN ?", ids).
|
||||
Delete(&model.Media{})
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
s.invalidateMediaCache(ctx)
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
|
||||
@@ -139,7 +140,7 @@ func (s *SchedulerService) Start(ctx context.Context) {
|
||||
// 首轮等满一个完整周期再跑,平时节奏不变。
|
||||
initialDelay = j.interval
|
||||
}
|
||||
go s.loopWithInitialDelay(ctx, j, initialDelay)
|
||||
helper.Go(s.log, "scheduler.loop."+j.name, func() { s.loopWithInitialDelay(ctx, j, initialDelay) })
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
)
|
||||
|
||||
// JobStatus is a snapshot suitable for the admin UI.
|
||||
@@ -142,7 +144,8 @@ func (s *SchedulerService) beginRun(j *scheduledJob) error {
|
||||
}
|
||||
|
||||
func (s *SchedulerService) runReserved(ctx context.Context, j *scheduledJob) error {
|
||||
err := j.run(ctx)
|
||||
// 任务 panic 转为 error,保证下方 running/lastErr 状态照常清理、调度循环存活。
|
||||
err := helper.Recover(s.log, "scheduler.job."+j.name, func() error { return j.run(ctx) })
|
||||
s.mu.Lock()
|
||||
j.lastRun = s.currentTime()
|
||||
if err != nil {
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
)
|
||||
|
||||
@@ -34,6 +35,13 @@ func (s *ScraperService) Start(ctx context.Context) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
// 启动自愈:进程中断遗留的 running 任务重置为 pending,否则永久卡死
|
||||
// (ClaimPending 只认 pending,重试按钮也拒绝 running)。
|
||||
if n, err := s.repo.ScrapeTask.ResetRunningToPending(ctx); err == nil && n > 0 && s.log != nil {
|
||||
s.log.Warn("scrape tasks reset from running to pending after restart", zap.Int64("count", n))
|
||||
} else if err != nil && s.log != nil {
|
||||
s.log.Warn("reset running scrape tasks failed", zap.Error(err))
|
||||
}
|
||||
go s.queueWorker(ctx)
|
||||
}
|
||||
|
||||
@@ -68,18 +76,33 @@ func (s *ScraperService) queueWorker(ctx context.Context) {
|
||||
defer wg.Done()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
// 任务已被 Claim 置为 running:停机前回写 pending,
|
||||
// 避免留下永久卡死的任务。
|
||||
s.requeueClaimedScrapeTask(t)
|
||||
return
|
||||
case sem <- struct{}{}:
|
||||
}
|
||||
defer func() { <-sem }()
|
||||
|
||||
s.processScrapeTask(ctx, t)
|
||||
// 刮削要解析远端元数据响应,单个任务 panic 不应拖垮队列 worker。
|
||||
helper.Run(s.log, "scraper.task", func() {
|
||||
s.processScrapeTask(ctx, t)
|
||||
})
|
||||
}(&tasks[i])
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
}
|
||||
|
||||
// requeueClaimedScrapeTask 把已认领但未开始执行的任务回写为 pending。
|
||||
func (s *ScraperService) requeueClaimedScrapeTask(t *model.ScrapeTask) {
|
||||
t.Status = model.ScrapeTaskPending
|
||||
t.StartedAt = nil
|
||||
if err := s.repo.ScrapeTask.Update(context.Background(), t); err != nil && s.log != nil {
|
||||
s.log.Warn("requeue claimed scrape task failed", zap.Error(err), zap.String("id", t.ID))
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ScraperService) processScrapeTask(ctx context.Context, task *model.ScrapeTask) {
|
||||
media, err := s.repo.Media.FindByID(ctx, task.MediaID)
|
||||
if err != nil || media == nil {
|
||||
@@ -216,8 +239,22 @@ func (s *ScraperService) EnqueueLibrary(ctx context.Context, libraryID string, o
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// 去重:排除已有 pending/running 任务的媒体,防止"先单集入队再点
|
||||
// 整库刮削"产生重复任务并被并发双刮(同一 Media 行被并发写两次)。
|
||||
mediaIDs := make([]string, 0, len(rows))
|
||||
for _, m := range rows {
|
||||
mediaIDs = append(mediaIDs, m.ID)
|
||||
}
|
||||
activeByMedia, err := s.repo.ScrapeTask.FindActiveByMediaIDs(ctx, mediaIDs)
|
||||
if err != nil {
|
||||
activeByMedia = nil // 去重查询失败不阻塞入队,仅退化为不去重
|
||||
}
|
||||
|
||||
tasks := make([]model.ScrapeTask, 0, len(rows))
|
||||
for _, m := range rows {
|
||||
if activeByMedia != nil && activeByMedia[m.ID] {
|
||||
continue
|
||||
}
|
||||
tasks = append(tasks, model.ScrapeTask{
|
||||
MediaID: m.ID,
|
||||
LibraryID: lib.ID,
|
||||
@@ -306,7 +343,16 @@ func (s *ScraperService) RetryScrapeTask(ctx context.Context, id string) error {
|
||||
return errors.New("刮削任务不存在")
|
||||
}
|
||||
if task.Status != model.ScrapeTaskFailed && task.Status != model.ScrapeTaskCanceled {
|
||||
return errors.New("只有失败或已取消的任务可以重试")
|
||||
// running 任务仅在其长时间无进展(>1h)时允许重试,作为卡死
|
||||
// 任务的逃生通道;正常执行中的任务仍拒绝重试以防双跑。
|
||||
stale := task.Status == model.ScrapeTaskRunning &&
|
||||
(task.StartedAt == nil || time.Since(*task.StartedAt) > time.Hour)
|
||||
if !stale {
|
||||
return errors.New("只有失败或已取消的任务可以重试")
|
||||
}
|
||||
if s.log != nil {
|
||||
s.log.Warn("retrying stuck running scrape task", zap.String("id", id))
|
||||
}
|
||||
}
|
||||
task.Status = model.ScrapeTaskPending
|
||||
task.Error = ""
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
|
||||
@@ -96,13 +97,14 @@ func (c *Container) Boot() {
|
||||
if err := c.APIConfig.SeedDefaults(c.stopCtx); err != nil {
|
||||
c.Log.Warn("api config seed failed", zap.Error(err))
|
||||
}
|
||||
go c.warmMediaSearchIndex(c.stopCtx)
|
||||
helper.Go(c.Log, "service.warmMediaSearchIndex", func() { c.warmMediaSearchIndex(c.stopCtx) })
|
||||
|
||||
// 启动调度器定时任务
|
||||
c.Scheduler.Start(c.stopCtx)
|
||||
|
||||
// 远程 Emby 挂载兼容迁移:旧账号无挂载时自动全量挂载
|
||||
// 远程 Emby 挂载兼容迁移:清理已删账号的残留挂载;旧账号无挂载时自动全量挂载
|
||||
if c.EmbyRemote != nil {
|
||||
c.EmbyRemote.CleanupOrphanMounts(c.stopCtx)
|
||||
c.EmbyRemote.AutoSeedMounts(c.stopCtx)
|
||||
}
|
||||
|
||||
@@ -119,7 +121,7 @@ func (c *Container) Boot() {
|
||||
// Mgo 保号规则巡检:默认关闭,由管理员通过 Telegram Bot 命令开启。
|
||||
// 每天触发一次评估;规则里的窗口可随机,不固定。
|
||||
if c.Device != nil {
|
||||
go c.runInactivitySweeper(c.stopCtx)
|
||||
helper.Go(c.Log, "service.inactivitySweeper", func() { c.runInactivitySweeper(c.stopCtx) })
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
@@ -47,12 +48,12 @@ func newServiceContainer(cfg *config.Config, log *zap.Logger, repos *repository.
|
||||
|
||||
func (b *serviceContainerBuilder) startRealtimeServices() {
|
||||
b.c.WSHub = NewHub(b.log)
|
||||
go b.c.WSHub.Run()
|
||||
helper.Go(b.log, "ws.hub", b.c.WSHub.Run)
|
||||
b.c.Tasks = NewTaskTrackerService(b.log, b.c.WSHub)
|
||||
b.c.SystemUpdate = NewSystemUpdateService(b.cfg, b.log, b.repos, b.c.Tasks, b.version)
|
||||
|
||||
b.c.SSEHub = NewSSEHub(b.log)
|
||||
go b.c.SSEHub.Run()
|
||||
helper.Go(b.log, "sse.hub", b.c.SSEHub.Run)
|
||||
}
|
||||
|
||||
func (b *serviceContainerBuilder) initProviderServices() {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -371,6 +372,16 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) {
|
||||
if err != nil || cfg["access_token"] == "" || cfg["refresh_token"] == "" {
|
||||
continue
|
||||
}
|
||||
// 临期才刷:115 refresh_token 是一次性轮转,无条件周期刷新会与
|
||||
// 运行中任务的自动刷新互相作废对方的凭据。24h 内刷新过(含运行
|
||||
// 中回调落库)就跳过;access_token 有效期远长于 24h。
|
||||
if last := strings.TrimSpace(cfg["token_refreshed_at"]); last != "" {
|
||||
if ts, perr := strconv.ParseInt(last, 10, 64); perr == nil {
|
||||
if time.Since(time.Unix(ts, 0)) < 24*time.Hour {
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
client := cloud115.NewOpenClient(cfg["app_id"], cfg["access_token"], cfg["refresh_token"])
|
||||
token, err := client.RefreshToken(cfg["refresh_token"])
|
||||
if err != nil {
|
||||
@@ -395,6 +406,7 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) {
|
||||
}
|
||||
cfg["access_token"] = s.crypto.Encrypt(token.AccessToken)
|
||||
cfg["refresh_token"] = s.crypto.Encrypt(token.RefreshToken)
|
||||
cfg["token_refreshed_at"] = strconv.FormatInt(time.Now().Unix(), 10)
|
||||
enc, err := s.strmAccountConfigJSON(cfg, false)
|
||||
if err != nil {
|
||||
continue
|
||||
@@ -410,6 +422,37 @@ func (s *StrmService) refresh115TokensOnce(ctx context.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
// persist115Tokens 把运行中任务自动刷新得到的新令牌加密写回账号配置,
|
||||
// 并记录刷新时间供定时刷新线程做临期判断。
|
||||
func (s *StrmService) persist115Tokens(accountID, accessToken, refreshToken string) {
|
||||
ctx := context.Background()
|
||||
acct, err := s.repo.StrmAccount.FindByID(ctx, accountID)
|
||||
if err != nil || acct == nil {
|
||||
return
|
||||
}
|
||||
cfg, err := s.strmAccountConfig(acct)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
cfg["access_token"] = s.crypto.Encrypt(accessToken)
|
||||
cfg["refresh_token"] = s.crypto.Encrypt(refreshToken)
|
||||
cfg["token_refreshed_at"] = strconv.FormatInt(time.Now().Unix(), 10)
|
||||
enc, err := s.strmAccountConfigJSON(cfg, false)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
acct.Config = enc
|
||||
now := time.Now()
|
||||
acct.LastTestAt = &now
|
||||
acct.LastTestResult = "ok"
|
||||
acct.LastTestOK = true
|
||||
if err := s.repo.StrmAccount.Update(ctx, acct); err != nil {
|
||||
s.log.Warn("persist refreshed 115 token failed", zap.Error(err), zap.String("account", acct.Name))
|
||||
return
|
||||
}
|
||||
s.log.Info("115 token refreshed and persisted", zap.String("account", acct.Name))
|
||||
}
|
||||
|
||||
// sync115RelayKey 把设置里的中继密钥同步给 cloud115(启动与设置保存时调用)。
|
||||
func (s *StrmService) sync115RelayKey(ctx context.Context) {
|
||||
cloud115.RelayEncryptionKey = s.strmSetting(ctx, Strm115RelayKeySetting)
|
||||
|
||||
@@ -5,7 +5,11 @@ import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
)
|
||||
|
||||
func TestStrmAccountConfigPreviewOf(t *testing.T) {
|
||||
@@ -67,6 +71,117 @@ func TestUpdateStrmAccountMergesConfigWithoutClearingSecrets(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteStrmAccountCascadesEmbyMounts(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newServiceTestDB(t, &model.StrmAccount{}, &model.EmbyMount{}, &model.StrmSyncPath{})
|
||||
repos := repository.New(db)
|
||||
svc := NewStrmService(nil, zap.NewNop(), repos, nil)
|
||||
|
||||
createEmbyAcct := func(id, name string) *model.StrmAccount {
|
||||
acct := &model.StrmAccount{
|
||||
Base: model.Base{ID: id},
|
||||
Name: name,
|
||||
Provider: model.StrmProviderEmbyRemote,
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.StrmAccount.Create(ctx, acct); err != nil {
|
||||
t.Fatalf("create account %s: %v", id, err)
|
||||
}
|
||||
return acct
|
||||
}
|
||||
createMounts := func(accountID string, viewIDs ...string) {
|
||||
for _, vid := range viewIDs {
|
||||
m := &model.EmbyMount{
|
||||
AccountID: accountID,
|
||||
RemoteViewID: vid,
|
||||
RemoteViewName: "库-" + vid,
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.EmbyMount.Create(ctx, m); err != nil {
|
||||
t.Fatalf("create mount %s: %v", vid, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
gone := createEmbyAcct("acct-gone", "要删除的账号")
|
||||
createMounts(gone.ID, "view-1", "view-2", "view-3")
|
||||
keep := createEmbyAcct("acct-keep", "保留的账号")
|
||||
createMounts(keep.ID, "view-a")
|
||||
|
||||
if err := svc.DeleteStrmAccount(ctx, gone.ID); err != nil {
|
||||
t.Fatalf("DeleteStrmAccount failed: %v", err)
|
||||
}
|
||||
|
||||
remaining, err := repos.StrmAccount.FindByID(ctx, gone.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("find account: %v", err)
|
||||
}
|
||||
if remaining != nil {
|
||||
t.Fatalf("account %s should be deleted", gone.ID)
|
||||
}
|
||||
goneMounts, err := repos.EmbyMount.ListByAccountID(ctx, gone.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list mounts: %v", err)
|
||||
}
|
||||
if len(goneMounts) != 0 {
|
||||
t.Fatalf("deleted account still has %d mounts (orphans)", len(goneMounts))
|
||||
}
|
||||
keepMounts, err := repos.EmbyMount.ListByAccountID(ctx, keep.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list mounts: %v", err)
|
||||
}
|
||||
if len(keepMounts) != 1 || keepMounts[0].RemoteViewID != "view-a" {
|
||||
t.Fatalf("keep account mounts = %#v, want 1 (view-a)", keepMounts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupOrphanMountsRemovesStaleRows(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newServiceTestDB(t, &model.StrmAccount{}, &model.EmbyMount{})
|
||||
repos := repository.New(db)
|
||||
svc := NewEmbyRemoteService(&config.Config{}, zap.NewNop(), repos, nil)
|
||||
|
||||
acct := &model.StrmAccount{
|
||||
Base: model.Base{ID: "acct-1"},
|
||||
Name: "emby",
|
||||
Provider: model.StrmProviderEmbyRemote,
|
||||
Enabled: true,
|
||||
}
|
||||
if err := repos.StrmAccount.Create(ctx, acct); err != nil {
|
||||
t.Fatalf("create account: %v", err)
|
||||
}
|
||||
for i, vid := range []string{"v1", "v2", "v-orphan-1", "v-orphan-2"} {
|
||||
m := &model.EmbyMount{
|
||||
Base: model.Base{ID: "mount-" + vid},
|
||||
AccountID: acct.ID,
|
||||
RemoteViewID: vid,
|
||||
Enabled: true,
|
||||
}
|
||||
if i >= 2 {
|
||||
// 模拟历史残留:挂载归属不存在的账号
|
||||
m.AccountID = "no-such-account"
|
||||
}
|
||||
if err := repos.EmbyMount.Create(ctx, m); err != nil {
|
||||
t.Fatalf("create mount %s: %v", vid, err)
|
||||
}
|
||||
}
|
||||
|
||||
svc.CleanupOrphanMounts(ctx)
|
||||
|
||||
left, err := repos.EmbyMount.List(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("list mounts: %v", err)
|
||||
}
|
||||
if len(left) != 2 {
|
||||
t.Fatalf("mounts after cleanup = %d, want 2 (orphans removed)", len(left))
|
||||
}
|
||||
for _, m := range left {
|
||||
if m.AccountID != acct.ID {
|
||||
t.Fatalf("mount %s still orphan (account %s)", m.ID, m.AccountID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func mustJSON(v any) string {
|
||||
data, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud115"
|
||||
@@ -66,7 +67,18 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
|
||||
return
|
||||
}
|
||||
defer s.releaseDownloadSlot(task.Provider)
|
||||
s.processDownloadTask(ctx, task)
|
||||
// 单个任务 panic 不应拖垮整个下载 worker,且 panic 时任务
|
||||
// 会永远停在 running:兜底走失败重试路径。
|
||||
completed := false
|
||||
helper.Run(s.log, "strm.downloadTask", func() {
|
||||
defer func() {
|
||||
if !completed {
|
||||
s.downloadTaskFailWithRetry(task, "任务执行异常中断")
|
||||
}
|
||||
}()
|
||||
s.processDownloadTask(ctx, task)
|
||||
completed = true
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
@@ -74,9 +86,19 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
|
||||
}
|
||||
|
||||
// requeueDownloadTask 把已认领但未实际执行的任务退回 pending,避免长期停留在 running。
|
||||
// 退回时必须设置 NextTryAt(WAF 冷却剩余时间):claim 只过滤 next_try_at
|
||||
// 已过期的任务,不设会让同一批任务被立刻再认领,形成 claim/requeue
|
||||
// 热循环(占用 SQLite 写锁并饿死上传队列)。
|
||||
func (s *StrmService) requeueDownloadTask(task *model.StrmDownloadTask) {
|
||||
task.Status = model.StrmTaskPending
|
||||
task.StartedAt = nil
|
||||
task.NextTryAt = nil
|
||||
if task.Provider == model.StrmProvider115 {
|
||||
if left := s.wafCooldownLeft(); left > 0 {
|
||||
next := time.Now().Add(left)
|
||||
task.NextTryAt = &next
|
||||
}
|
||||
}
|
||||
if err := s.repo.StrmDownload.Update(context.Background(), task); err != nil {
|
||||
s.log.Warn("requeue strm download task failed", zap.Error(err), zap.String("id", task.ID))
|
||||
}
|
||||
@@ -93,8 +115,16 @@ func (s *StrmService) processDownloadTask(ctx context.Context, task *model.StrmD
|
||||
task.Status = status
|
||||
task.Error = message
|
||||
task.FinishedAt = &now
|
||||
if err := s.repo.StrmDownload.Update(context.Background(), task); err != nil {
|
||||
// 条件化收尾:用户取消会直接把 running 改为 canceled,无条件
|
||||
// Update 会把已取消任务覆盖回 done。
|
||||
if ok, err := s.repo.StrmDownload.UpdateIfRunning(context.Background(), task.ID, map[string]any{
|
||||
"status": status,
|
||||
"error": message,
|
||||
"finished_at": &now,
|
||||
}); err != nil {
|
||||
s.log.Warn("update strm download task failed", zap.Error(err))
|
||||
} else if !ok {
|
||||
s.log.Info("strm download task already closed elsewhere", zap.String("id", task.ID))
|
||||
}
|
||||
}
|
||||
acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID)
|
||||
@@ -147,7 +177,19 @@ func (s *StrmService) uploadWorker(ctx context.Context) {
|
||||
continue
|
||||
}
|
||||
for i := range tasks {
|
||||
s.processUploadTask(ctx, &tasks[i])
|
||||
t := &tasks[i]
|
||||
// 与下载侧一致:单任务 panic 不损失 worker 线程,且兜底走
|
||||
// 失败重试路径(否则任务永久 running)。
|
||||
completed := false
|
||||
helper.Run(s.log, "strm.uploadTask", func() {
|
||||
defer func() {
|
||||
if !completed {
|
||||
s.uploadTaskFailWithRetry(t, "任务执行异常中断")
|
||||
}
|
||||
}()
|
||||
s.processUploadTask(ctx, t)
|
||||
completed = true
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -158,8 +200,15 @@ func (s *StrmService) processUploadTask(ctx context.Context, task *model.StrmUpl
|
||||
task.Status = status
|
||||
task.Error = message
|
||||
task.FinishedAt = &now
|
||||
if err := s.repo.StrmUpload.Update(context.Background(), task); err != nil {
|
||||
// 条件化收尾:与下载侧一致,防止覆盖已取消任务。
|
||||
if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{
|
||||
"status": status,
|
||||
"error": message,
|
||||
"finished_at": &now,
|
||||
}); err != nil {
|
||||
s.log.Warn("update strm upload task failed", zap.Error(err))
|
||||
} else if !ok {
|
||||
s.log.Info("strm upload task already closed elsewhere", zap.String("id", task.ID))
|
||||
}
|
||||
}
|
||||
if task.Provider == model.StrmProvider115 {
|
||||
@@ -209,8 +258,15 @@ func (s *StrmService) processUpload115(ctx context.Context, task *model.StrmUplo
|
||||
task.Status = status
|
||||
task.Error = message
|
||||
task.FinishedAt = &now
|
||||
if err := s.repo.StrmUpload.Update(context.Background(), task); err != nil {
|
||||
// 条件化收尾:与下载侧一致,防止覆盖已取消任务。
|
||||
if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{
|
||||
"status": status,
|
||||
"error": message,
|
||||
"finished_at": &now,
|
||||
}); err != nil {
|
||||
s.log.Warn("update strm upload task failed", zap.Error(err))
|
||||
} else if !ok {
|
||||
s.log.Info("strm upload task already closed elsewhere", zap.String("id", task.ID))
|
||||
}
|
||||
}
|
||||
acct, err := s.repo.StrmAccount.FindByID(ctx, task.AccountID)
|
||||
@@ -249,7 +305,19 @@ func (s *StrmService) downloadTaskFailWithRetry(task *model.StrmDownloadTask, me
|
||||
if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) {
|
||||
return
|
||||
}
|
||||
_ = s.repo.StrmDownload.Update(context.Background(), task)
|
||||
// 条件化写入:任务已被取消(DB 中不再是 running)时不得覆盖回 pending,
|
||||
// 否则用户刚取消的任务会“复活”并自动重试。
|
||||
if ok, err := s.repo.StrmDownload.UpdateIfRunning(context.Background(), task.ID, map[string]any{
|
||||
"status": task.Status,
|
||||
"error": task.Error,
|
||||
"retry_count": task.RetryCount,
|
||||
"next_try_at": task.NextTryAt,
|
||||
"finished_at": task.FinishedAt,
|
||||
}); err != nil {
|
||||
s.log.Warn("fail strm download task failed", zap.Error(err), zap.String("id", task.ID))
|
||||
} else if !ok {
|
||||
s.log.Info("strm download task already closed elsewhere, skip retry overwrite", zap.String("id", task.ID))
|
||||
}
|
||||
}
|
||||
|
||||
// uploadTaskFailWithRetry 上传失败任务按退避重试,超过上限标记 failed。
|
||||
@@ -257,7 +325,17 @@ func (s *StrmService) uploadTaskFailWithRetry(task *model.StrmUploadTask, messag
|
||||
if !retryTask(&task.RetryCount, &task.Status, &task.Error, &task.NextTryAt, &task.FinishedAt, message) {
|
||||
return
|
||||
}
|
||||
_ = s.repo.StrmUpload.Update(context.Background(), task)
|
||||
if ok, err := s.repo.StrmUpload.UpdateIfRunning(context.Background(), task.ID, map[string]any{
|
||||
"status": task.Status,
|
||||
"error": task.Error,
|
||||
"retry_count": task.RetryCount,
|
||||
"next_try_at": task.NextTryAt,
|
||||
"finished_at": task.FinishedAt,
|
||||
}); err != nil {
|
||||
s.log.Warn("fail strm upload task failed", zap.Error(err), zap.String("id", task.ID))
|
||||
} else if !ok {
|
||||
s.log.Info("strm upload task already closed elsewhere, skip retry overwrite", zap.String("id", task.ID))
|
||||
}
|
||||
}
|
||||
|
||||
// retryTask 失败状态机:重试次数不足则回 pending 并设置退避时间,否则 failed。
|
||||
|
||||
@@ -24,9 +24,11 @@ import (
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud115"
|
||||
)
|
||||
|
||||
// strm 全局设置键(存于 Setting 表,strm.* 前缀)。
|
||||
@@ -91,8 +93,8 @@ type StrmService struct {
|
||||
oauthSessions map[string]*strm115AuthSession
|
||||
wafUntil time.Time // 115 风控/限流熔断截止时间(由 mu 保护)
|
||||
|
||||
downloadSem115 chan struct{} // 115 换直链+下载并发上限(风控兜底)
|
||||
downloadSemDAV chan struct{} // WebDAV/OpenList/CloudDrive2 元数据下载并发上限
|
||||
downloadSem115 chan struct{} // 115 换直链+下载并发上限(风控兜底)
|
||||
downloadSemDAV chan struct{} // WebDAV/OpenList/CloudDrive2 元数据下载并发上限
|
||||
downloadSemOnce sync.Once
|
||||
}
|
||||
|
||||
@@ -168,6 +170,11 @@ func NewStrmService(cfg *config.Config, log *zap.Logger, repos *repository.Conta
|
||||
|
||||
// Start 启动下载/上传队列 worker、定时同步巡检、115 token 刷新与队列清理。
|
||||
func (s *StrmService) Start(ctx context.Context) {
|
||||
// baseCtx 挂到服务生命周期 ctx 上(Start 由启动流程传入 stopCtx):
|
||||
// 此前硬编码 context.Background(),Stop() 关 stopCh 后 worker 会退出,
|
||||
// 但进行中的全量同步(可能持续数小时)完全不受停机控制,优雅停机
|
||||
// 窗口内仍在批量写库/写盘。
|
||||
s.baseCtx = ctx
|
||||
s.sync115RelayKey(ctx)
|
||||
s.recoverInterruptedSyncs(ctx)
|
||||
downloadThreads := s.strmIntSetting(ctx, StrmSettingDownloadThreads, 6)
|
||||
@@ -186,14 +193,26 @@ func (s *StrmService) Start(ctx context.Context) {
|
||||
uploadThreads = 4
|
||||
}
|
||||
for i := 0; i < downloadThreads; i++ {
|
||||
go s.downloadWorker(ctx)
|
||||
helper.Go(s.log, "strm.downloadWorker", func() { s.downloadWorker(ctx) })
|
||||
}
|
||||
for i := 0; i < uploadThreads; i++ {
|
||||
go s.uploadWorker(ctx)
|
||||
helper.Go(s.log, "strm.uploadWorker", func() { s.uploadWorker(ctx) })
|
||||
}
|
||||
go s.cronLoop(ctx)
|
||||
go s.queueCleanupLoop(ctx)
|
||||
go s.refresh115TokensLoop(ctx)
|
||||
// 队列任务自愈:进程崩溃/停机遗留的 running 任务重置为 pending,
|
||||
// 否则永久卡死并会通过 GetActiveLocalPathMap 阻塞该文件的重复下载。
|
||||
if n, err := s.repo.StrmDownload.ResetRunningToPending(ctx); err == nil && n > 0 {
|
||||
s.log.Warn("strm download tasks reset from running to pending after restart", zap.Int64("count", n))
|
||||
} else if err != nil {
|
||||
s.log.Warn("reset running strm download tasks failed", zap.Error(err))
|
||||
}
|
||||
if n, err := s.repo.StrmUpload.ResetRunningToPending(ctx); err == nil && n > 0 {
|
||||
s.log.Warn("strm upload tasks reset from running to pending after restart", zap.Int64("count", n))
|
||||
} else if err != nil {
|
||||
s.log.Warn("reset running strm upload tasks failed", zap.Error(err))
|
||||
}
|
||||
helper.Go(s.log, "strm.cronLoop", func() { s.cronLoop(ctx) })
|
||||
helper.Go(s.log, "strm.queueCleanupLoop", func() { s.queueCleanupLoop(ctx) })
|
||||
helper.Go(s.log, "strm.refresh115TokensLoop", func() { s.refresh115TokensLoop(ctx) })
|
||||
s.log.Info("strm service started",
|
||||
zap.Int("download_threads", downloadThreads),
|
||||
zap.Int("upload_threads", uploadThreads))
|
||||
@@ -378,13 +397,13 @@ func (s *StrmService) UpdateStrmAccount(ctx context.Context, id, name string, en
|
||||
if enabled != nil {
|
||||
acct.Enabled = *enabled
|
||||
}
|
||||
if len(config) > 0 {
|
||||
enc, err := s.mergeStrmAccountConfig(acct.Config, config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
acct.Config = enc
|
||||
if len(config) > 0 {
|
||||
enc, err := s.mergeStrmAccountConfig(acct.Config, config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
acct.Config = enc
|
||||
}
|
||||
if err := s.repo.StrmAccount.Update(ctx, acct); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -405,6 +424,11 @@ func (s *StrmService) DeleteStrmAccount(ctx context.Context, id string) error {
|
||||
if err := s.repo.StrmAccount.Delete(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
// 级联清理远程 Emby 挂载:否则留下孤儿挂载,挂载计数/列表仍会显示。
|
||||
// 账号已删,挂载清理失败只记日志,不让删除请求报错。
|
||||
if _, err := s.repo.EmbyMount.DeleteByAccountID(ctx, id); err != nil && s.log != nil {
|
||||
s.log.Warn("delete emby mounts for account failed", zap.String("account", id), zap.Error(err))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -451,6 +475,15 @@ func (s *StrmService) providerFor(ctx context.Context, acct *model.StrmAccount)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 115 开放平台:运行中自动刷新得到的新令牌必须落库。否则长任务
|
||||
// 里的新 token 只存在于内存,定时刷新线程又用 DB 里的旧
|
||||
// refresh_token 再刷(一次性轮转),两者互相作废,最终把有效账号
|
||||
// 标成“授权已失效”。
|
||||
if oc, ok := provider.(interface{ OpenClient() *cloud115.OpenClient }); ok {
|
||||
oc.OpenClient().OnTokenRefreshed = func(accessToken, refreshToken string) {
|
||||
s.persist115Tokens(acct.ID, accessToken, refreshToken)
|
||||
}
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
|
||||
+208
-123
@@ -6,6 +6,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -18,6 +19,7 @@ import (
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/helper"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud"
|
||||
"github.com/truewhile/MeBox/internal/service/cloud115"
|
||||
@@ -96,7 +98,8 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...
|
||||
StartedAt: &now,
|
||||
}
|
||||
if err := s.repo.StrmSyncRecord.Create(ctx, rec); err != nil {
|
||||
s.clearRunning(pathID)
|
||||
s.clearRunning(pathID, cancel)
|
||||
cancel()
|
||||
return err
|
||||
}
|
||||
status := model.StrmSyncRecordRunning
|
||||
@@ -105,7 +108,7 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...
|
||||
p.LastSyncMessage = "同步进行中"
|
||||
_ = s.repo.StrmSyncPath.Update(ctx, p)
|
||||
|
||||
go s.runSync(runCtx, p, rec)
|
||||
helper.Go(s.log, "strm.sync", func() { s.runSync(runCtx, p, rec, cancel) })
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -113,11 +116,10 @@ func (s *StrmService) StartSync(ctx context.Context, pathID string, syncType ...
|
||||
func (s *StrmService) CancelSync(ctx context.Context, pathID string) error {
|
||||
s.mu.Lock()
|
||||
cancel, exists := s.running[pathID]
|
||||
if exists {
|
||||
delete(s.running, pathID)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
// 不在这里预删 running 标记:runSync 退出时的 clearRunning 会按
|
||||
// cancel 身份校验后删除,避免旧同步收尾误删新同步的标记。
|
||||
if exists && cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
@@ -141,12 +143,24 @@ func (s *StrmService) IsSyncRunning(pathID string) bool {
|
||||
return exists
|
||||
}
|
||||
|
||||
func (s *StrmService) clearRunning(pathID string) {
|
||||
// clearRunning 清除同步的运行标记;仅当 map 中登记的 cancel 与本次同步
|
||||
// 一致时才删除,防止慢收尾的旧同步把随后启动的新同步标记误删掉。
|
||||
func (s *StrmService) clearRunning(pathID string, cancel context.CancelFunc) {
|
||||
s.mu.Lock()
|
||||
delete(s.running, pathID)
|
||||
if cur, ok := s.running[pathID]; ok {
|
||||
if cancel == nil || cur == nil || sameCancelFunc(cur, cancel) {
|
||||
delete(s.running, pathID)
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// sameCancelFunc 比较两个 cancel 是否为同一实例(每次 WithCancel 返回
|
||||
// 独立闭包,函数指针即身份)。约定 running 表只登记 StartSync 的 cancel。
|
||||
func sameCancelFunc(a, b context.CancelFunc) bool {
|
||||
return reflect.ValueOf(a).Pointer() == reflect.ValueOf(b).Pointer()
|
||||
}
|
||||
|
||||
// ListRemoteDir 列出网盘账号某目录下的条目(供前端目录选择器使用)。
|
||||
func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string) ([]cloud.FileEntry, error) {
|
||||
acct, err := s.repo.StrmAccount.FindByID(ctx, accountID)
|
||||
@@ -168,8 +182,9 @@ func (s *StrmService) ListRemoteDir(ctx context.Context, accountID, dir string)
|
||||
}
|
||||
|
||||
// runSync 执行同步主体;结束时更新记录与目录状态。
|
||||
func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord) {
|
||||
defer s.clearRunning(p.ID)
|
||||
func (s *StrmService) runSync(ctx context.Context, p *model.StrmSyncPath, rec *model.StrmSyncRecord, cancel context.CancelFunc) {
|
||||
defer s.clearRunning(p.ID, cancel)
|
||||
defer cancel()
|
||||
|
||||
cfg, err := s.strmEffectiveConfig(ctx, p)
|
||||
if err != nil {
|
||||
@@ -333,24 +348,34 @@ func (st *strmSyncState) walkRemote() error {
|
||||
ctx, cancel := context.WithCancel(st.ctx)
|
||||
defer cancel()
|
||||
|
||||
queue := make(chan dirTask, 512)
|
||||
var pending atomic.Int64
|
||||
// 工作队列用「互斥锁 + 条件变量 + 动态 slice」实现,而不是有界
|
||||
// channel:有界缓冲下所有 worker 可能同时阻塞在发送上、无人接收,
|
||||
// closer 又在等 pending 归零,形成永久死锁。push 永不阻塞即可保证
|
||||
// 有进度就一定有推进。
|
||||
// pending 计数 = 尚未处理完的任务数(在 work 里或正在被 List)。
|
||||
var (
|
||||
walkMu sync.Mutex
|
||||
walkCond = sync.NewCond(&walkMu)
|
||||
work []dirTask
|
||||
pending int
|
||||
)
|
||||
push := func(t dirTask) {
|
||||
walkMu.Lock()
|
||||
work = append(work, t)
|
||||
pending++
|
||||
walkCond.Signal()
|
||||
walkMu.Unlock()
|
||||
}
|
||||
// ctx 取消时唤醒所有等待中的 worker 让其退出。
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
walkMu.Lock()
|
||||
walkCond.Broadcast()
|
||||
walkMu.Unlock()
|
||||
}()
|
||||
|
||||
// 根目录入队
|
||||
pending.Add(1)
|
||||
queue <- dirTask{id: root, rel: ""}
|
||||
|
||||
// 当队列中所有目录都被消费(pending 归零)或出错时关闭 channel,
|
||||
// 让 worker 全部退出。
|
||||
go func() {
|
||||
for {
|
||||
if ctx.Err() != nil || pending.Load() == 0 {
|
||||
close(queue)
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
}()
|
||||
push(dirTask{id: root, rel: ""})
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
@@ -361,38 +386,64 @@ func (st *strmSyncState) walkRemote() error {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for task := range queue {
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
entries, err := st.provider.List(ctx, task.id)
|
||||
if err != nil {
|
||||
errMu.Lock()
|
||||
if firstErr == nil {
|
||||
firstErr = fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err)
|
||||
}
|
||||
errMu.Unlock()
|
||||
cancel()
|
||||
return
|
||||
}
|
||||
for _, entry := range entries {
|
||||
cleanName := cleanEntryName(entry.Name, entry.IsDir)
|
||||
rel := cleanName
|
||||
if task.rel != "" {
|
||||
rel = task.rel + "/" + cleanName
|
||||
}
|
||||
if entry.IsDir {
|
||||
pending.Add(1)
|
||||
select {
|
||||
case queue <- dirTask{id: entry.ID, rel: rel}:
|
||||
case <-ctx.Done():
|
||||
pending.Add(-1)
|
||||
// worker 解析远端响应 panic 时取消整个同步,让其余 worker
|
||||
// 正常收尾;正常退出不取消。
|
||||
if err := helper.Recover(st.s.log, "strm.sync.walkRemote", func() error {
|
||||
for {
|
||||
walkMu.Lock()
|
||||
for len(work) == 0 {
|
||||
if ctx.Err() != nil || pending == 0 {
|
||||
walkMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
} else {
|
||||
st.processRemoteFile(entry, rel)
|
||||
walkCond.Wait()
|
||||
}
|
||||
task := work[0]
|
||||
work = work[1:]
|
||||
walkMu.Unlock()
|
||||
|
||||
if ctx.Err() != nil {
|
||||
walkMu.Lock()
|
||||
pending--
|
||||
walkCond.Broadcast()
|
||||
walkMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
entries, err := st.provider.List(ctx, task.id)
|
||||
if err != nil {
|
||||
errMu.Lock()
|
||||
if firstErr == nil {
|
||||
firstErr = fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err)
|
||||
}
|
||||
errMu.Unlock()
|
||||
walkMu.Lock()
|
||||
pending--
|
||||
walkCond.Broadcast()
|
||||
walkMu.Unlock()
|
||||
cancel()
|
||||
return nil
|
||||
}
|
||||
for _, entry := range entries {
|
||||
cleanName := cleanEntryName(entry.Name, entry.IsDir)
|
||||
rel := cleanName
|
||||
if task.rel != "" {
|
||||
rel = task.rel + "/" + cleanName
|
||||
}
|
||||
if entry.IsDir {
|
||||
push(dirTask{id: entry.ID, rel: rel})
|
||||
} else {
|
||||
st.processRemoteFile(entry, rel)
|
||||
}
|
||||
}
|
||||
walkMu.Lock()
|
||||
pending--
|
||||
if pending == 0 {
|
||||
walkCond.Broadcast()
|
||||
}
|
||||
walkMu.Unlock()
|
||||
}
|
||||
pending.Add(-1)
|
||||
}); err != nil {
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -560,22 +611,28 @@ func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for t := range taskCh {
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
files, _, err := open115.GetFsListFlat(ctx, rootCID, t.offset, pageSize)
|
||||
if err != nil {
|
||||
errMu.Lock()
|
||||
if fetchErr == nil {
|
||||
fetchErr = err
|
||||
// 分页拉取 panic 时取消整个同步;正常退出不取消。
|
||||
if err := helper.Recover(st.s.log, "strm.sync.walk115.page", func() error {
|
||||
for t := range taskCh {
|
||||
if ctx.Err() != nil {
|
||||
return nil
|
||||
}
|
||||
errMu.Unlock()
|
||||
return
|
||||
files, _, err := open115.GetFsListFlat(ctx, rootCID, t.offset, pageSize)
|
||||
if err != nil {
|
||||
errMu.Lock()
|
||||
if fetchErr == nil {
|
||||
fetchErr = err
|
||||
}
|
||||
errMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
filesMu.Lock()
|
||||
allFiles = append(allFiles, files...)
|
||||
filesMu.Unlock()
|
||||
}
|
||||
filesMu.Lock()
|
||||
allFiles = append(allFiles, files...)
|
||||
filesMu.Unlock()
|
||||
return nil
|
||||
}); err != nil {
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -632,63 +689,70 @@ func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error {
|
||||
pwg.Add(1)
|
||||
go func() {
|
||||
defer pwg.Done()
|
||||
for pid := range pidCh {
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
if _, loaded := st.dirCache.Load(pid); loaded {
|
||||
if n := doneDirs.Add(1); n%20 == 0 || n == int64(totalDirs) {
|
||||
// 解析目录详情 panic 时中止整个同步(避免带着损坏的相对路径
|
||||
// 继续执行);正常退出不取消。
|
||||
if err := helper.Recover(st.s.log, "strm.sync.walk115.dirTree", func() error {
|
||||
for pid := range pidCh {
|
||||
if ctx.Err() != nil {
|
||||
return nil
|
||||
}
|
||||
if _, loaded := st.dirCache.Load(pid); loaded {
|
||||
if n := doneDirs.Add(1); n%20 == 0 || n == int64(totalDirs) {
|
||||
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (%d/%d)...", n, totalDirs))
|
||||
}
|
||||
continue
|
||||
}
|
||||
detail, err := open115.GetFsDetailByCid(ctx, pid)
|
||||
if err != nil {
|
||||
// 目录详情解析失败会导致下游文件 rel 无法还原真实父路径,
|
||||
// seen key 与磁盘路径对不上:增量 prune 会误删本地文件、上传会
|
||||
// 误传本地未变文件、下载会重复下载。这里不是降级容错,而是
|
||||
// 直接中止整个同步——宁可本次同步失败,也不带着损坏的相对路径
|
||||
// 继续执行造成大规模误删/误传/重下(参考用户反馈"云盘没动却重下重传")。
|
||||
errMu.Lock()
|
||||
if firstErr == nil {
|
||||
firstErr = fmt.Errorf("115: 解析目录树失败(file_id=%s):%w", pid, err)
|
||||
}
|
||||
errMu.Unlock()
|
||||
st.scanIncomplete.Store(true)
|
||||
cancel()
|
||||
return nil
|
||||
} else if detail != nil {
|
||||
// 解析相对路径
|
||||
relPath := cleanDirRel(detail.RelativePath(rootCID))
|
||||
st.dirCache.Store(pid, relPath)
|
||||
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, pid, relPath)
|
||||
|
||||
// 顺便解析并缓存 detail.Paths 中包含的中间各层级目录
|
||||
for _, ancestor := range detail.Paths {
|
||||
if ancestor.FileId == "0" || ancestor.FileId == rootCID {
|
||||
continue
|
||||
}
|
||||
if _, loaded := st.dirCache.Load(ancestor.FileId); !loaded {
|
||||
subDetail := &cloud115.RemoteFileDetail{
|
||||
FileId: ancestor.FileId,
|
||||
FileName: ancestor.Name,
|
||||
Paths: nil,
|
||||
}
|
||||
for _, p := range detail.Paths {
|
||||
subDetail.Paths = append(subDetail.Paths, p)
|
||||
if p.FileId == ancestor.FileId {
|
||||
break
|
||||
}
|
||||
}
|
||||
ancestorRel := cleanDirRel(subDetail.RelativePath(rootCID))
|
||||
st.dirCache.Store(ancestor.FileId, ancestorRel)
|
||||
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, ancestor.FileId, ancestorRel)
|
||||
}
|
||||
}
|
||||
}
|
||||
if n := doneDirs.Add(1); n%10 == 0 || n == int64(totalDirs) {
|
||||
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (%d/%d)...", n, totalDirs))
|
||||
}
|
||||
continue
|
||||
}
|
||||
detail, err := open115.GetFsDetailByCid(ctx, pid)
|
||||
if err != nil {
|
||||
// 目录详情解析失败会导致下游文件 rel 无法还原真实父路径,
|
||||
// seen key 与磁盘路径对不上:增量 prune 会误删本地文件、上传会
|
||||
// 误传本地未变文件、下载会重复下载。这里不是降级容错,而是
|
||||
// 直接中止整个同步——宁可本次同步失败,也不带着损坏的相对路径
|
||||
// 继续执行造成大规模误删/误传/重下(参考用户反馈"云盘没动却重下重传")。
|
||||
errMu.Lock()
|
||||
if firstErr == nil {
|
||||
firstErr = fmt.Errorf("115: 解析目录树失败(file_id=%s):%w", pid, err)
|
||||
}
|
||||
errMu.Unlock()
|
||||
st.scanIncomplete.Store(true)
|
||||
cancel()
|
||||
return
|
||||
} else if detail != nil {
|
||||
// 解析相对路径
|
||||
relPath := cleanDirRel(detail.RelativePath(rootCID))
|
||||
st.dirCache.Store(pid, relPath)
|
||||
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, pid, relPath)
|
||||
|
||||
// 顺便解析并缓存 detail.Paths 中包含的中间各层级目录
|
||||
for _, ancestor := range detail.Paths {
|
||||
if ancestor.FileId == "0" || ancestor.FileId == rootCID {
|
||||
continue
|
||||
}
|
||||
if _, loaded := st.dirCache.Load(ancestor.FileId); !loaded {
|
||||
subDetail := &cloud115.RemoteFileDetail{
|
||||
FileId: ancestor.FileId,
|
||||
FileName: ancestor.Name,
|
||||
Paths: nil,
|
||||
}
|
||||
for _, p := range detail.Paths {
|
||||
subDetail.Paths = append(subDetail.Paths, p)
|
||||
if p.FileId == ancestor.FileId {
|
||||
break
|
||||
}
|
||||
}
|
||||
ancestorRel := cleanDirRel(subDetail.RelativePath(rootCID))
|
||||
st.dirCache.Store(ancestor.FileId, ancestorRel)
|
||||
_ = st.s.repo.StrmDirCache.Set(ctx, st.p.ID, ancestor.FileId, ancestorRel)
|
||||
}
|
||||
}
|
||||
}
|
||||
if n := doneDirs.Add(1); n%10 == 0 || n == int64(totalDirs) {
|
||||
st.updateSyncMessage(fmt.Sprintf("正在解析目录树 (%d/%d)...", n, totalDirs))
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -1308,6 +1372,10 @@ func (st *strmSyncState) updateSyncMessage(msg string) {
|
||||
func (s *StrmService) cronLoop(ctx context.Context) {
|
||||
ticker := time.NewTicker(60 * time.Second)
|
||||
defer ticker.Stop()
|
||||
// 记录上次检查到的分钟:一轮循环若被慢操作拖过 60s(远端 List 慢、
|
||||
// 串行 StartSync、DB 忙),ticker 会丢掉中间的 tick,命中排程的分钟
|
||||
// 若只按"当前分钟相等"判定就会被静默跳过。逐分钟回放补触。
|
||||
last := time.Now().Truncate(time.Minute)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
@@ -1315,8 +1383,18 @@ func (s *StrmService) cronLoop(ctx context.Context) {
|
||||
case <-s.stopCh:
|
||||
return
|
||||
case now := <-ticker.C:
|
||||
now = now.Truncate(time.Minute)
|
||||
paths, err := s.repo.StrmSyncPath.List(ctx)
|
||||
if err != nil {
|
||||
last = now
|
||||
continue
|
||||
}
|
||||
due := make([]time.Time, 0, 2)
|
||||
for m := last.Add(time.Minute); !m.After(now); m = m.Add(time.Minute) {
|
||||
due = append(due, m)
|
||||
}
|
||||
last = now
|
||||
if len(due) == 0 {
|
||||
continue
|
||||
}
|
||||
for i := range paths {
|
||||
@@ -1324,7 +1402,14 @@ func (s *StrmService) cronLoop(ctx context.Context) {
|
||||
if !p.Enabled || !p.EnableCron || strings.TrimSpace(p.Cron) == "" {
|
||||
continue
|
||||
}
|
||||
if !cronMatches(p.Cron, now) {
|
||||
matched := false
|
||||
for _, m := range due {
|
||||
if cronMatches(p.Cron, m) {
|
||||
matched = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !matched {
|
||||
continue
|
||||
}
|
||||
s.mu.Lock()
|
||||
|
||||
@@ -23,6 +23,8 @@ import (
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
@@ -35,6 +37,21 @@ type SubtitleService struct {
|
||||
log *zap.Logger
|
||||
repo *repository.Container
|
||||
cfg *config.Config
|
||||
|
||||
// 目录发现是 Emby 条目列表的热路径(每个媒体源一次 DB 查询 + 最多 5 次
|
||||
// os.ReadDir),而字幕文件极少变化:按 media_id 做短 TTL 缓存。
|
||||
cacheMu sync.Mutex
|
||||
discovery map[string]subtitleDiscoveryEntry
|
||||
}
|
||||
|
||||
const (
|
||||
subtitleDiscoveryTTL = 2 * time.Minute
|
||||
subtitleDiscoveryCacheCap = 4096
|
||||
)
|
||||
|
||||
type subtitleDiscoveryEntry struct {
|
||||
tracks []SubtitleTrack
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
// NewSubtitleService is the constructor.
|
||||
@@ -73,6 +90,50 @@ func (s *SubtitleService) DiscoverExternalOnly(ctx context.Context, mediaID stri
|
||||
}
|
||||
|
||||
func (s *SubtitleService) discover(ctx context.Context, mediaID string) ([]SubtitleTrack, error) {
|
||||
if tracks, ok := s.cachedDiscovery(mediaID); ok {
|
||||
return tracks, nil
|
||||
}
|
||||
tracks, err := s.discoverUncached(ctx, mediaID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.rememberDiscovery(mediaID, tracks)
|
||||
return tracks, nil
|
||||
}
|
||||
|
||||
func (s *SubtitleService) cachedDiscovery(mediaID string) ([]SubtitleTrack, bool) {
|
||||
now := time.Now()
|
||||
s.cacheMu.Lock()
|
||||
defer s.cacheMu.Unlock()
|
||||
entry, ok := s.discovery[mediaID]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if now.After(entry.expiresAt) {
|
||||
delete(s.discovery, mediaID)
|
||||
return nil, false
|
||||
}
|
||||
// 返回副本,避免调用方修改缓存内容。
|
||||
return append([]SubtitleTrack(nil), entry.tracks...), true
|
||||
}
|
||||
|
||||
func (s *SubtitleService) rememberDiscovery(mediaID string, tracks []SubtitleTrack) {
|
||||
now := time.Now()
|
||||
s.cacheMu.Lock()
|
||||
defer s.cacheMu.Unlock()
|
||||
if s.discovery == nil {
|
||||
s.discovery = make(map[string]subtitleDiscoveryEntry)
|
||||
}
|
||||
if len(s.discovery) >= subtitleDiscoveryCacheCap {
|
||||
s.discovery = make(map[string]subtitleDiscoveryEntry)
|
||||
}
|
||||
s.discovery[mediaID] = subtitleDiscoveryEntry{
|
||||
tracks: append([]SubtitleTrack(nil), tracks...),
|
||||
expiresAt: now.Add(subtitleDiscoveryTTL),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SubtitleService) discoverUncached(ctx context.Context, mediaID string) ([]SubtitleTrack, error) {
|
||||
m, err := s.repo.Media.FindByID(ctx, mediaID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user