yo 优化
This commit is contained in:
truewhile
2026-08-26 08:43:36 +08:00
parent b676733af7
commit 5f323eb2ce
3 changed files with 506 additions and 228 deletions
+278 -152
View File
@@ -3,6 +3,7 @@ package repository
import ( import (
"context" "context"
"errors" "errors"
"sync"
"time" "time"
"gorm.io/gorm" "gorm.io/gorm"
@@ -10,13 +11,17 @@ import (
"github.com/ShukeBta/MMTL/internal/model" "github.com/ShukeBta/MMTL/internal/model"
) )
var strmClaimMu sync.Mutex
// ─── StrmAccount ─────────────────────────────────────────────────────────────── // ─── StrmAccount ───────────────────────────────────────────────────────────────
// StrmAccountRepository persists model.StrmAccount. // StrmAccountRepository persists model.StrmAccount.
type StrmAccountRepository struct{ db *gorm.DB } type StrmAccountRepository struct{ db *gorm.DB }
func (r *StrmAccountRepository) Create(ctx context.Context, a *model.StrmAccount) error { 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) { 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 { 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{ return withSQLiteBusyRetry(ctx, func() error {
"name": a.Name, return r.db.WithContext(ctx).Model(&model.StrmAccount{}).Where("id = ?", a.ID).Updates(map[string]any{
"provider": a.Provider, "name": a.Name,
"config": a.Config, "provider": a.Provider,
"enabled": a.Enabled, "config": a.Config,
"last_test_at": a.LastTestAt, "enabled": a.Enabled,
"last_test_result": a.LastTestResult, "last_test_at": a.LastTestAt,
"last_test_ok": a.LastTestOK, "last_test_result": a.LastTestResult,
"updated_at": time.Now(), "last_test_ok": a.LastTestOK,
}).Error "updated_at": time.Now(),
}).Error
})
} }
func (r *StrmAccountRepository) Delete(ctx context.Context, id string) 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 ────────────────────────────────────────────────────────────── // ─── StrmSyncPath ──────────────────────────────────────────────────────────────
@@ -60,7 +69,9 @@ func (r *StrmAccountRepository) Delete(ctx context.Context, id string) error {
type StrmSyncPathRepository struct{ db *gorm.DB } type StrmSyncPathRepository struct{ db *gorm.DB }
func (r *StrmSyncPathRepository) Create(ctx context.Context, p *model.StrmSyncPath) error { 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) { 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 { 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{ return withSQLiteBusyRetry(ctx, func() error {
"name": p.Name, return r.db.WithContext(ctx).Model(&model.StrmSyncPath{}).Where("id = ?", p.ID).Updates(map[string]any{
"account_id": p.AccountID, "name": p.Name,
"provider": p.Provider, "account_id": p.AccountID,
"remote_path": p.RemotePath, "provider": p.Provider,
"local_path": p.LocalPath, "remote_path": p.RemotePath,
"strm_base_url": p.StrmBaseURL, "local_path": p.LocalPath,
"video_ext": p.VideoExt, "strm_base_url": p.StrmBaseURL,
"meta_ext": p.MetaExt, "video_ext": p.VideoExt,
"exclude_name": p.ExcludeName, "meta_ext": p.MetaExt,
"min_video_size_mb": p.MinVideoSizeMB, "exclude_name": p.ExcludeName,
"add_path": p.AddPath, "min_video_size_mb": p.MinVideoSizeMB,
"download_meta": p.DownloadMeta, "add_path": p.AddPath,
"upload_meta": p.UploadMeta, "download_meta": p.DownloadMeta,
"upload_meta": p.UploadMeta,
"delete_dir": p.DeleteDir, "delete_dir": p.DeleteDir,
"cron": p.Cron, "cron": p.Cron,
"enable_cron": p.EnableCron, "enable_cron": p.EnableCron,
"sync_mode": p.SyncMode, "sync_mode": p.SyncMode,
"enabled": p.Enabled, "enabled": p.Enabled,
"last_sync_at": p.LastSyncAt, "last_sync_at": p.LastSyncAt,
"last_sync_status": p.LastSyncStatus, "last_sync_status": p.LastSyncStatus,
"last_sync_message": p.LastSyncMessage, "last_sync_message": p.LastSyncMessage,
"updated_at": time.Now(), "updated_at": time.Now(),
}).Error }).Error
})
} }
func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) 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 ──────────────────────────────────────────────────────────── // ─── StrmSyncRecord ────────────────────────────────────────────────────────────
@@ -118,24 +133,28 @@ func (r *StrmSyncPathRepository) Delete(ctx context.Context, id string) error {
type StrmSyncRecordRepository struct{ db *gorm.DB } type StrmSyncRecordRepository struct{ db *gorm.DB }
func (r *StrmSyncRecordRepository) Create(ctx context.Context, rec *model.StrmSyncRecord) error { 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 { 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{ return withSQLiteBusyRetry(ctx, func() error {
"sync_type": rec.SyncType, return r.db.WithContext(ctx).Model(&model.StrmSyncRecord{}).Where("id = ?", rec.ID).Updates(map[string]any{
"status": rec.Status, "sync_type": rec.SyncType,
"total": rec.Total, "status": rec.Status,
"new_strm": rec.NewStrm, "total": rec.Total,
"new_meta": rec.NewMeta, "new_strm": rec.NewStrm,
"uploaded": rec.Uploaded, "new_meta": rec.NewMeta,
"pruned": rec.Pruned, "uploaded": rec.Uploaded,
"skipped": rec.Skipped, "pruned": rec.Pruned,
"message": rec.Message, "skipped": rec.Skipped,
"started_at": rec.StartedAt, "message": rec.Message,
"finished_at": rec.FinishedAt, "started_at": rec.StartedAt,
"updated_at": time.Now(), "finished_at": rec.FinishedAt,
}).Error "updated_at": time.Now(),
}).Error
})
} }
func (r *StrmSyncRecordRepository) List(ctx context.Context, syncPathID string, limit int) ([]model.StrmSyncRecord, 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 } type StrmDownloadTaskRepository struct{ db *gorm.DB }
func (r *StrmDownloadTaskRepository) Create(ctx context.Context, t *model.StrmDownloadTask) error { 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) { 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. // ClaimPendingDownload picks the oldest pending task and marks it running.
// Returns (nil, nil) when the queue is empty. // Returns (nil, nil) when the queue is empty.
func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, limit int) ([]model.StrmDownloadTask, error) { func (r *StrmDownloadTaskRepository) ClaimPendingDownload(ctx context.Context, limit int) ([]model.StrmDownloadTask, error) {
strmClaimMu.Lock()
defer strmClaimMu.Unlock()
var rows []model.StrmDownloadTask var rows []model.StrmDownloadTask
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { err := withSQLiteBusyRetry(ctx, func() error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
return err Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
} return err
if len(rows) == 0 { }
return nil if len(rows) == 0 {
} return nil
ids := make([]string, 0, len(rows)) }
now := time.Now() ids := make([]string, 0, len(rows))
for i := range rows { now := time.Now()
ids = append(ids, rows[i].ID) for i := range rows {
rows[i].Status = model.StrmTaskRunning ids = append(ids, rows[i].ID)
rows[i].StartedAt = &now 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 return tx.Model(&model.StrmDownloadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
})
}) })
if err != nil { if err != nil {
return nil, err 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 { 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{ return withSQLiteBusyRetry(ctx, func() error {
"status": t.Status, return r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"error": t.Error, "status": t.Status,
"retry_count": t.RetryCount, "error": t.Error,
"next_try_at": t.NextTryAt, "retry_count": t.RetryCount,
"started_at": t.StartedAt, "next_try_at": t.NextTryAt,
"finished_at": t.FinishedAt, "started_at": t.StartedAt,
"updated_at": time.Now(), "finished_at": t.FinishedAt,
}).Error "updated_at": time.Now(),
}).Error
})
} }
func (r *StrmDownloadTaskRepository) Delete(ctx context.Context, id string) 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 清空全部已完成下载任务。 // ClearDone 清空全部已完成下载任务。
func (r *StrmDownloadTaskRepository) ClearDone(ctx context.Context) (int64, error) { func (r *StrmDownloadTaskRepository) ClearDone(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Where("status = ?", model.StrmTaskDone).Delete(&model.StrmDownloadTask{}) var count int64
return res.RowsAffected, res.Error 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 清空全部已完成与失败下载任务。 // ClearFinished 清空全部已完成与失败下载任务。
func (r *StrmDownloadTaskRepository) ClearFinished(ctx context.Context) (int64, error) { func (r *StrmDownloadTaskRepository) ClearFinished(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Where("status IN ?", []string{model.StrmTaskDone, model.StrmTaskFailed}). var count int64
Delete(&model.StrmDownloadTask{}) err := withSQLiteBusyRetry(ctx, func() error {
return res.RowsAffected, res.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 把所有失败任务重置回待处理,清空错误与重试计数。 // RetryAllFailed 把所有失败任务重置回待处理,清空错误与重试计数。
func (r *StrmDownloadTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) { func (r *StrmDownloadTaskRepository) RetryAllFailed(ctx context.Context) (int64, error) {
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). var count int64
Where("status = ?", model.StrmTaskFailed). err := withSQLiteBusyRetry(ctx, func() error {
Updates(map[string]any{ res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
"status": model.StrmTaskPending, Where("status = ?", model.StrmTaskFailed).
"error": "", Updates(map[string]any{
"retry_count": 0, "status": model.StrmTaskPending,
"next_try_at": nil, "error": "",
"started_at": nil, "retry_count": 0,
"finished_at": nil, "next_try_at": nil,
"updated_at": time.Now(), "started_at": nil,
}) "finished_at": nil,
return res.RowsAffected, res.Error "updated_at": time.Now(),
})
count = res.RowsAffected
return res.Error
})
return count, err
} }
// CancelPending 批量取消所有排队中的任务。 // CancelPending 批量取消所有排队中的任务。
func (r *StrmDownloadTaskRepository) CancelPending(ctx context.Context) (int64, error) { func (r *StrmDownloadTaskRepository) CancelPending(ctx context.Context) (int64, error) {
now := time.Now() now := time.Now()
res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}). var count int64
Where("status = ?", model.StrmTaskPending). err := withSQLiteBusyRetry(ctx, func() error {
Updates(map[string]any{ res := r.db.WithContext(ctx).Model(&model.StrmDownloadTask{}).
"status": model.StrmTaskCanceled, Where("status = ?", model.StrmTaskPending).
"error": "已批量取消", Updates(map[string]any{
"finished_at": now, "status": model.StrmTaskCanceled,
"updated_at": now, "error": "已批量取消",
}) "finished_at": now,
return res.RowsAffected, res.Error "updated_at": now,
})
count = res.RowsAffected
return res.Error
})
return count, err
} }
// CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。 // CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。
@@ -301,10 +363,28 @@ func (r *StrmDownloadTaskRepository) CountActive(ctx context.Context, syncPathID
return count 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 { func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error {
return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", return withSQLiteBusyRetry(ctx, func() error {
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?",
Delete(&model.StrmDownloadTask{}).Error []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmDownloadTask{}).Error
})
} }
// ─── StrmUploadTask ──────────────────────────────────────────────────────────── // ─── StrmUploadTask ────────────────────────────────────────────────────────────
@@ -313,7 +393,21 @@ func (r *StrmDownloadTaskRepository) DeleteFinishedOlderThan(ctx context.Context
type StrmUploadTaskRepository struct{ db *gorm.DB } type StrmUploadTaskRepository struct{ db *gorm.DB }
func (r *StrmUploadTaskRepository) Create(ctx context.Context, t *model.StrmUploadTask) error { 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) { 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. // ClaimPendingUpload picks the oldest pending task and marks it running.
// Returns (nil, nil) when the queue is empty. // Returns (nil, nil) when the queue is empty.
func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit int) ([]model.StrmUploadTask, error) { func (r *StrmUploadTaskRepository) ClaimPendingUpload(ctx context.Context, limit int) ([]model.StrmUploadTask, error) {
strmClaimMu.Lock()
defer strmClaimMu.Unlock()
var rows []model.StrmUploadTask var rows []model.StrmUploadTask
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { err := withSQLiteBusyRetry(ctx, func() error {
if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()). return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil { if err := tx.Where("status = ? AND (next_try_at IS NULL OR next_try_at <= ?)", model.StrmTaskPending, time.Now()).
return err Order("created_at asc").Limit(limit).Find(&rows).Error; err != nil {
} return err
if len(rows) == 0 { }
return nil if len(rows) == 0 {
} return nil
ids := make([]string, 0, len(rows)) }
now := time.Now() ids := make([]string, 0, len(rows))
for i := range rows { now := time.Now()
ids = append(ids, rows[i].ID) for i := range rows {
rows[i].Status = model.StrmTaskRunning ids = append(ids, rows[i].ID)
rows[i].StartedAt = &now 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 return tx.Model(&model.StrmUploadTask{}).Where("id IN ?", ids).
Updates(map[string]any{"status": model.StrmTaskRunning, "started_at": now}).Error
})
}) })
if err != nil { if err != nil {
return nil, err 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 { 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{ return withSQLiteBusyRetry(ctx, func() error {
"status": t.Status, return r.db.WithContext(ctx).Model(&model.StrmUploadTask{}).Where("id = ?", t.ID).Updates(map[string]any{
"error": t.Error, "status": t.Status,
"retry_count": t.RetryCount, "error": t.Error,
"next_try_at": t.NextTryAt, "retry_count": t.RetryCount,
"started_at": t.StartedAt, "next_try_at": t.NextTryAt,
"finished_at": t.FinishedAt, "started_at": t.StartedAt,
"updated_at": time.Now(), "finished_at": t.FinishedAt,
}).Error "updated_at": time.Now(),
}).Error
})
} }
func (r *StrmUploadTaskRepository) Delete(ctx context.Context, id string) 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 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。 // CountActive 统计某同步目录下目标仍在排队/进行的任务数(用于去重)。
@@ -433,10 +536,28 @@ func (r *StrmUploadTaskRepository) CountActive(ctx context.Context, syncPathID,
return count 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 { func (r *StrmUploadTaskRepository) DeleteFinishedOlderThan(ctx context.Context, before time.Time) error {
return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?", return withSQLiteBusyRetry(ctx, func() error {
[]string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before). return r.db.WithContext(ctx).Where("status IN ? AND finished_at < ?",
Delete(&model.StrmUploadTask{}).Error []string{model.StrmTaskDone, model.StrmTaskFailed, model.StrmTaskCanceled}, before).
Delete(&model.StrmUploadTask{}).Error
})
} }
// ─── StrmDirCache ───────────────────────────────────────────────────────────── // ─── 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 { func (r *StrmDirCacheRepository) Set(ctx context.Context, syncPathID, dirID, path string) error {
var row model.StrmDirCache return withSQLiteBusyRetry(ctx, func() error {
err := r.db.WithContext(ctx).Where("sync_path_id = ? AND dir_id = ?", syncPathID, dirID).First(&row).Error var row model.StrmDirCache
if errors.Is(err, gorm.ErrRecordNotFound) { err := r.db.WithContext(ctx).Where("sync_path_id = ? AND dir_id = ?", syncPathID, dirID).First(&row).Error
row = model.StrmDirCache{ if errors.Is(err, gorm.ErrRecordNotFound) {
SyncPathID: syncPathID, row = model.StrmDirCache{
DirID: dirID, SyncPathID: syncPathID,
Path: path, 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
if err != nil { }
return err return r.db.WithContext(ctx).Model(&model.StrmDirCache{}).Where("id = ?", row.ID).Updates(map[string]any{
} "path": path,
return r.db.WithContext(ctx).Model(&model.StrmDirCache{}).Where("id = ?", row.ID).Updates(map[string]any{ "updated_at": time.Now(),
"path": path, }).Error
"updated_at": time.Now(), })
}).Error
} }
func (r *StrmDirCacheRepository) DeleteBySyncPathID(ctx context.Context, syncPathID string) 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
})
} }
+163 -73
View File
@@ -34,12 +34,17 @@ type strmSyncState struct {
rec *model.StrmSyncRecord rec *model.StrmSyncRecord
syncType string syncType string
mu sync.Mutex mu sync.Mutex
processed int // 已处理文件计数(用于定期落库进度) processed int // 已处理文件计数(用于定期落库进度)
seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频 lastProgressFlush time.Time // 上次进度落库时间
seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据 seenVideo map[string]bool // "v:"+去掉扩展名的相对路径 → 远端存在该视频
remoteMeta map[string]int64 // 远端元数据大小(上传比对用) seenMeta map[string]bool // "m:"+相对路径 → 远端存在该元数据
dirCache sync.Map // dirID (string) -> relativePath (string) 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 启动一次同步(异步执行,同一目录同时只允许一个任务)。 // StartSync 启动一次同步(异步执行,同一目录同时只允许一个任务)。
@@ -227,6 +232,21 @@ func (st *strmSyncState) run() error {
if err := ensureLocalDir(st.p.LocalPath); err != nil { if err := ensureLocalDir(st.p.LocalPath); err != nil {
return fmt.Errorf("创建输出目录失败:%w", err) 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 st.provider != nil {
if open115, ok := st.provider.(cloud.OpenAPI115Provider); ok && st.p.Provider == model.StrmProvider115 { if open115, ok := st.provider.(cloud.OpenAPI115Provider); ok && st.p.Provider == model.StrmProvider115 {
if err := st.walk115Flat(open115.OpenClient()); err != nil { if err := st.walk115Flat(open115.OpenClient()); err != nil {
@@ -242,11 +262,13 @@ func (st *strmSyncState) run() error {
return err return err
} }
} }
st.flushPendingDownloads()
st.flushProgress() st.flushProgress()
if st.cfg.UploadMeta && st.provider != nil && st.p.Provider != model.StrmProvider115 { if st.cfg.UploadMeta && st.provider != nil && st.p.Provider != model.StrmProvider115 {
if err := st.scanLocalMetaForUpload(); err != nil { if err := st.scanLocalMetaForUpload(); err != nil {
return err return err
} }
st.flushPendingUploads()
} }
if err := st.pruneLocal(); err != nil { if err := st.pruneLocal(); err != nil {
return err return err
@@ -265,6 +287,7 @@ const strmScanWorkers = 8
// 多个 worker 并行执行 List(受全局 115 令牌桶限流约束),子目录动态 // 多个 worker 并行执行 List(受全局 115 令牌桶限流约束),子目录动态
// 入队;任一目录失败则取消其余 worker 并返回错误(与旧串行版语义一致)。 // 入队;任一目录失败则取消其余 worker 并返回错误(与旧串行版语义一致)。
func (st *strmSyncState) walkRemote() error { func (st *strmSyncState) walkRemote() error {
defer st.flushPendingDownloads()
root := strings.TrimSpace(st.p.RemotePath) root := strings.TrimSpace(st.p.RemotePath)
if root == "" { if root == "" {
root = "/" root = "/"
@@ -403,6 +426,7 @@ func (st *strmSyncState) isMetaExt(ext string) bool {
// walk115Flat 使用 115 开放平台扁平化分页批量拉取机制与目录拓扑缓存(参考 QMediaSync)。 // walk115Flat 使用 115 开放平台扁平化分页批量拉取机制与目录拓扑缓存(参考 QMediaSync)。
// 极大地降低 API 请求次数并支持毫秒级/秒级增量同步。 // 极大地降低 API 请求次数并支持毫秒级/秒级增量同步。
func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error { func (st *strmSyncState) walk115Flat(open115 *cloud115.OpenClient) error {
defer st.flushPendingDownloads()
ctx := st.ctx ctx := st.ctx
rootCID := strings.TrimSpace(st.p.RemotePath) rootCID := strings.TrimSpace(st.p.RemotePath)
if rootCID == "" { if rootCID == "" {
@@ -738,6 +762,36 @@ func (st *strmSyncState) recordRemoteMeta(entry cloud.FileEntry, rel string) {
st.mu.Unlock() 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 元数据入下载队列(本地已存在且大小一致则跳过)。 // handleMeta 元数据入下载队列(本地已存在且大小一致则跳过)。
func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) { func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
st.recordRemoteMeta(entry, rel) st.recordRemoteMeta(entry, rel)
@@ -750,10 +804,22 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
st.touchProgress() st.touchProgress()
return 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() st.touchProgress()
return return
} }
st.activeDownloadPaths[target] = true
st.mu.Unlock()
task := &model.StrmDownloadTask{ task := &model.StrmDownloadTask{
SyncPathID: st.p.ID, SyncPathID: st.p.ID,
AccountID: st.p.AccountID, AccountID: st.p.AccountID,
@@ -772,13 +838,16 @@ func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
if st.p.Provider != model.StrmProvider115 { if st.p.Provider != model.StrmProvider115 {
task.RemoteRef = entry.ID 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.mu.Lock()
st.pendingDownloads = append(st.pendingDownloads, task)
shouldFlush := len(st.pendingDownloads) >= 100
st.rec.NewMeta++ st.rec.NewMeta++
st.mu.Unlock() st.mu.Unlock()
if shouldFlush {
st.flushPendingDownloads()
}
st.touchProgress() st.touchProgress()
} }
@@ -872,67 +941,84 @@ func (st *strmSyncState) walkLocalSource() error {
}) })
} }
// scanLocalMetaForUpload 扫描本地元数据,与远端比对后入上传队列。 // scanLocalMetaForUpload 扫描本地元数据,与远端比对后入上传队列。
func (st *strmSyncState) scanLocalMetaForUpload() error { func (st *strmSyncState) scanLocalMetaForUpload() error {
localRoot := filepath.Clean(st.p.LocalPath) defer st.flushPendingUploads()
return filepath.WalkDir(localRoot, func(path string, d os.DirEntry, err error) error { if st.activeUploadPaths == nil {
if err != 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 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 远端元数据目标路径 = 同步目录远端根 + 相对路径。 // remoteUploadPath 远端元数据目标路径 = 同步目录远端根 + 相对路径。
func (st *strmSyncState) remoteUploadPath(rel string) string { func (st *strmSyncState) remoteUploadPath(rel string) string {
@@ -1018,12 +1104,16 @@ func (st *strmSyncState) pruneLocal() error {
return nil return nil
} }
// touchProgress 每处理若干个文件落库一次进度。 // touchProgress 进度计数并限流防抖落库(避免高频写 SQLite 导致锁竞争)。
func (st *strmSyncState) touchProgress() { func (st *strmSyncState) touchProgress() {
st.mu.Lock() st.mu.Lock()
st.rec.Total++ st.rec.Total++
st.processed++ 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() st.mu.Unlock()
if flush { if flush {
st.flushProgress() st.flushProgress()
+65 -3
View File
@@ -6,6 +6,7 @@ import (
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"testing" "testing"
"time" "time"
@@ -493,7 +494,68 @@ func TestWalkRemoteConcurrent(t *testing.T) {
if walkErr != nil { if walkErr != nil {
t.Fatal(walkErr) t.Fatal(walkErr)
} }
if strmCount != 5 { if strmCount != 5 {
t.Errorf("生成的 .strm 数量 = %d,期望 5", strmCount) 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)
}
}