subscription_test.go
| 1 | package repository |
| 2 | |
| 3 | import ( |
| 4 | "database/sql" |
| 5 | "testing" |
| 6 | "time" |
| 7 | |
| 8 | "vidarchive/internal/models" |
| 9 | ) |
| 10 | |
| 11 | func TestSubscriptionRepositoryCRUD(t *testing.T) { |
| 12 | db := setupTestDB(t) |
| 13 | defer db.Close() |
| 14 | |
| 15 | repo := NewSubscriptionRepository(db) |
| 16 | |
| 17 | sub := &models.Subscription{ |
| 18 | Name: "News", |
| 19 | URL: "https://example.com/playlist", |
| 20 | Enabled: true, |
| 21 | RefreshMode: "overwrite", |
| 22 | ScheduleKind: "daily", |
| 23 | CronExpr: "0 3 * * *", |
| 24 | OutputDir: "subscriptions/news", |
| 25 | PruneRemoved: true, |
| 26 | } |
| 27 | if err := repo.Create(sub); err != nil { |
| 28 | t.Fatalf("create: %v", err) |
| 29 | } |
| 30 | if sub.ID == 0 { |
| 31 | t.Fatal("expected an assigned ID after create") |
| 32 | } |
| 33 | |
| 34 | got, err := repo.GetByID(sub.ID) |
| 35 | if err != nil { |
| 36 | t.Fatalf("get: %v", err) |
| 37 | } |
| 38 | if got.Name != "News" || got.URL != sub.URL || !got.Enabled || got.RefreshMode != "overwrite" || |
| 39 | got.OutputDir != "subscriptions/news" || !got.PruneRemoved { |
| 40 | t.Errorf("round-trip mismatch: %+v", got) |
| 41 | } |
| 42 | |
| 43 | got.Name = "Updated" |
| 44 | got.RefreshMode = "skip" |
| 45 | if err := repo.Update(got); err != nil { |
| 46 | t.Fatalf("update: %v", err) |
| 47 | } |
| 48 | again, _ := repo.GetByID(sub.ID) |
| 49 | if again.Name != "Updated" || again.RefreshMode != "skip" { |
| 50 | t.Errorf("update not persisted: %+v", again) |
| 51 | } |
| 52 | |
| 53 | if err := repo.SetEnabled(sub.ID, false); err != nil { |
| 54 | t.Fatalf("set enabled: %v", err) |
| 55 | } |
| 56 | if again, _ = repo.GetByID(sub.ID); again.Enabled { |
| 57 | t.Error("expected disabled after SetEnabled(false)") |
| 58 | } |
| 59 | |
| 60 | if err := repo.Delete(sub.ID); err != nil { |
| 61 | t.Fatalf("delete: %v", err) |
| 62 | } |
| 63 | if _, err := repo.GetByID(sub.ID); err == nil { |
| 64 | t.Error("expected error fetching deleted subscription") |
| 65 | } |
| 66 | } |
| 67 | |
| 68 | func TestSubscriptionRepositoryGetDueAndMarkRun(t *testing.T) { |
| 69 | db := setupTestDB(t) |
| 70 | defer db.Close() |
| 71 | |
| 72 | repo := NewSubscriptionRepository(db) |
| 73 | now := time.Now() |
| 74 | |
| 75 | // Due: next run in the past. |
| 76 | due := &models.Subscription{ |
| 77 | Name: "due", URL: "u1", Enabled: true, RefreshMode: "overwrite", |
| 78 | ScheduleKind: "daily", CronExpr: "0 3 * * *", OutputDir: "a", |
| 79 | NextRunAt: nullTime(now.Add(-time.Hour)), |
| 80 | } |
| 81 | // Not due: next run in the future. |
| 82 | future := &models.Subscription{ |
| 83 | Name: "future", URL: "u2", Enabled: true, RefreshMode: "overwrite", |
| 84 | ScheduleKind: "daily", CronExpr: "0 3 * * *", OutputDir: "b", |
| 85 | NextRunAt: nullTime(now.Add(time.Hour)), |
| 86 | } |
| 87 | // Disabled: never due even though its next run is in the past. |
| 88 | disabled := &models.Subscription{ |
| 89 | Name: "disabled", URL: "u3", Enabled: false, RefreshMode: "overwrite", |
| 90 | ScheduleKind: "daily", CronExpr: "0 3 * * *", OutputDir: "c", |
| 91 | NextRunAt: nullTime(now.Add(-time.Hour)), |
| 92 | } |
| 93 | for _, s := range []*models.Subscription{due, future, disabled} { |
| 94 | if err := repo.Create(s); err != nil { |
| 95 | t.Fatalf("create %s: %v", s.Name, err) |
| 96 | } |
| 97 | } |
| 98 | |
| 99 | gotDue, err := repo.GetDue(now) |
| 100 | if err != nil { |
| 101 | t.Fatalf("get due: %v", err) |
| 102 | } |
| 103 | if len(gotDue) != 1 || gotDue[0].Name != "due" { |
| 104 | names := make([]string, len(gotDue)) |
| 105 | for i, s := range gotDue { |
| 106 | names[i] = s.Name |
| 107 | } |
| 108 | t.Fatalf("expected only the due subscription, got %v", names) |
| 109 | } |
| 110 | |
| 111 | next := now.Add(24 * time.Hour) |
| 112 | if err := repo.MarkRun(due.ID, now, next, "queued"); err != nil { |
| 113 | t.Fatalf("mark run: %v", err) |
| 114 | } |
| 115 | after, _ := repo.GetByID(due.ID) |
| 116 | if !after.LastRunAt.Valid || after.LastStatus.String != "queued" { |
| 117 | t.Errorf("MarkRun not persisted: %+v", after) |
| 118 | } |
| 119 | // After marking run, it should no longer be due. |
| 120 | if gotDue, _ = repo.GetDue(now); len(gotDue) != 0 { |
| 121 | t.Errorf("expected no due subscriptions after MarkRun, got %d", len(gotDue)) |
| 122 | } |
| 123 | } |
| 124 | |
| 125 | func nullTime(t time.Time) sql.NullTime { |
| 126 | return sql.NullTime{Time: t, Valid: true} |
| 127 | } |
| 128 |