download.go
⎇
Raw
1package repository
2
3import (
4 "context"
5 "database/sql"
6 "fmt"
7
8 "vidarchive/internal/models"
9)
10
11// downloadColumns is the canonical select list, kept in the order scanDownload
12// expects so the two can't drift apart.
13const downloadColumns = `id, url, status, logs, error_message, preset_id, format_override,
14 custom_flags, output_dir, subscription_id, started_at, completed_at, created_at`
15
16type DownloadRepository struct {
17 db *sql.DB
18}
19
20func NewDownloadRepository(db *sql.DB) *DownloadRepository {
21 return &DownloadRepository{db: db}
22}
23
24func (r *DownloadRepository) Create(d *models.Download) error {
25 result, err := r.db.Exec(
26 `INSERT INTO downloads (url, status, preset_id, format_override, custom_flags, output_dir, subscription_id)
27 VALUES (?, ?, ?, ?, ?, ?, ?)`,
28 d.URL, d.Status, d.PresetID, d.FormatOverride, d.CustomFlags, d.OutputDir, d.SubscriptionID,
29 )
30 if err != nil {
31 return err
32 }
33 id, err := result.LastInsertId()
34 if err != nil {
35 return fmt.Errorf("download last insert id: %w", err)
36 }
37 d.ID = id
38 return nil
39}
40
41func (r *DownloadRepository) GetByID(id int64) (*models.Download, error) {
42 row := r.db.QueryRow(`SELECT `+downloadColumns+` FROM downloads WHERE id = ?`, id)
43 return scanDownload(row)
44}
45
46func (r *DownloadRepository) GetAll(status, sortBy string) ([]*models.Download, error) {
47 query := `SELECT ` + downloadColumns + ` FROM downloads WHERE 1=1`
48 var args []interface{}
49
50 if status != "" && status != "all" {
51 query += ` AND status = ?`
52 args = append(args, status)
53 }
54
55 switch sortBy {
56 case "status":
57 query += ` ORDER BY status, created_at DESC`
58 default:
59 query += ` ORDER BY created_at DESC`
60 }
61
62 return r.queryDownloads(query, args...)
63}
64
65func (r *DownloadRepository) GetQueued(limit int) ([]*models.Download, error) {
66 return r.queryDownloads(
67 `SELECT `+downloadColumns+` FROM downloads WHERE status = 'queued' ORDER BY created_at ASC LIMIT ?`,
68 limit,
69 )
70}
71
72// CountByStatus returns the number of downloads per status, aggregated in SQL so
73// callers that only need totals don't load every row's logs. It honours ctx: the
74// pool is limited to one connection, so a caller with a deadline (the health
75// probe) must be able to give up while a long write holds it.
76func (r *DownloadRepository) CountByStatus(ctx context.Context) (map[string]int, error) {
77 rows, err := r.db.QueryContext(ctx, `SELECT status, COUNT(*) FROM downloads GROUP BY status`)
78 if err != nil {
79 return nil, err
80 }
81 defer rows.Close()
82
83 counts := make(map[string]int)
84 for rows.Next() {
85 var status string
86 var n int
87 if err := rows.Scan(&status, &n); err != nil {
88 return nil, err
89 }
90 counts[status] = n
91 }
92 return counts, rows.Err()
93}
94
95// IDsByStatus returns the ids of downloads in the given status.
96func (r *DownloadRepository) IDsByStatus(status string) ([]int64, error) {
97 rows, err := r.db.Query(`SELECT id FROM downloads WHERE status = ?`, status)
98 if err != nil {
99 return nil, err
100 }
101 defer rows.Close()
102
103 var ids []int64
104 for rows.Next() {
105 var id int64
106 if err := rows.Scan(&id); err != nil {
107 return nil, err
108 }
109 ids = append(ids, id)
110 }
111 return ids, rows.Err()
112}
113
114func (r *DownloadRepository) queryDownloads(query string, args ...interface{}) ([]*models.Download, error) {
115 rows, err := r.db.Query(query, args...)
116 if err != nil {
117 return nil, err
118 }
119 defer rows.Close()
120
121 var downloads []*models.Download
122 for rows.Next() {
123 d, err := scanDownload(rows)
124 if err != nil {
125 return nil, err
126 }
127 downloads = append(downloads, d)
128 }
129 return downloads, rows.Err()
130}
131
132func (r *DownloadRepository) AppendLogs(id int64, logs string) error {
133 _, err := r.db.Exec(
134 `UPDATE downloads SET logs = COALESCE(logs, '') || ? WHERE id = ?`,
135 logs, id,
136 )
137 return err
138}
139
140// MarkStarted atomically transitions a download from 'queued' to 'downloading'.
141// It reports whether this call actually claimed it: false means another worker
142// already started it, so the caller must not process it again.
143func (r *DownloadRepository) MarkStarted(id int64) (bool, error) {
144 res, err := r.db.Exec(
145 `UPDATE downloads SET status = 'downloading', started_at = CURRENT_TIMESTAMP WHERE id = ? AND status = 'queued'`,
146 id,
147 )
148 if err != nil {
149 return false, err
150 }
151 n, err := res.RowsAffected()
152 if err != nil {
153 return false, err
154 }
155 return n > 0, nil
156}
157
158func (r *DownloadRepository) MarkCompleted(id int64, status string) error {
159 _, err := r.db.Exec(
160 `UPDATE downloads SET status = ?, completed_at = CURRENT_TIMESTAMP WHERE id = ?`,
161 status, id,
162 )
163 return err
164}
165
166func (r *DownloadRepository) MarkError(id int64, errMsg string) error {
167 _, err := r.db.Exec(
168 `UPDATE downloads SET status = 'error', error_message = ? WHERE id = ?`,
169 errMsg, id,
170 )
171 return err
172}
173
174// HasActiveForSubscription reports whether the given subscription already has a
175// download that is queued or in progress, so the scheduler can avoid stacking a
176// second run on top of one that hasn't finished.
177func (r *DownloadRepository) HasActiveForSubscription(subID int64) (bool, error) {
178 var n int
179 err := r.db.QueryRow(
180 `SELECT COUNT(*) FROM downloads WHERE subscription_id = ? AND status IN ('queued', 'downloading')`,
181 subID,
182 ).Scan(&n)
183 if err != nil {
184 return false, err
185 }
186 return n > 0, nil
187}
188
189func (r *DownloadRepository) Delete(id int64) error {
190 _, err := r.db.Exec(`DELETE FROM downloads WHERE id = ?`, id)
191 return err
192}
193
194func (r *DownloadRepository) DeleteAll() error {
195 _, err := r.db.Exec(`DELETE FROM downloads`)
196 return err
197}
198
199func (r *DownloadRepository) UpdateStatusWhere(oldStatus, newStatus string) error {
200 _, err := r.db.Exec(
201 `UPDATE downloads SET status = ? WHERE status = ?`,
202 newStatus, oldStatus,
203 )
204 return err
205}
206
207func scanDownload(row interface{ Scan(...interface{}) error }) (*models.Download, error) {
208 var d models.Download
209 err := row.Scan(
210 &d.ID, &d.URL, &d.Status,
211 &d.Logs, &d.ErrorMessage, &d.PresetID, &d.FormatOverride, &d.CustomFlags, &d.OutputDir, &d.SubscriptionID,
212 &d.StartedAt, &d.CompletedAt, &d.CreatedAt,
213 )
214 if err != nil {
215 return nil, err
216 }
217 return &d, nil
218}
219