diff --git a/internal/repository/strm_repository.go b/internal/repository/strm_repository.go index 6e25ad8..55a111f 100644 --- a/internal/repository/strm_repository.go +++ b/internal/repository/strm_repository.go @@ -3,6 +3,7 @@ package repository import ( "context" "errors" + "sync" "time" "gorm.io/gorm" @@ -10,13 +11,17 @@ import ( "github.com/ShukeBta/MMTL/internal/model" ) +var strmClaimMu sync.Mutex + // ─── StrmAccount ─────────────────────────────────────────────────────────────── // StrmAccountRepository persists model.StrmAccount. type StrmAccountRepository struct{ db *gorm.DB } func (r *StrmAccountRepository) Create(ctx context.Context, a *model.StrmAccount) error { - return r.db.WithContext(ctx).Create(a).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(a).Error + }) } func (r *StrmAccountRepository) FindByID(ctx context.Context, id string) (*model.StrmAccount, error) { @@ -38,20 +43,24 @@ func (r *StrmAccountRepository) List(ctx context.Context) ([]model.StrmAccount, } func (r *StrmAccountRepository) Update(ctx context.Context, a *model.StrmAccount) error { - return r.db.WithContext(ctx).Model(&model.StrmAccount{}).Where("id = ?", a.ID).Updates(map[string]any{ - "name": a.Name, - "provider": a.Provider, - "config": a.Config, - "enabled": a.Enabled, - "last_test_at": a.LastTestAt, - "last_test_result": a.LastTestResult, - "last_test_ok": a.LastTestOK, - "updated_at": time.Now(), - }).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.StrmAccount{}).Where("id = ?", a.ID).Updates(map[string]any{ + "name": a.Name, + "provider": a.Provider, + "config": a.Config, + "enabled": a.Enabled, + "last_test_at": a.LastTestAt, + "last_test_result": a.LastTestResult, + "last_test_ok": a.LastTestOK, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmAccountRepository) Delete(ctx context.Context, id string) error { - return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmAccount{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmAccount{}).Error + }) } // ─── StrmSyncPath ────────────────────────────────────────────────────────────── @@ -60,7 +69,9 @@ func (r *StrmAccountRepository) Delete(ctx context.Context, id string) error { type StrmSyncPathRepository struct{ db *gorm.DB } func (r *StrmSyncPathRepository) Create(ctx context.Context, p *model.StrmSyncPath) error { - return r.db.WithContext(ctx).Create(p).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(p).Error + }) } func (r *StrmSyncPathRepository) FindByID(ctx context.Context, id string) (*model.StrmSyncPath, error) { @@ -82,34 +93,38 @@ func (r *StrmSyncPathRepository) List(ctx context.Context) ([]model.StrmSyncPath } func (r *StrmSyncPathRepository) Update(ctx context.Context, p *model.StrmSyncPath) error { - return r.db.WithContext(ctx).Model(&model.StrmSyncPath{}).Where("id = ?", p.ID).Updates(map[string]any{ - "name": p.Name, - "account_id": p.AccountID, - "provider": p.Provider, - "remote_path": p.RemotePath, - "local_path": p.LocalPath, - "strm_base_url": p.StrmBaseURL, - "video_ext": p.VideoExt, - "meta_ext": p.MetaExt, - "exclude_name": p.ExcludeName, - "min_video_size_mb": p.MinVideoSizeMB, - "add_path": p.AddPath, - "download_meta": p.DownloadMeta, - "upload_meta": p.UploadMeta, + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.StrmSyncPath{}).Where("id = ?", p.ID).Updates(map[string]any{ + "name": p.Name, + "account_id": p.AccountID, + "provider": p.Provider, + "remote_path": p.RemotePath, + "local_path": p.LocalPath, + "strm_base_url": p.StrmBaseURL, + "video_ext": p.VideoExt, + "meta_ext": p.MetaExt, + "exclude_name": p.ExcludeName, + "min_video_size_mb": p.MinVideoSizeMB, + "add_path": p.AddPath, + "download_meta": p.DownloadMeta, + "upload_meta": p.UploadMeta, "delete_dir": p.DeleteDir, "cron": p.Cron, "enable_cron": p.EnableCron, "sync_mode": p.SyncMode, "enabled": p.Enabled, "last_sync_at": p.LastSyncAt, - "last_sync_status": p.LastSyncStatus, - "last_sync_message": p.LastSyncMessage, - "updated_at": time.Now(), - }).Error + "last_sync_status": p.LastSyncStatus, + "last_sync_message": p.LastSyncMessage, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) error { - return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmSyncPath{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmSyncPath{}).Error + }) } // ─── StrmSyncRecord ──────────────────────────────────────────────────────────── @@ -118,24 +133,28 @@ func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) error { type StrmSyncRecordRepository struct{ db *gorm.DB } func (r *StrmSyncRecordRepository) Create(ctx context.Context, rec *model.StrmSyncRecord) error { - return r.db.WithContext(ctx).Create(rec).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(rec).Error + }) } func (r *StrmSyncRecordRepository) Update(ctx context.Context, rec *model.StrmSyncRecord) error { - return r.db.WithContext(ctx).Model(&model.StrmSyncRecord{}).Where("id = ?", rec.ID).Updates(map[string]any{ - "sync_type": rec.SyncType, - "status": rec.Status, - "total": rec.Total, - "new_strm": rec.NewStrm, - "new_meta": rec.NewMeta, - "uploaded": rec.Uploaded, - "pruned": rec.Pruned, - "skipped": rec.Skipped, - "message": rec.Message, - "started_at": rec.StartedAt, - "finished_at": rec.FinishedAt, - "updated_at": time.Now(), - }).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.StrmSyncRecord{}).Where("id = ?", rec.ID).Updates(map[string]any{ + "sync_type": rec.SyncType, + "status": rec.Status, + "total": rec.Total, + "new_strm": rec.NewStrm, + "new_meta": rec.NewMeta, + "uploaded": rec.Uploaded, + "pruned": rec.Pruned, + "skipped": rec.Skipped, + "message": rec.Message, + "started_at": rec.StartedAt, + "finished_at": rec.FinishedAt, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmSyncRecordRepository) List(ctx context.Context, syncPathID string, limit int) ([]model.StrmSyncRecord, error) { @@ -157,7 +176,21 @@ func (r *StrmSyncRecordRepository) List(ctx context.Context, syncPathID string, type StrmDownloadTaskRepository struct{ db *gorm.DB } func (r *StrmDownloadTaskRepository) Create(ctx context.Context, t *model.StrmDownloadTask) error { - return r.db.WithContext(ctx).Create(t).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(t).Error + }) +} + +func (r *StrmDownloadTaskRepository) CreateInBatches(ctx context.Context, tasks []*model.StrmDownloadTask, batchSize int) error { + if len(tasks) == 0 { + return nil + } + if batchSize <= 0 { + batchSize = 100 + } + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).CreateInBatches(tasks, batchSize).Error + }) } func (r *StrmDownloadTaskRepository) FindByID(ctx context.Context, id string) (*model.StrmDownloadTask, error) { @@ -207,24 +240,29 @@ func (r *StrmDownloadTaskRepository) CountByStatus(ctx context.Context) (map[str // ClaimPendingDownload picks the oldest pending task and marks it running. // Returns (nil, nil) when the queue is empty. func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, limit int) ([]model.StrmDownloadTask, error) { + strmClaimMu.Lock() + defer strmClaimMu.Unlock() + var rows []model.StrmDownloadTask - err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). - Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { - return err - } - if len(rows) == 0 { - return nil - } - ids := make([]string, 0, len(rows)) - now := time.Now() - for i := range rows { - ids = append(ids, rows[i].ID) - rows[i].Status = model.StrmTaskRunning - rows[i].StartedAt = &now - } - return tx.Model(&model.StrmDownloadTask{}).Where("id IN ?", ids). - Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error + err := withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). + Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { + return err + } + if len(rows) == 0 { + return nil + } + ids := make([]string, 0, len(rows)) + now := time.Now() + for i := range rows { + ids = append(ids, rows[i].ID) + rows[i].Status = model.StrmTaskRunning + rows[i].StartedAt = &now + } + return tx.Model(&model.StrmDownloadTask{}).Where("id IN ?", ids). + Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error + }) }) if err != nil { return nil, err @@ -233,62 +271,86 @@ func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, l } func (r *StrmDownloadTaskRepository) Update(ctx context.Context, t *model.StrmDownloadTask) error { - return r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).Where("id = ?", t.ID).Updates(map[string]any{ - "status": t.Status, - "error": t.Error, - "retry_count": t.RetryCount, - "next_try_at": t.NextTryAt, - "started_at": t.StartedAt, - "finished_at": t.FinishedAt, - "updated_at": time.Now(), - }).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).Where("id = ?", t.ID).Updates(map[string]any{ + "status": t.Status, + "error": t.Error, + "retry_count": t.RetryCount, + "next_try_at": t.NextTryAt, + "started_at": t.StartedAt, + "finished_at": t.FinishedAt, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) error { - return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmDownloadTask{}).Error + }) } // ClearDone 清空全部已完成下载任务。 func (r *StrmDownloadTaskRepository) ClearDone(ctx context.Context) (int64, error) { - res := r.db.WithContext(ctx).Where("status = ?", model.StrmTaskDone).Delete(&model.StrmDownloadTask{}) - return res.RowsAffected, res.Error + var count int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Where("status = ?", model.StrmTaskDone).Delete(&model.StrmDownloadTask{}) + count = res.RowsAffected + return res.Error + }) + return count, err } // ClearFinished 清空全部已完成与失败下载任务。 func (r *StrmDownloadTaskRepository) ClearFinished(ctx context.Context) (int64, error) { - res := r.db.WithContext(ctx).Where("status IN ?", []string{model.StrmTaskDone, model.StrmTaskFailed}). - Delete(&model.StrmDownloadTask{}) - return res.RowsAffected, res.Error + var count int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Where("status IN ?", []string{model.StrmTaskDone, model.StrmTaskFailed}). + Delete(&model.StrmDownloadTask{}) + count = res.RowsAffected + return res.Error + }) + return count, err } // RetryAllFailed 把所有失败任务重置回待处理,清空错误与重试计数。 func (r *StrmDownloadTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) { - res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). - Where("status = ?", model.StrmTaskFailed). - Updates(map[string]any{ - "status": model.StrmTaskPending, - "error": "", - "retry_count": 0, - "next_try_at": nil, - "started_at": nil, - "finished_at": nil, - "updated_at": time.Now(), - }) - return res.RowsAffected, res.Error + var count int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). + Where("status = ?", model.StrmTaskFailed). + Updates(map[string]any{ + "status": model.StrmTaskPending, + "error": "", + "retry_count": 0, + "next_try_at": nil, + "started_at": nil, + "finished_at": nil, + "updated_at": time.Now(), + }) + count = res.RowsAffected + return res.Error + }) + return count, err } // CancelPending 批量取消所有排队中的任务。 func (r *StrmDownloadTaskRepository) CancelPending(ctx context.Context) (int64, error) { now := time.Now() - res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). - Where("status = ?", model.StrmTaskPending). - Updates(map[string]any{ - "status": model.StrmTaskCanceled, - "error": "已批量取消", - "finished_at": now, - "updated_at": now, - }) - return res.RowsAffected, res.Error + var count int64 + err := withSQLiteBusyRetry(ctx, func() error { + res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). + Where("status = ?", model.StrmTaskPending). + Updates(map[string]any{ + "status": model.StrmTaskCanceled, + "error": "已批量取消", + "finished_at": now, + "updated_at": now, + }) + count = res.RowsAffected + return res.Error + }) + return count, err } // CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。 @@ -301,10 +363,28 @@ func (r *StrmDownloadTaskRepository) CountActive(ctx context.Context, syncPathID return count } +// GetActiveLocalPathMap 一次性获取某同步目录下正在排队或执行中的 local_path 集合,供同步时 O(1) 内存去重。 +func (r *StrmDownloadTaskRepository) GetActiveLocalPathMap(ctx context.Context, syncPathID string) (map[string]bool, error) { + var paths []string + err := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). + Where("sync_path_id = ? AND status IN ?", syncPathID, []string{model.StrmTaskPending, model.StrmTaskRunning}). + Pluck("local_path", &paths).Error + if err != nil { + return nil, err + } + out := make(map[string]bool, len(paths)) + for _, p := range paths { + out[p] = true + } + return out, nil +} + func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error { - return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", - []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). - Delete(&model.StrmDownloadTask{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", + []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). + Delete(&model.StrmDownloadTask{}).Error + }) } // ─── StrmUploadTask ──────────────────────────────────────────────────────────── @@ -313,7 +393,21 @@ func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context type StrmUploadTaskRepository struct{ db *gorm.DB } func (r *StrmUploadTaskRepository) Create(ctx context.Context, t *model.StrmUploadTask) error { - return r.db.WithContext(ctx).Create(t).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(t).Error + }) +} + +func (r *StrmUploadTaskRepository) CreateInBatches(ctx context.Context, tasks []*model.StrmUploadTask, batchSize int) error { + if len(tasks) == 0 { + return nil + } + if batchSize <= 0 { + batchSize = 100 + } + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).CreateInBatches(tasks, batchSize).Error + }) } func (r *StrmUploadTaskRepository) FindByID(ctx context.Context, id string) (*model.StrmUploadTask, error) { @@ -382,24 +476,29 @@ func (r *StrmUploadTaskRepository) CountByStatus(ctx context.Context) (map[strin // ClaimPendingUpload picks the oldest pending task and marks it running. // Returns (nil, nil) when the queue is empty. func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit int) ([]model.StrmUploadTask, error) { + strmClaimMu.Lock() + defer strmClaimMu.Unlock() + var rows []model.StrmUploadTask - err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). - Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { - return err - } - if len(rows) == 0 { - return nil - } - ids := make([]string, 0, len(rows)) - now := time.Now() - for i := range rows { - ids = append(ids, rows[i].ID) - rows[i].Status = model.StrmTaskRunning - rows[i].StartedAt = &now - } - return tx.Model(&model.StrmUploadTask{}).Where("id IN ?", ids). - Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error + err := withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). + Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { + return err + } + if len(rows) == 0 { + return nil + } + ids := make([]string, 0, len(rows)) + now := time.Now() + for i := range rows { + ids = append(ids, rows[i].ID) + rows[i].Status = model.StrmTaskRunning + rows[i].StartedAt = &now + } + return tx.Model(&model.StrmUploadTask{}).Where("id IN ?", ids). + Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error + }) }) if err != nil { return nil, err @@ -408,19 +507,23 @@ func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit } func (r *StrmUploadTaskRepository) Update(ctx context.Context, t *model.StrmUploadTask) error { - return r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).Where("id = ?", t.ID).Updates(map[string]any{ - "status": t.Status, - "error": t.Error, - "retry_count": t.RetryCount, - "next_try_at": t.NextTryAt, - "started_at": t.StartedAt, - "finished_at": t.FinishedAt, - "updated_at": time.Now(), - }).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).Where("id = ?", t.ID).Updates(map[string]any{ + "status": t.Status, + "error": t.Error, + "retry_count": t.RetryCount, + "next_try_at": t.NextTryAt, + "started_at": t.StartedAt, + "finished_at": t.FinishedAt, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) error { - return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("id = ?", id).Delete(&model.StrmUploadTask{}).Error + }) } // CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。 @@ -433,10 +536,28 @@ func (r *StrmUploadTaskRepository) CountActive(ctx context.Context, syncPathID, return count } +// GetActiveLocalPathMap 一次性获取某同步目录下正在排队或执行中的 local_path 集合,供同步时 O(1) 内存去重。 +func (r *StrmUploadTaskRepository) GetActiveLocalPathMap(ctx context.Context, syncPathID string) (map[string]bool, error) { + var paths []string + err := r.db.WithContext(ctx).Model(&model.StrmUploadTask{}). + Where("sync_path_id = ? AND status IN ?", syncPathID, []string{model.StrmTaskPending, model.StrmTaskRunning}). + Pluck("local_path", &paths).Error + if err != nil { + return nil, err + } + out := make(map[string]bool, len(paths)) + for _, p := range paths { + out[p] = true + } + return out, nil +} + func (r *StrmUploadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error { - return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", - []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). - Delete(&model.StrmUploadTask{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", + []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). + Delete(&model.StrmUploadTask{}).Error + }) } // ─── StrmDirCache ───────────────────────────────────────────────────────────── @@ -451,26 +572,31 @@ func (r *StrmDirCacheRepository) ListBySyncPathID(ctx context.Context, syncPathI } func (r *StrmDirCacheRepository) Set(ctx context.Context, syncPathID, dirID, path string) error { - var row model.StrmDirCache - err := r.db.WithContext(ctx).Where("sync_path_id = ? AND dir_id = ?", syncPathID, dirID).First(&row).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - row = model.StrmDirCache{ - SyncPathID: syncPathID, - DirID: dirID, - Path: path, + return withSQLiteBusyRetry(ctx, func() error { + var row model.StrmDirCache + err := r.db.WithContext(ctx).Where("sync_path_id = ? AND dir_id = ?", syncPathID, dirID).First(&row).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + row = model.StrmDirCache{ + SyncPathID: syncPathID, + DirID: dirID, + Path: path, + } + return r.db.WithContext(ctx).Create(&row).Error } - return r.db.WithContext(ctx).Create(&row).Error - } - if err != nil { - return err - } - return r.db.WithContext(ctx).Model(&model.StrmDirCache{}).Where("id = ?", row.ID).Updates(map[string]any{ - "path": path, - "updated_at": time.Now(), - }).Error + if err != nil { + return err + } + return r.db.WithContext(ctx).Model(&model.StrmDirCache{}).Where("id = ?", row.ID).Updates(map[string]any{ + "path": path, + "updated_at": time.Now(), + }).Error + }) } func (r *StrmDirCacheRepository) DeleteBySyncPathID(ctx context.Context, syncPathID string) error { - return r.db.WithContext(ctx).Where("sync_path_id = ?", syncPathID).Delete(&model.StrmDirCache{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("sync_path_id = ?", syncPathID).Delete(&model.StrmDirCache{}).Error + }) } + diff --git a/internal/service/strm_sync.go b/internal/service/strm_sync.go index 7ca89d3..4c34ebe 100644 --- a/internal/service/strm_sync.go +++ b/internal/service/strm_sync.go @@ -34,12 +34,17 @@ type strmSyncState struct { rec *model.StrmSyncRecord syncType string - mu sync.Mutex - processed int // 已处理文件计数(用于定期落库进度) - seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频 - seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据 - remoteMeta map[string]int64 // 远端元数据大小(上传比对用) - dirCache sync.Map // dirID (string) -> relativePath (string) + mu sync.Mutex + processed int // 已处理文件计数(用于定期落库进度) + lastProgressFlush time.Time // 上次进度落库时间 + seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频 + seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据 + remoteMeta map[string]int64 // 远端元数据大小(上传比对用) + activeDownloadPaths map[string]bool // 本地已在排队/进行的下载任务路径(内存去重) + activeUploadPaths map[string]bool // 本地已在排队/进行的上传任务路径(内存去重) + pendingDownloads []*model.StrmDownloadTask + pendingUploads []*model.StrmUploadTask + dirCache sync.Map // dirID (string) -> relativePath (string) } // StartSync 启动一次同步(异步执行,同一目录同时只允许一个任务)。 @@ -227,6 +232,21 @@ func (st *strmSyncState) run() error { if err := ensureLocalDir(st.p.LocalPath); err != nil { return fmt.Errorf("创建输出目录失败:%w", err) } + if st.cfg.DownloadMeta { + if active, err := st.s.repo.StrmDownload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil { + st.activeDownloadPaths = active + } else { + st.activeDownloadPaths = map[string]bool{} + } + } + if st.cfg.UploadMeta { + if active, err := st.s.repo.StrmUpload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil { + st.activeUploadPaths = active + } else { + st.activeUploadPaths = map[string]bool{} + } + } + if st.provider != nil { if open115, ok := st.provider.(cloud.OpenAPI115Provider); ok && st.p.Provider == model.StrmProvider115 { if err := st.walk115Flat(open115.OpenClient()); err != nil { @@ -242,11 +262,13 @@ func (st *strmSyncState) run() error { return err } } + st.flushPendingDownloads() st.flushProgress() if st.cfg.UploadMeta && st.provider != nil && st.p.Provider != model.StrmProvider115 { if err := st.scanLocalMetaForUpload(); err != nil { return err } + st.flushPendingUploads() } if err := st.pruneLocal(); err != nil { return err @@ -265,6 +287,7 @@ const strmScanWorkers = 8 // 多个 worker 并行执行 List(受全局 115 令牌桶限流约束),子目录动态 // 入队;任一目录失败则取消其余 worker 并返回错误(与旧串行版语义一致)。 func (st *strmSyncState) walkRemote() error { + defer st.flushPendingDownloads() root := strings.TrimSpace(st.p.RemotePath) if root == "" { root = "/" @@ -403,6 +426,7 @@ func (st *strmSyncState) isMetaExt(ext string) bool { // walk115Flat 使用 115 开放平台扁平化分页批量拉取机制与目录拓扑缓存(参考 QMediaSync)。 // 极大地降低 API 请求次数并支持毫秒级/秒级增量同步。 func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error { + defer st.flushPendingDownloads() ctx := st.ctx rootCID := strings.TrimSpace(st.p.RemotePath) if rootCID == "" { @@ -738,6 +762,36 @@ func (st *strmSyncState) recordRemoteMeta(entry cloud.FileEntry, rel string) { st.mu.Unlock() } +func (st *strmSyncState) flushPendingDownloads() { + st.mu.Lock() + if len(st.pendingDownloads) == 0 { + st.mu.Unlock() + return + } + batch := st.pendingDownloads + st.pendingDownloads = nil + st.mu.Unlock() + + if err := st.s.repo.StrmDownload.CreateInBatches(st.ctx, batch, 100); err != nil { + st.s.log.Warn("batch enqueue strm download tasks failed", zap.Error(err)) + } +} + +func (st *strmSyncState) flushPendingUploads() { + st.mu.Lock() + if len(st.pendingUploads) == 0 { + st.mu.Unlock() + return + } + batch := st.pendingUploads + st.pendingUploads = nil + st.mu.Unlock() + + if err := st.s.repo.StrmUpload.CreateInBatches(st.ctx, batch, 100); err != nil { + st.s.log.Warn("batch enqueue strm upload tasks failed", zap.Error(err)) + } +} + // handleMeta 元数据入下载队列(本地已存在且大小一致则跳过)。 func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) { st.recordRemoteMeta(entry, rel) @@ -750,10 +804,22 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) { st.touchProgress() return } - if st.taskExists("download", st.p.ID, target) { + st.mu.Lock() + if st.activeDownloadPaths == nil { + if active, err := st.s.repo.StrmDownload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil { + st.activeDownloadPaths = active + } else { + st.activeDownloadPaths = map[string]bool{} + } + } + if st.activeDownloadPaths[target] { + st.mu.Unlock() st.touchProgress() return } + st.activeDownloadPaths[target] = true + st.mu.Unlock() + task := &model.StrmDownloadTask{ SyncPathID: st.p.ID, AccountID: st.p.AccountID, @@ -772,13 +838,16 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) { if st.p.Provider != model.StrmProvider115 { task.RemoteRef = entry.ID } - if err := st.s.repo.StrmDownload.Create(st.ctx, task); err != nil { - st.s.log.Warn("enqueue strm download task failed", zap.Error(err)) - return - } + st.mu.Lock() + st.pendingDownloads = append(st.pendingDownloads, task) + shouldFlush := len(st.pendingDownloads) >= 100 st.rec.NewMeta++ st.mu.Unlock() + + if shouldFlush { + st.flushPendingDownloads() + } st.touchProgress() } @@ -872,67 +941,84 @@ func (st *strmSyncState) walkLocalSource() error { }) } -// scanLocalMetaForUpload 扫描本地元数据,与远端比对后入上传队列。 -func (st *strmSyncState) scanLocalMetaForUpload() error { - localRoot := filepath.Clean(st.p.LocalPath) - return filepath.WalkDir(localRoot, func(path string, d os.DirEntry, err error) error { - if err != nil { + // scanLocalMetaForUpload 扫描本地元数据,与远端比对后入上传队列。 + func (st *strmSyncState) scanLocalMetaForUpload() error { + defer st.flushPendingUploads() + if st.activeUploadPaths == nil { + if active, err := st.s.repo.StrmUpload.GetActiveLocalPathMap(st.ctx, st.p.ID); err == nil { + st.activeUploadPaths = active + } else { + st.activeUploadPaths = map[string]bool{} + } + } + localRoot := filepath.Clean(st.p.LocalPath) + return filepath.WalkDir(localRoot, func(path string, d os.DirEntry, err error) error { + if err != nil { + return nil + } + if path == localRoot { + return nil + } + select { + case <-st.ctx.Done(): + return st.ctx.Err() + default: + } + if d.IsDir() { + return nil + } + rel, err := filepath.Rel(localRoot, path) + if err != nil { + return nil + } + rel = filepath.ToSlash(rel) + ext := strings.ToLower(filepath.Ext(rel)) + if !st.isMetaExt(ext) { + return nil + } + info, err := d.Info() + if err != nil { + return nil + } + st.mu.Lock() + _, exists := st.remoteMeta["m:"+rel] + st.mu.Unlock() + if exists { + // 网盘端已存在该元数据文件,跳过上传 + return nil + } + st.mu.Lock() + if st.activeUploadPaths != nil && st.activeUploadPaths[path] { + st.mu.Unlock() + return nil + } + if st.activeUploadPaths != nil { + st.activeUploadPaths[path] = true + } + st.mu.Unlock() + + task := &model.StrmUploadTask{ + SyncPathID: st.p.ID, + AccountID: st.p.AccountID, + Provider: st.p.Provider, + FileName: filepath.Base(rel), + LocalPath: path, + RemotePath: st.remoteUploadPath(rel), + Size: info.Size(), + Status: model.StrmTaskPending, + } + st.mu.Lock() + st.pendingUploads = append(st.pendingUploads, task) + shouldFlush := len(st.pendingUploads) >= 100 + st.rec.Uploaded++ + st.mu.Unlock() + + if shouldFlush { + st.flushPendingUploads() + } return nil - } - if path == localRoot { - return nil - } - select { - case <-st.ctx.Done(): - return st.ctx.Err() - default: - } - if d.IsDir() { - return nil - } - rel, err := filepath.Rel(localRoot, path) - if err != nil { - return nil - } - rel = filepath.ToSlash(rel) - ext := strings.ToLower(filepath.Ext(rel)) - if !st.isMetaExt(ext) { - return nil - } - info, err := d.Info() - if err != nil { - return nil - } - st.mu.Lock() - _, exists := st.remoteMeta["m:"+rel] - st.mu.Unlock() - if exists { - // 网盘端已存在该元数据文件,跳过上传 - return nil - } - if st.taskExists("upload", st.p.ID, path) { - return nil - } - task := &model.StrmUploadTask{ - SyncPathID: st.p.ID, - AccountID: st.p.AccountID, - Provider: st.p.Provider, - FileName: filepath.Base(rel), - LocalPath: path, - RemotePath: st.remoteUploadPath(rel), - Size: info.Size(), - Status: model.StrmTaskPending, - } - if err := st.s.repo.StrmUpload.Create(st.ctx, task); err != nil { - st.s.log.Warn("enqueue strm upload task failed", zap.Error(err)) - return nil - } - st.mu.Lock() - st.rec.Uploaded++ - st.mu.Unlock() - return nil - }) -} + }) + } // remoteUploadPath 远端元数据目标路径 = 同步目录远端根 + 相对路径。 func (st *strmSyncState) remoteUploadPath(rel string) string { @@ -1018,12 +1104,16 @@ func (st *strmSyncState) pruneLocal() error { return nil } -// touchProgress 每处理若干个文件落库一次进度。 +// touchProgress 进度计数并限流防抖落库(避免高频写 SQLite 导致锁竞争)。 func (st *strmSyncState) touchProgress() { st.mu.Lock() st.rec.Total++ st.processed++ - flush := st.processed%100 == 0 + now := time.Now() + flush := st.processed%100 == 0 || (st.processed%20 == 0 && now.Sub(st.lastProgressFlush) >= 2*time.Second) + if flush { + st.lastProgressFlush = now + } st.mu.Unlock() if flush { st.flushProgress() diff --git a/internal/service/strm_sync_test.go b/internal/service/strm_sync_test.go index 8abe99f..c5254a7 100644 --- a/internal/service/strm_sync_test.go +++ b/internal/service/strm_sync_test.go @@ -6,6 +6,7 @@ import ( "os" "path/filepath" "strings" + "sync" "testing" "time" @@ -493,7 +494,68 @@ func TestWalkRemoteConcurrent(t *testing.T) { if walkErr != nil { t.Fatal(walkErr) } - if strmCount != 5 { - t.Errorf("生成的 .strm 数量 = %d,期望 5", strmCount) + if strmCount != 5 { + t.Errorf("生成的 .strm 数量 = %d,期望 5", strmCount) + } } -} + + // TestStrmBatchEnqueueAndConcurrentClaim 测试大规模批量入库及多协程并发认领无死锁 + func TestStrmBatchEnqueueAndConcurrentClaim(t *testing.T) { + svc := testStrmService(t) + ctx := context.Background() + + // 1. 批量插入 200 个下载任务 + tasks := make([]*model.StrmDownloadTask, 0, 200) + for i := 0; i < 200; i++ { + tasks = append(tasks, &model.StrmDownloadTask{ + SyncPathID: "test-sync-path", + AccountID: "test-acct", + Provider: model.StrmProvider115, + FileName: filepath.Base(string(rune('a'+i%26))) + ".nfo", + LocalPath: filepath.Join(t.TempDir(), string(rune('a'+i%26)), "test.nfo"), + Status: model.StrmTaskPending, + }) + } + if err := svc.repo.StrmDownload.CreateInBatches(ctx, tasks, 50); err != nil { + t.Fatalf("CreateInBatches failed: %v", err) + } + + // 2. 验证 ActiveLocalPathMap + activeMap, err := svc.repo.StrmDownload.GetActiveLocalPathMap(ctx, "test-sync-path") + if err != nil { + t.Fatalf("GetActiveLocalPathMap failed: %v", err) + } + if len(activeMap) == 0 { + t.Fatal("expected active local path map to have entries") + } + + // 3. 模拟 6 个 worker 并发 ClaimPendingDownload + claimedCount := 0 + var claimMu sync.Mutex + var wg sync.WaitGroup + for w := 0; w < 6; w++ { + wg.Add(1) + go func() { + defer wg.Done() + for { + batch, err := svc.repo.StrmDownload.ClaimPendingDownload(ctx, 10) + if err != nil { + t.Errorf("concurrent ClaimPendingDownload failed: %v", err) + return + } + if len(batch) == 0 { + return + } + claimMu.Lock() + claimedCount += len(batch) + claimMu.Unlock() + } + }() + } + wg.Wait() + + if claimedCount != 200 { + t.Fatalf("expected all 200 tasks claimed, got %d", claimedCount) + } + } +