From c0d743d89d69f90cb0d097f71b57a207e6003eae Mon Sep 17 00:00:00 2001 From: ShukeBta <272197458+ShukeBta@users.noreply.github.com> Date: Mon, 10 Aug 2026 19:30:17 +0800 Subject: [PATCH] fix(subscriptions): reject concurrent duplicate rules Fixes #66 --- internal/database/database_test.go | 57 +++++++++ internal/database/schema_migration.go | 3 + .../database/schema_subscription_identity.go | 54 +++++++++ internal/handler/subscription_errors.go | 22 ++++ internal/handler/subscription_errors_test.go | 32 +++++ internal/handler/subscription_extra.go | 8 +- internal/handler/subscriptions.go | 6 + internal/model/download_subscription.go | 1 + internal/model/subscription_identity.go | 90 ++++++++++++++ internal/model/subscription_identity_test.go | 45 +++++++ .../repository/subscription_repository.go | 23 ++++ internal/service/subscription.go | 9 ++ internal/service/subscription_archive.go | 16 +++ internal/service/subscription_identity.go | 89 ++++++++++++++ .../service/subscription_identity_test.go | 114 ++++++++++++++++++ .../service/subscription_metadata_prepare.go | 7 +- web/src/pages/SubscriptionForm.tsx | 12 +- web/src/pages/SubscriptionsPage.tsx | 6 + 18 files changed, 584 insertions(+), 10 deletions(-) create mode 100644 internal/database/schema_subscription_identity.go create mode 100644 internal/handler/subscription_errors.go create mode 100644 internal/handler/subscription_errors_test.go create mode 100644 internal/model/subscription_identity.go create mode 100644 internal/model/subscription_identity_test.go create mode 100644 internal/service/subscription_identity.go create mode 100644 internal/service/subscription_identity_test.go diff --git a/internal/database/database_test.go b/internal/database/database_test.go index c874cc7..4c3d6af 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -97,6 +97,63 @@ func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T } } +func TestEnsureSubscriptionIdentityUniquenessArchivesDuplicatesAndAddsIndex(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.Subscription{}); err != nil { + t.Fatal(err) + } + createdAt := time.Now().Add(-time.Hour) + rows := []model.Subscription{ + {UserID: "user-1", Name: "Example", FeedURL: "site-search://search?keyword=Example", Filter: "Example", Resolution: "1080p", Priority: 50}, + {UserID: "user-1", Name: "Example", FeedURL: "site-search://search?keyword=Example", Filter: "Example", Resolution: "1080p", Priority: 50}, + } + for i := range rows { + rows[i].CreatedAt = createdAt.Add(time.Duration(i) * time.Minute) + if err := db.Select("*").Omit("DeletedAt").Create(&rows[i]).Error; err != nil { + t.Fatal(err) + } + } + + if err := ensureSubscriptionIdentityUniqueness(db); err != nil { + t.Fatal(err) + } + var active []model.Subscription + if err := db.Where("archived_at IS NULL").Order("created_at asc").Find(&active).Error; err != nil { + t.Fatal(err) + } + if len(active) != 1 || active[0].ID != rows[0].ID { + t.Fatalf("active subscriptions = %#v, want earliest row only", active) + } + var archived model.Subscription + if err := db.First(&archived, "id = ?", rows[1].ID).Error; err != nil { + t.Fatal(err) + } + if archived.ArchivedAt == nil || archived.ArchiveReason != duplicateSubscriptionMigrationReason { + t.Fatalf("duplicate was not archived by migration: %#v", archived) + } + + duplicate := rows[0] + duplicate.ID = "" + duplicate.CreatedAt = time.Time{} + duplicate.UpdatedAt = time.Time{} + model.RefreshSubscriptionIdentity(&duplicate) + if err := db.Select("*").Omit("DeletedAt").Create(&duplicate).Error; err == nil { + t.Fatal("expected active identity index to reject a duplicate rule") + } + different := rows[0] + different.ID = "" + different.CreatedAt = time.Time{} + different.UpdatedAt = time.Time{} + different.Resolution = "2160p" + model.RefreshSubscriptionIdentity(&different) + if err := db.Select("*").Omit("DeletedAt").Create(&different).Error; err != nil { + t.Fatalf("different rule should be allowed: %v", err) + } +} + func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) { db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { diff --git a/internal/database/schema_migration.go b/internal/database/schema_migration.go index a74d340..0ed7db7 100644 --- a/internal/database/schema_migration.go +++ b/internal/database/schema_migration.go @@ -17,6 +17,9 @@ func AutoMigrate(db *gorm.DB) error { if err := enforceTelegramBindingOneToOne(db); err != nil { return err } + if err := ensureSubscriptionIdentityUniqueness(db); err != nil { + return err + } if err := ensurePerformanceIndexes(db); err != nil { return err } diff --git a/internal/database/schema_subscription_identity.go b/internal/database/schema_subscription_identity.go new file mode 100644 index 0000000..aa9c88e --- /dev/null +++ b/internal/database/schema_subscription_identity.go @@ -0,0 +1,54 @@ +package database + +import ( + "time" + + "gorm.io/gorm" + + "github.com/ShukeBta/MediaStationGo/internal/model" +) + +const duplicateSubscriptionMigrationReason = "迁移合并重复订阅规则" + +func ensureSubscriptionIdentityUniqueness(db *gorm.DB) error { + if db == nil || !db.Migrator().HasTable(&model.Subscription{}) { + return nil + } + return db.Transaction(func(tx *gorm.DB) error { + var rows []model.Subscription + if err := tx.Unscoped().Order("created_at asc, id asc").Find(&rows).Error; err != nil { + return err + } + seen := make(map[string]string) + for i := range rows { + row := &rows[i] + key := model.RefreshSubscriptionIdentity(row) + updates := map[string]any{"identity_key": key} + if !row.DeletedAt.Valid && row.ArchivedAt == nil { + activeKey := row.UserID + "\x00" + key + if _, duplicate := seen[activeKey]; duplicate { + archivedAt := row.UpdatedAt + if archivedAt.IsZero() { + archivedAt = row.CreatedAt + } + if archivedAt.IsZero() { + archivedAt = time.Now() + } + updates["enabled"] = false + updates["archived_at"] = &archivedAt + updates["archive_reason"] = duplicateSubscriptionMigrationReason + } else { + seen[activeKey] = row.ID + } + } + if err := tx.Unscoped().Model(&model.Subscription{}).Where("id = ?", row.ID).Updates(updates).Error; err != nil { + return err + } + } + return tx.Exec(` +CREATE UNIQUE INDEX IF NOT EXISTS idx_subscriptions_user_identity_active +ON subscriptions(user_id, identity_key) +WHERE deleted_at IS NULL AND archived_at IS NULL AND identity_key <> '' +`).Error + }) +} diff --git a/internal/handler/subscription_errors.go b/internal/handler/subscription_errors.go new file mode 100644 index 0000000..8973e93 --- /dev/null +++ b/internal/handler/subscription_errors.go @@ -0,0 +1,22 @@ +package handler + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" + + "github.com/ShukeBta/MediaStationGo/internal/service" +) + +func writeSubscriptionConflict(c *gin.Context, err error) bool { + if !errors.Is(err, service.ErrSubscriptionAlreadyExists) { + return false + } + body := gin.H{"error": "subscription already exists"} + if existingID := service.SubscriptionAlreadyExistsID(err); existingID != "" { + body["existing_id"] = existingID + } + c.JSON(http.StatusConflict, body) + return true +} diff --git a/internal/handler/subscription_errors_test.go b/internal/handler/subscription_errors_test.go new file mode 100644 index 0000000..29aff2d --- /dev/null +++ b/internal/handler/subscription_errors_test.go @@ -0,0 +1,32 @@ +package handler + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + + "github.com/ShukeBta/MediaStationGo/internal/service" +) + +func TestWriteSubscriptionConflictReturns409WithExistingID(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + err := &service.SubscriptionAlreadyExistsError{ExistingID: "existing-sub"} + if !writeSubscriptionConflict(c, err) { + t.Fatal("expected conflict to be handled") + } + if w.Code != http.StatusConflict { + t.Fatalf("status = %d body=%s", w.Code, w.Body.String()) + } + var body map[string]any + if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil { + t.Fatal(err) + } + if body["existing_id"] != "existing-sub" { + t.Fatalf("body = %#v", body) + } +} diff --git a/internal/handler/subscription_extra.go b/internal/handler/subscription_extra.go index e0bcb87..b61bfa7 100644 --- a/internal/handler/subscription_extra.go +++ b/internal/handler/subscription_extra.go @@ -58,14 +58,14 @@ func updateSubscriptionHandler(svc *service.Container) gin.HandlerFunc { c.Status(http.StatusNoContent) return } - if err := svc.Repo.DB.WithContext(c.Request.Context()). - Model(&model.Subscription{}). - Where("id = ?", c.Param("id")). - Updates(updates).Error; err != nil { + if err := svc.Subscription.Update(c.Request.Context(), c.Param("id"), updates); err != nil { logSubscriptionWarn(svc, "subscription update failed", zap.String("user_id", subscriptionRequestUserID(c)), zap.String("subscription_id", c.Param("id")), zap.Error(err)) + if writeSubscriptionConflict(c, err) { + return + } c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } diff --git a/internal/handler/subscriptions.go b/internal/handler/subscriptions.go index 423dd1b..a3e74e0 100644 --- a/internal/handler/subscriptions.go +++ b/internal/handler/subscriptions.go @@ -97,6 +97,9 @@ func createSubscriptionHandler(svc *service.Container) gin.HandlerFunc { zap.String("feed_kind", subscriptionFeedKind(req.FeedURL)), zap.Bool("enabled", enabled), zap.Error(err)) + if writeSubscriptionConflict(c, err) { + return + } c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } @@ -199,6 +202,9 @@ func restoreSubscriptionHandler(svc *service.Container) gin.HandlerFunc { zap.String("user_id", subscriptionRequestUserID(c)), zap.String("subscription_id", c.Param("id")), zap.Error(err)) + if writeSubscriptionConflict(c, err) { + return + } c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } diff --git a/internal/model/download_subscription.go b/internal/model/download_subscription.go index f1d45f9..3b230bd 100644 --- a/internal/model/download_subscription.go +++ b/internal/model/download_subscription.go @@ -34,6 +34,7 @@ type DownloadTask struct { type Subscription struct { Base UserID string `gorm:"index;size:36" json:"user_id"` + IdentityKey string `gorm:"size:64" json:"-"` Name string `gorm:"size:128;not null" json:"name"` FeedURL string `gorm:"size:2048;not null" json:"feed_url"` Filter string `gorm:"size:512" json:"filter"` diff --git a/internal/model/subscription_identity.go b/internal/model/subscription_identity.go new file mode 100644 index 0000000..1cad050 --- /dev/null +++ b/internal/model/subscription_identity.go @@ -0,0 +1,90 @@ +package model + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "strings" +) + +type subscriptionIdentityPayload struct { + Name string `json:"name"` + FeedURL string `json:"feed_url"` + Filter string `json:"filter"` + MediaType string `json:"media_type"` + MediaCategory string `json:"media_category"` + SavePath string `json:"save_path"` + SearchMode string `json:"search_mode"` + IMDBID string `json:"imdb_id"` + Resolution string `json:"resolution"` + Quality string `json:"quality"` + Effects string `json:"effects"` + ReleaseGroups string `json:"release_groups"` + ExcludeWords string `json:"exclude_words"` + MinSeeders int `json:"min_seeders"` + MaxSeeders int `json:"max_seeders"` + MinSizeGB float64 `json:"min_size_gb"` + MaxSizeGB float64 `json:"max_size_gb"` + FreeOnly bool `json:"free_only"` + WashEnabled bool `json:"wash_enabled"` + WashPriority string `json:"wash_priority"` + Priority int `json:"priority"` +} + +// SubscriptionIdentityKey identifies one functional subscription rule. It +// intentionally excludes display metadata and runtime/archive state because +// those fields are backfilled or changed automatically after creation. +func SubscriptionIdentityKey(sub *Subscription) string { + if sub == nil { + return "" + } + payload := subscriptionIdentityPayload{ + Name: subscriptionIdentityFold(sub.Name), + FeedURL: strings.TrimSpace(sub.FeedURL), + Filter: strings.TrimSpace(sub.Filter), + MediaType: subscriptionIdentityFold(sub.MediaType), + MediaCategory: subscriptionIdentityFold(sub.MediaCategory), + SavePath: strings.TrimSpace(sub.SavePath), + SearchMode: subscriptionIdentityFold(sub.SearchMode), + IMDBID: subscriptionIdentityFold(sub.IMDBID), + Resolution: subscriptionIdentityFold(sub.Resolution), + Quality: subscriptionIdentityFold(sub.Quality), + Effects: subscriptionIdentityList(sub.Effects), + ReleaseGroups: subscriptionIdentityList(sub.ReleaseGroups), + ExcludeWords: subscriptionIdentityList(sub.ExcludeWords), + MinSeeders: sub.MinSeeders, + MaxSeeders: sub.MaxSeeders, + MinSizeGB: sub.MinSizeGB, + MaxSizeGB: sub.MaxSizeGB, + FreeOnly: sub.FreeOnly, + WashEnabled: sub.WashEnabled, + WashPriority: subscriptionIdentityFold(sub.WashPriority), + Priority: sub.Priority, + } + raw, _ := json.Marshal(payload) + sum := sha256.Sum256(raw) + return hex.EncodeToString(sum[:]) +} + +func RefreshSubscriptionIdentity(sub *Subscription) string { + key := SubscriptionIdentityKey(sub) + if sub != nil { + sub.IdentityKey = key + } + return key +} + +func subscriptionIdentityFold(value string) string { + return strings.ToLower(strings.TrimSpace(value)) +} + +func subscriptionIdentityList(value string) string { + parts := strings.Split(value, ",") + out := make([]string, 0, len(parts)) + for _, part := range parts { + if part = subscriptionIdentityFold(part); part != "" { + out = append(out, part) + } + } + return strings.Join(out, ",") +} diff --git a/internal/model/subscription_identity_test.go b/internal/model/subscription_identity_test.go new file mode 100644 index 0000000..361f1c0 --- /dev/null +++ b/internal/model/subscription_identity_test.go @@ -0,0 +1,45 @@ +package model + +import "testing" + +func TestSubscriptionIdentityKeyNormalizesEquivalentRules(t *testing.T) { + base := Subscription{ + Name: " Example Show ", + FeedURL: "site-search://search?keyword=Example", + Filter: " Example.*Show ", + MediaType: "TV", + MediaCategory: " 欧美剧 ", + SearchMode: "KEYWORD", + Resolution: "1080P", + Effects: " HDR, Atmos ", + ExcludeWords: " CAM, TS ", + WashPriority: "Balanced", + Priority: 50, + } + equivalent := base + equivalent.Name = "example show" + equivalent.MediaType = "tv" + equivalent.MediaCategory = "欧美剧" + equivalent.SearchMode = "keyword" + equivalent.Resolution = "1080p" + equivalent.Effects = "hdr,atmos" + equivalent.ExcludeWords = "cam,ts" + equivalent.WashPriority = "balanced" + if got, want := SubscriptionIdentityKey(&equivalent), SubscriptionIdentityKey(&base); got != want { + t.Fatalf("equivalent keys differ: %s != %s", got, want) + } + + differentRule := base + differentRule.Resolution = "2160p" + if SubscriptionIdentityKey(&differentRule) == SubscriptionIdentityKey(&base) { + t.Fatal("different resolution should produce a different identity") + } + + displayOnly := base + displayOnly.PosterURL = "https://example.test/poster.jpg" + displayOnly.Overview = "new overview" + displayOnly.TotalEpisodes = 24 + if SubscriptionIdentityKey(&displayOnly) != SubscriptionIdentityKey(&base) { + t.Fatal("display/runtime metadata should not change identity") + } +} diff --git a/internal/repository/subscription_repository.go b/internal/repository/subscription_repository.go index fe3a610..3a557b7 100644 --- a/internal/repository/subscription_repository.go +++ b/internal/repository/subscription_repository.go @@ -2,6 +2,7 @@ package repository import ( "context" + "errors" "time" "gorm.io/gorm" @@ -17,6 +18,28 @@ func (r *SubscriptionRepository) Create(ctx context.Context, s *model.Subscripti return r.db.WithContext(ctx).Select("*").Omit("DeletedAt").Create(s).Error } +// FindActiveByIdentity returns an unarchived, non-deleted rule with the same +// per-user functional identity. excludeID is used while editing/restoring. +func (r *SubscriptionRepository) FindActiveByIdentity(ctx context.Context, userID, identityKey, excludeID string) (*model.Subscription, error) { + if identityKey == "" { + return nil, nil + } + q := r.db.WithContext(ctx). + Where("user_id = ? AND identity_key = ? AND archived_at IS NULL", userID, identityKey) + if excludeID != "" { + q = q.Where("id <> ?", excludeID) + } + var sub model.Subscription + err := q.First(&sub).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &sub, nil +} + // List returns active subscription rules. Archived rows live in history and are // intentionally excluded from scheduler polling and the active management list. func (r *SubscriptionRepository) List(ctx context.Context) ([]model.Subscription, error) { diff --git a/internal/service/subscription.go b/internal/service/subscription.go index d12e2fb..f4e788e 100644 --- a/internal/service/subscription.go +++ b/internal/service/subscription.go @@ -98,8 +98,17 @@ func (s *SubscriptionService) Create(ctx context.Context, sub *model.Subscriptio return errors.New("name and feed_url required") } normalizeSubscriptionDefaults(sub) + model.RefreshSubscriptionIdentity(sub) + if duplicate, err := s.subscriptionDuplicate(ctx, sub, ""); err != nil { + return err + } else if duplicate != nil { + return newSubscriptionAlreadyExistsError(duplicate.ID) + } enabled := sub.Enabled if err := s.repo.Subscription.Create(ctx, sub); err != nil { + if duplicate, lookupErr := s.subscriptionDuplicate(ctx, sub, ""); lookupErr == nil && duplicate != nil { + return newSubscriptionAlreadyExistsError(duplicate.ID) + } return err } if !enabled { diff --git a/internal/service/subscription_archive.go b/internal/service/subscription_archive.go index 172e61b..76b8a29 100644 --- a/internal/service/subscription_archive.go +++ b/internal/service/subscription_archive.go @@ -23,6 +23,18 @@ func (s *SubscriptionService) Restore(ctx context.Context, id string) (*model.Su if err := s.repo.DB.WithContext(ctx).Unscoped().Where("id = ?", id).First(&sub).Error; err != nil { return nil, err } + sub.Enabled = true + sub.ArchivedAt = nil + sub.ArchiveReason = "" + sub.DeletedAt.Valid = false + sub.TotalEpisodes = 0 + normalizeSubscriptionDefaults(&sub) + model.RefreshSubscriptionIdentity(&sub) + if duplicate, err := s.subscriptionDuplicate(ctx, &sub, sub.ID); err != nil { + return nil, err + } else if duplicate != nil { + return nil, newSubscriptionAlreadyExistsError(duplicate.ID) + } if err := s.repo.DB.WithContext(ctx).Unscoped().Model(&model.Subscription{}). Where("id = ?", id). Updates(map[string]any{ @@ -30,12 +42,16 @@ func (s *SubscriptionService) Restore(ctx context.Context, id string) (*model.Su "archived_at": nil, "archive_reason": "", "deleted_at": nil, + "identity_key": sub.IdentityKey, // 重置为 0:此前可能被 feed 低估并锁死(updateSubscriptionTotalEpisodes // 只增不减,resolveSubscriptionTotalEpisodes 见 >0 即不再回查元数据)。 // 归零后下次 run 会从 TMDb/豆瓣等权威源重算真实总集数,避免恢复后 // 因"误判已无缺集"而不再搜索资源。 "total_episodes": 0, }).Error; err != nil { + if duplicate, lookupErr := s.subscriptionDuplicate(ctx, &sub, sub.ID); lookupErr == nil && duplicate != nil { + return nil, newSubscriptionAlreadyExistsError(duplicate.ID) + } return nil, err } if s.repo.Setting != nil { diff --git a/internal/service/subscription_identity.go b/internal/service/subscription_identity.go new file mode 100644 index 0000000..c11dc6b --- /dev/null +++ b/internal/service/subscription_identity.go @@ -0,0 +1,89 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "fmt" + + "github.com/ShukeBta/MediaStationGo/internal/model" +) + +var ErrSubscriptionAlreadyExists = errors.New("subscription already exists") + +type SubscriptionAlreadyExistsError struct { + ExistingID string +} + +func (e *SubscriptionAlreadyExistsError) Error() string { + if e == nil || e.ExistingID == "" { + return ErrSubscriptionAlreadyExists.Error() + } + return fmt.Sprintf("%s: %s", ErrSubscriptionAlreadyExists, e.ExistingID) +} + +func (e *SubscriptionAlreadyExistsError) Unwrap() error { + return ErrSubscriptionAlreadyExists +} + +func newSubscriptionAlreadyExistsError(existingID string) error { + return &SubscriptionAlreadyExistsError{ExistingID: existingID} +} + +func SubscriptionAlreadyExistsID(err error) string { + var conflict *SubscriptionAlreadyExistsError + if errors.As(err, &conflict) && conflict != nil { + return conflict.ExistingID + } + return "" +} + +func (s *SubscriptionService) subscriptionDuplicate(ctx context.Context, sub *model.Subscription, excludeID string) (*model.Subscription, error) { + if s == nil || s.repo == nil || s.repo.Subscription == nil || sub == nil { + return nil, nil + } + return s.repo.Subscription.FindActiveByIdentity(ctx, sub.UserID, sub.IdentityKey, excludeID) +} + +// Update applies API patch fields while recomputing the functional identity. +// The database partial unique index remains the final concurrency guard. +func (s *SubscriptionService) Update(ctx context.Context, id string, updates map[string]any) error { + if s == nil || s.repo == nil || s.repo.DB == nil { + return errors.New("subscription service unavailable") + } + var sub model.Subscription + if err := s.repo.DB.WithContext(ctx).Where("id = ?", id).First(&sub).Error; err != nil { + return err + } + raw, err := json.Marshal(updates) + if err != nil { + return err + } + if err := json.Unmarshal(raw, &sub); err != nil { + return err + } + if sub.Name == "" || sub.FeedURL == "" { + return errors.New("name and feed_url required") + } + normalizeSubscriptionDefaults(&sub) + model.RefreshSubscriptionIdentity(&sub) + if duplicate, err := s.subscriptionDuplicate(ctx, &sub, sub.ID); err != nil { + return err + } else if duplicate != nil { + return newSubscriptionAlreadyExistsError(duplicate.ID) + } + + updates["search_mode"] = sub.SearchMode + updates["resolution"] = sub.Resolution + updates["wash_priority"] = sub.WashPriority + updates["priority"] = sub.Priority + updates["identity_key"] = sub.IdentityKey + if err := s.repo.DB.WithContext(ctx).Model(&model.Subscription{}). + Where("id = ?", sub.ID).Updates(updates).Error; err != nil { + if duplicate, lookupErr := s.subscriptionDuplicate(ctx, &sub, sub.ID); lookupErr == nil && duplicate != nil { + return newSubscriptionAlreadyExistsError(duplicate.ID) + } + return err + } + return nil +} diff --git a/internal/service/subscription_identity_test.go b/internal/service/subscription_identity_test.go new file mode 100644 index 0000000..ca9bcb6 --- /dev/null +++ b/internal/service/subscription_identity_test.go @@ -0,0 +1,114 @@ +package service + +import ( + "errors" + "path/filepath" + "sync" + "testing" + "time" + + "go.uber.org/zap" + + "github.com/ShukeBta/MediaStationGo/internal/config" + "github.com/ShukeBta/MediaStationGo/internal/database" + "github.com/ShukeBta/MediaStationGo/internal/model" + "github.com/ShukeBta/MediaStationGo/internal/repository" +) + +func TestSubscriptionCreateRejectsConcurrentDuplicateRules(t *testing.T) { + cfg := &config.Config{} + cfg.Database.Type = "sqlite" + cfg.Database.DBPath = filepath.Join(t.TempDir(), "subscriptions.db") + cfg.Database.WALMode = true + cfg.Database.BusyTimeout = 5000 + cfg.Database.MaxOpenConns = 4 + db, err := database.Open(cfg, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + if sqlDB, dbErr := db.DB(); dbErr == nil { + t.Cleanup(func() { _ = sqlDB.Close() }) + } + if err := database.AutoMigrate(db); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + svc := NewSubscriptionService(cfg, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop())) + + start := make(chan struct{}) + results := make(chan error, 2) + var wg sync.WaitGroup + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + <-start + results <- svc.Create(t.Context(), &model.Subscription{ + UserID: "user-1", + Name: "Example Show", + FeedURL: "site-search://search?keyword=Example+Show", + Filter: "Example Show", + MediaType: "tv", + Resolution: "1080p", + Enabled: true, + }) + }() + } + close(start) + wg.Wait() + close(results) + + var created, conflicts int + for createErr := range results { + switch { + case createErr == nil: + created++ + case errors.Is(createErr, ErrSubscriptionAlreadyExists): + conflicts++ + default: + t.Fatalf("unexpected create error: %v", createErr) + } + } + if created != 1 || conflicts != 1 { + t.Fatalf("created=%d conflicts=%d, want 1/1", created, conflicts) + } + var count int64 + if err := db.Model(&model.Subscription{}).Where("archived_at IS NULL").Count(&count).Error; err != nil { + t.Fatal(err) + } + if count != 1 { + t.Fatalf("active subscription count = %d, want 1", count) + } +} + +func TestSubscriptionUpdateAndRestoreReturnExistingConflict(t *testing.T) { + db := newServiceTestDB(t, &model.Subscription{}) + if err := database.AutoMigrate(db); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + svc := NewSubscriptionService(&config.Config{}, zap.NewNop(), repos, nil, nil, NewHub(zap.NewNop())) + first := &model.Subscription{UserID: "user-1", Name: "Example", FeedURL: "rss://example", Filter: "Example", Resolution: "1080p", Enabled: true} + second := &model.Subscription{UserID: "user-1", Name: "Example", FeedURL: "rss://example", Filter: "Example", Resolution: "2160p", Enabled: true} + if err := svc.Create(t.Context(), first); err != nil { + t.Fatal(err) + } + if err := svc.Create(t.Context(), second); err != nil { + t.Fatal(err) + } + if err := svc.Update(t.Context(), second.ID, map[string]any{"resolution": "1080p"}); !errors.Is(err, ErrSubscriptionAlreadyExists) || SubscriptionAlreadyExistsID(err) != first.ID { + t.Fatalf("update error = %v, want conflict with %s", err, first.ID) + } + + archivedAt := time.Now() + archived := *first + archived.ID = "" + archived.ArchivedAt = &archivedAt + archived.Enabled = false + if err := repos.Subscription.Create(t.Context(), &archived); err != nil { + t.Fatal(err) + } + if _, err := svc.Restore(t.Context(), archived.ID); !errors.Is(err, ErrSubscriptionAlreadyExists) || SubscriptionAlreadyExistsID(err) != first.ID { + t.Fatalf("restore error = %v, want conflict with %s", err, first.ID) + } +} diff --git a/internal/service/subscription_metadata_prepare.go b/internal/service/subscription_metadata_prepare.go index 27cf13a..91cd7ef 100644 --- a/internal/service/subscription_metadata_prepare.go +++ b/internal/service/subscription_metadata_prepare.go @@ -16,7 +16,12 @@ func (s *SubscriptionService) prepareSubscriptionForRun(ctx context.Context, sub } normalizeSubscriptionDefaults(sub) updates := map[string]any{} - if s.fillSubscriptionRunMetadata(ctx, sub, updates); len(updates) > 0 && s.repo != nil && s.repo.DB != nil { + s.fillSubscriptionRunMetadata(ctx, sub, updates) + if identityKey := model.SubscriptionIdentityKey(sub); identityKey != sub.IdentityKey { + sub.IdentityKey = identityKey + updates["identity_key"] = identityKey + } + if len(updates) > 0 && s.repo != nil && s.repo.DB != nil { if err := s.repo.DB.WithContext(ctx).Model(&model.Subscription{}).Where("id = ?", sub.ID).Updates(updates).Error; err != nil && s.log != nil { s.log.Debug("subscription metadata prepare persist failed", zap.String("id", sub.ID), zap.Error(err)) } diff --git a/web/src/pages/SubscriptionForm.tsx b/web/src/pages/SubscriptionForm.tsx index d745b19..0b5cd85 100644 --- a/web/src/pages/SubscriptionForm.tsx +++ b/web/src/pages/SubscriptionForm.tsx @@ -1,17 +1,18 @@ import { FormEvent } from 'react' -import { Plus, Save } from 'lucide-react' +import { Loader2, Plus, Save } from 'lucide-react' import type { SubscriptionFormValues } from './subscriptionFormModel' interface SubscriptionFormProps { values: SubscriptionFormValues editing: boolean + busy: boolean onSubmit: (event: FormEvent) => void onCancelEdit: () => void onChange: (key: K, value: SubscriptionFormValues[K]) => void } -export function SubscriptionForm({ values, editing, onSubmit, onCancelEdit, onChange }: SubscriptionFormProps) { +export function SubscriptionForm({ values, editing, busy, onSubmit, onCancelEdit, onChange }: SubscriptionFormProps) { return (
只下载免费资源 - {editing && (