diff --git a/internal/database/database.go b/internal/database/database.go index 96f69e1..4180acb 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -5,6 +5,7 @@ package database import ( "fmt" "path/filepath" + "sync" "github.com/glebarez/sqlite" "go.uber.org/zap" @@ -38,6 +39,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) { if err != nil { return nil, fmt.Errorf("gorm open: %w", err) } + installSQLiteWriteGate(db) sqlDB, err := db.DB() if err != nil { return nil, fmt.Errorf("gorm sqldb: %w", err) @@ -51,6 +53,39 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) { return db, nil } +func installSQLiteWriteGate(db *gorm.DB) { + if db == nil { + return + } + gate := &sqliteWriteGate{} + lock := func(tx *gorm.DB) { + gate.Lock() + } + unlock := func(tx *gorm.DB) { + gate.Unlock() + } + _ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock) + _ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock) + _ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock) + _ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock) + _ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock) + _ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock) + _ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock) + _ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock) +} + +type sqliteWriteGate struct { + mu sync.Mutex +} + +func (g *sqliteWriteGate) Lock() { + g.mu.Lock() +} + +func (g *sqliteWriteGate) Unlock() { + g.mu.Unlock() +} + func buildDSN(cfg *config.Config) string { dbPath := cfg.Database.DBPath if !filepath.IsAbs(dbPath) { diff --git a/internal/repository/repository.go b/internal/repository/repository.go index e15a868..898be8d 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -309,6 +309,12 @@ func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB { // 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending') // 永远捞不到数据。 func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { + return withSQLiteBusyRetry(ctx, func() error { + return r.upsert(ctx, m) + }) +} + +func (r *MediaRepository) upsert(ctx context.Context, m *model.Media) error { var existing model.Media err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error if errors.Is(err, gorm.ErrRecordNotFound) { @@ -328,15 +334,16 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { } // 已存在:仅刷新文件层面的字段。 - updates := map[string]any{ - "size_bytes": m.SizeBytes, - "duration_sec": m.DurationSec, - "width": m.Width, - "height": m.Height, - "video_codec": m.VideoCodec, - "audio_codec": m.AudioCodec, - "container": m.Container, - "deleted_at": nil, + updates := map[string]any{} + setIfChanged(updates, "size_bytes", existing.SizeBytes, m.SizeBytes) + setIfChanged(updates, "duration_sec", existing.DurationSec, m.DurationSec) + setIfChanged(updates, "width", existing.Width, m.Width) + setIfChanged(updates, "height", existing.Height, m.Height) + setIfChanged(updates, "video_codec", existing.VideoCodec, m.VideoCodec) + setIfChanged(updates, "audio_codec", existing.AudioCodec, m.AudioCodec) + setIfChanged(updates, "container", existing.Container, m.Container) + if existing.DeletedAt.Valid { + updates["deleted_at"] = nil } // 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。 if m.FileID != "" && m.FileID != existing.FileID { @@ -347,56 +354,56 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { // 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题, // 避免覆盖刮削结果。 if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" { - updates["title"] = m.Title + setIfChanged(updates, "title", existing.Title, m.Title) if m.Year > 0 { - updates["year"] = m.Year + setIfChanged(updates, "year", existing.Year, m.Year) } } } if m.ScrapeStatus == "matched" { - updates["scrape_status"] = m.ScrapeStatus + setIfChanged(updates, "scrape_status", existing.ScrapeStatus, m.ScrapeStatus) if m.OriginalName != "" { - updates["original_name"] = m.OriginalName + setIfChanged(updates, "original_name", existing.OriginalName, m.OriginalName) } if m.PosterURL != "" { - updates["poster_url"] = m.PosterURL + setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL) } if m.BackdropURL != "" { - updates["backdrop_url"] = m.BackdropURL + setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL) } if m.Overview != "" { - updates["overview"] = m.Overview + setIfChanged(updates, "overview", existing.Overview, m.Overview) } if m.Rating > 0 { - updates["rating"] = m.Rating + setIfChanged(updates, "rating", existing.Rating, m.Rating) } if m.Year > 0 { - updates["year"] = m.Year + setIfChanged(updates, "year", existing.Year, m.Year) } if m.TMDbID > 0 { - updates["tm_db_id"] = m.TMDbID + setIfChanged(updates, "tm_db_id", existing.TMDbID, m.TMDbID) } if m.BangumiID > 0 { - updates["bangumi_id"] = m.BangumiID + setIfChanged(updates, "bangumi_id", existing.BangumiID, m.BangumiID) } if m.Languages != "" { - updates["languages"] = m.Languages + setIfChanged(updates, "languages", existing.Languages, m.Languages) } if m.Countries != "" { - updates["countries"] = m.Countries + setIfChanged(updates, "countries", existing.Countries, m.Countries) } if m.Genres != "" { - updates["genres"] = m.Genres + setIfChanged(updates, "genres", existing.Genres, m.Genres) } - if m.NSFW { + if m.NSFW && !existing.NSFW { updates["nsfw"] = true } } if m.PosterURL != "" { - updates["poster_url"] = m.PosterURL + setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL) } if m.BackdropURL != "" { - updates["backdrop_url"] = m.BackdropURL + setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL) } if lib := m.LibraryID; lib != "" && lib != existing.LibraryID { updates["library_id"] = m.LibraryID @@ -408,9 +415,13 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { updates["episode_num"] = m.EpisodeNum } if m.STRMURL != "" { - updates["strm_url"] = m.STRMURL + setIfChanged(updates, "strm_url", existing.STRMURL, m.STRMURL) } + if len(updates) == 0 { + *m = existing + return nil + } if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}). Where("id = ?", existing.ID).Updates(updates).Error; err != nil { return err @@ -421,6 +432,12 @@ func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error { return nil } +func setIfChanged[T comparable](updates map[string]any, key string, current, next T) { + if current != next { + updates[key] = next + } +} + // FindByID returns the media row or (nil, nil). func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) { var m model.Media diff --git a/internal/repository/repository_test.go b/internal/repository/repository_test.go index ac5b18e..cd471e9 100644 --- a/internal/repository/repository_test.go +++ b/internal/repository/repository_test.go @@ -2,6 +2,7 @@ package repository import ( "testing" + "time" "github.com/glebarez/sqlite" "gorm.io/gorm" @@ -10,6 +11,65 @@ import ( "github.com/ShukeBta/MediaStationGo/internal/model" ) +func TestMediaUpsertSkipsUnchangedExistingRow(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) + lib := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true} + if err := repos.Library.Create(t.Context(), &lib); err != nil { + t.Fatal(err) + } + media := model.Media{ + LibraryID: lib.ID, + Title: "已有影片", + Path: "/media/movie/existing.mkv", + SizeBytes: 1024, + DurationSec: 60, + Width: 1920, + Height: 1080, + VideoCodec: "h264", + AudioCodec: "aac", + Container: "matroska,webm", + ScrapeStatus: "pending", + } + if err := repos.Media.Upsert(t.Context(), &media); err != nil { + t.Fatal(err) + } + var before model.Media + if err := repos.DB.Where("path = ?", media.Path).First(&before).Error; err != nil { + t.Fatal(err) + } + time.Sleep(10 * time.Millisecond) + again := model.Media{ + LibraryID: lib.ID, + Title: before.Title, + Path: before.Path, + SizeBytes: before.SizeBytes, + DurationSec: before.DurationSec, + Width: before.Width, + Height: before.Height, + VideoCodec: before.VideoCodec, + AudioCodec: before.AudioCodec, + Container: before.Container, + ScrapeStatus: before.ScrapeStatus, + } + if err := repos.Media.Upsert(t.Context(), &again); err != nil { + t.Fatal(err) + } + var after model.Media + if err := repos.DB.Where("path = ?", media.Path).First(&after).Error; err != nil { + t.Fatal(err) + } + if !after.UpdatedAt.Equal(before.UpdatedAt) { + t.Fatalf("unchanged upsert touched updated_at: before=%s after=%s", before.UpdatedAt, after.UpdatedAt) + } +} + func TestMediaSearchFilteredSupportsChineseFuzzyTerms(t *testing.T) { db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { diff --git a/internal/service/downloads.go b/internal/service/downloads.go index 55f4dd5..67f2b90 100644 --- a/internal/service/downloads.go +++ b/internal/service/downloads.go @@ -946,10 +946,16 @@ func (d *DownloadService) syncDownloadTaskProgress(ctx context.Context, torrent if strings.TrimSpace(status) == "" { status = matched.Status } - updates := map[string]any{"progress": torrent.Progress} - if status != "" { + updates := map[string]any{} + if math.Abs(float64(matched.Progress-torrent.Progress)) > 0.0001 { + updates["progress"] = torrent.Progress + } + if status != "" && status != matched.Status { updates["status"] = status } + if len(updates) == 0 { + return + } _ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", matched.ID).Updates(updates).Error } diff --git a/internal/service/downloads_test.go b/internal/service/downloads_test.go index 2a6aa1a..dfdeec9 100644 --- a/internal/service/downloads_test.go +++ b/internal/service/downloads_test.go @@ -10,6 +10,7 @@ import ( "strings" "sync/atomic" "testing" + "time" "github.com/glebarez/sqlite" "go.uber.org/zap" @@ -47,6 +48,46 @@ func TestDownloadViewsDoNotExposePrivateURL(t *testing.T) { } } +func TestSyncDownloadTaskProgressSkipsUnchangedCompletedTask(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.DownloadTask{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + task := &model.DownloadTask{ + Source: "qbittorrent", + URL: "magnet:?xt=urn:btih:test", + Title: "Already.Done.S01E01", + SavePath: "/downloads", + Status: "completed", + Progress: 1, + } + if err := repos.Download.Create(t.Context(), task); err != nil { + t.Fatal(err) + } + var before model.DownloadTask + if err := db.First(&before, "id = ?", task.ID).Error; err != nil { + t.Fatal(err) + } + time.Sleep(10 * time.Millisecond) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.syncDownloadTaskProgress(t.Context(), QBitTorrent{ + Name: task.Title, + Progress: 1, + State: "completed", + }, tasksByIdentity([]model.DownloadTask{before})) + var after model.DownloadTask + if err := db.First(&after, "id = ?", task.ID).Error; err != nil { + t.Fatal(err) + } + if !after.UpdatedAt.Equal(before.UpdatedAt) { + t.Fatalf("unchanged completed torrent touched updated_at: before=%s after=%s", before.UpdatedAt, after.UpdatedAt) + } +} + func TestDownloadCompleteAutoOrganizesContentPath(t *testing.T) { root := t.TempDir() src := filepath.Join(root, "downloads", "国产剧", "狂飙.S01E01.2023.1080p.mkv") diff --git a/internal/service/token_svc.go b/internal/service/token_svc.go index a08a384..7791b08 100644 --- a/internal/service/token_svc.go +++ b/internal/service/token_svc.go @@ -6,6 +6,7 @@ import ( "crypto/rand" "encoding/hex" "errors" + "sync" "time" "github.com/golang-jwt/jwt/v5" @@ -37,14 +38,16 @@ type Claims struct { // TokenService 处理双令牌认证(Access Token + Refresh Token)。 type TokenService struct { - cfg *config.Config - log *zap.Logger - repo *repository.Container + cfg *config.Config + log *zap.Logger + repo *repository.Container + delayedStoreMu sync.Mutex + delayedStores map[string]struct{} } // NewTokenService 创建令牌服务实例。 func NewTokenService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *TokenService { - return &TokenService{cfg: cfg, log: log, repo: repo} + return &TokenService{cfg: cfg, log: log, repo: repo, delayedStores: make(map[string]struct{})} } // TokenPair 包含访问令牌和刷新令牌。 @@ -110,7 +113,9 @@ func (s *TokenService) issuePair(ctx context.Context, userID, role, tier string, zap.String("user_id", userID), zap.Error(err)) } - go s.storeRefreshTokenEventually(userID, tokenHash, rt.ExpiresAt) + if s.trackDelayedStore(userID, tokenHash) { + go s.storeRefreshTokenEventually(userID, tokenHash, rt.ExpiresAt) + } } return &TokenPair{ @@ -132,9 +137,12 @@ func (s *TokenService) storeRefreshToken(ctx context.Context, rt *model.RefreshT } func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, expiresAt time.Time) { - delay := 500 * time.Millisecond - for attempt := 1; attempt <= 30; attempt++ { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer s.untrackDelayedStore(userID, tokenHash) + delay := 5 * time.Second + for attempt := 1; attempt <= 8; attempt++ { + timer := time.NewTimer(delay) + <-timer.C + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) err := s.storeRefreshToken(ctx, &model.RefreshToken{ UserID: userID, TokenHash: tokenHash, @@ -150,15 +158,13 @@ func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, exp } return } - if s.log != nil && (attempt == 1 || attempt%10 == 0) { + if s.log != nil && (attempt == 1 || attempt == 4 || attempt == 8) { s.log.Warn("refresh token delayed store still waiting", zap.String("user_id", userID), zap.Int("attempt", attempt), zap.Error(err)) } - timer := time.NewTimer(delay) - <-timer.C - if delay < 10*time.Second { + if delay < 60*time.Second { delay *= 2 } } @@ -167,6 +173,33 @@ func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, exp } } +func (s *TokenService) trackDelayedStore(userID, tokenHash string) bool { + if s == nil { + return false + } + key := userID + "\x00" + tokenHash + s.delayedStoreMu.Lock() + defer s.delayedStoreMu.Unlock() + if s.delayedStores == nil { + s.delayedStores = make(map[string]struct{}) + } + if _, ok := s.delayedStores[key]; ok { + return false + } + s.delayedStores[key] = struct{}{} + return true +} + +func (s *TokenService) untrackDelayedStore(userID, tokenHash string) { + if s == nil { + return + } + key := userID + "\x00" + tokenHash + s.delayedStoreMu.Lock() + delete(s.delayedStores, key) + s.delayedStoreMu.Unlock() +} + func (s *TokenService) maxActiveRefreshTokens(ctx context.Context) int { cfg := loadBotConfig(ctx, s.repo) if cfg.MaxLoggedClients < 1 {