download.go
⎇
Raw
1package service
2
3import (
4 "bufio"
5 "database/sql"
6 "encoding/json"
7 "fmt"
8 "log"
9 "os"
10 "os/exec"
11 "path/filepath"
12 "regexp"
13 "sort"
14 "strconv"
15 "strings"
16 "sync"
17 "syscall"
18 "time"
19
20 "github.com/BurntSushi/toml"
21 "github.com/gabriel-vasile/mimetype"
22
23 "vidarchive/internal/config"
24 "vidarchive/internal/models"
25 "vidarchive/internal/repository"
26)
27
28type DownloadService struct {
29 repo *repository.DownloadRepository
30 librarySvc *LibraryService
31 presetSvc *PresetService
32 settingsSvc *SettingsService
33 cfg *config.Config
34 cache *ProgressCache
35 processMu sync.Mutex
36 processes map[int64]*os.Process
37}
38
39func NewDownloadService(repo *repository.DownloadRepository, librarySvc *LibraryService, presetSvc *PresetService, settingsSvc *SettingsService, cfg *config.Config) *DownloadService {
40 return &DownloadService{
41 repo: repo,
42 librarySvc: librarySvc,
43 presetSvc: presetSvc,
44 settingsSvc: settingsSvc,
45 cfg: cfg,
46 cache: NewProgressCache(),
47 processes: make(map[int64]*os.Process),
48 }
49}
50
51func (s *DownloadService) Create(url string, presetID *int64, formatOverride, customFlags, outputDir string) (*models.Download, error) {
52 d := &models.Download{
53 URL: url,
54 Status: "queued",
55 FormatOverride: formatOverride,
56 CustomFlags: customFlags,
57 OutputDir: sql.NullString{String: outputDir, Valid: outputDir != ""},
58 }
59
60 if presetID != nil {
61 d.PresetID = sqlNullInt64(*presetID)
62 }
63
64 if err := s.repo.Create(d); err != nil {
65 return nil, err
66 }
67 return d, nil
68}
69
70func (s *DownloadService) GetByID(id int64) (*models.Download, error) {
71 d, err := s.repo.GetByID(id)
72 if err != nil {
73 return nil, err
74 }
75 if logs := s.cache.Snapshot(id); logs != "" {
76 d.Logs = sql.NullString{String: logs, Valid: true}
77 }
78 return d, nil
79}
80
81func (s *DownloadService) GetAll(status, sortBy string) ([]*models.Download, error) {
82 downloads, err := s.repo.GetAll(status, sortBy)
83 if err != nil {
84 return nil, err
85 }
86
87 for _, d := range downloads {
88 if logs := s.cache.Snapshot(d.ID); logs != "" {
89 d.Logs = sql.NullString{String: logs, Valid: true}
90 }
91 }
92
93 return downloads, nil
94}
95
96func (s *DownloadService) GetQueued(limit int) ([]*models.Download, error) {
97 return s.repo.GetQueued(limit)
98}
99
100func (s *DownloadService) Delete(id int64) error {
101 s.killProcess(id)
102 s.cache.Delete(id)
103 return s.repo.Delete(id)
104}
105
106func (s *DownloadService) killProcess(id int64) {
107 s.processMu.Lock()
108 proc, ok := s.processes[id]
109 delete(s.processes, id)
110 s.processMu.Unlock()
111
112 if !ok || proc == nil {
113 return
114 }
115
116 _ = syscall.Kill(-proc.Pid, syscall.SIGKILL)
117}
118
119func (s *DownloadService) DeleteAll() error {
120 return s.repo.DeleteAll()
121}
122
123func (s *DownloadService) ResetStalledDownloads() error {
124 return s.repo.UpdateStatusWhere("downloading", "queued")
125}
126
127func (s *DownloadService) ListFormats(url string) ([]*models.FormatInfo, error) {
128 cmd := exec.Command(s.cfg.YTDLPPath, "-F", "--no-warnings", url)
129 output, err := cmd.CombinedOutput()
130 if err != nil {
131 return nil, fmt.Errorf("yt-dlp -F failed: %w\nOutput: %s", err, string(output))
132 }
133
134 return parseFormatList(string(output)), nil
135}
136
137// ExecuteDownload runs the download for d. The bool reports whether this call
138// actually processed it: false means another worker already claimed it (Submit
139// and the queue checker can both enqueue the same row within the 2s poll window),
140// so the caller should not log it as completed.
141func (s *DownloadService) ExecuteDownload(d *models.Download) (bool, error) {
142 // Atomically claim the download. If it's no longer queued, another worker
143 // already took it — bail rather than download it twice.
144 claimed, err := s.repo.MarkStarted(d.ID)
145 if err != nil {
146 return false, err
147 }
148 if !claimed {
149 return false, nil
150 }
151
152 s.cache.Set(d.ID, &LiveDownload{LastUpdate: time.Now()})
153 defer s.cache.Delete(d.ID)
154
155 var preset *models.Preset
156
157 if d.PresetID.Valid {
158 preset, err = s.presetSvc.GetByID(d.PresetID.Int64)
159 if err != nil {
160 preset, _ = s.presetSvc.GetDefault()
161 }
162 } else {
163 preset, _ = s.presetSvc.GetDefault()
164 }
165
166 if preset == nil {
167 preset = &models.Preset{}
168 }
169
170 tempDownloadDir := filepath.Join(s.cfg.TempDir, fmt.Sprintf("%d", d.ID))
171 if err := os.MkdirAll(tempDownloadDir, 0755); err != nil {
172 return false, fmt.Errorf("create temp download dir: %w", err)
173 }
174
175 args := s.presetSvc.BuildArgs(preset, d.FormatOverride, d.CustomFlags)
176
177 cookies, err := s.settingsSvc.GetCookies()
178 if err == nil && strings.TrimSpace(cookies) != "" {
179 tmpFile, err := os.CreateTemp("", "cookies-*.txt")
180 if err == nil {
181 tmpFile.WriteString(cookies)
182 tmpFile.Close()
183 args = append(args, "--cookies", tmpFile.Name())
184 defer os.Remove(tmpFile.Name())
185 }
186 }
187
188 args = append(args, "-P", tempDownloadDir)
189 args = append(args, "-o", "item-%(autonumber)05d/%(title)s.%(ext)s")
190 args = append(args, d.URL)
191
192 cmd := exec.Command(s.cfg.YTDLPPath, args...)
193 cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
194
195 stdout, err := cmd.StdoutPipe()
196 if err != nil {
197 s.finalizeError(d.ID, err)
198 return false, err
199 }
200 cmd.Stderr = cmd.Stdout
201
202 if err := cmd.Start(); err != nil {
203 s.finalizeError(d.ID, err)
204 return false, err
205 }
206
207 s.processMu.Lock()
208 s.processes[d.ID] = cmd.Process
209 s.processMu.Unlock()
210 defer func() {
211 s.processMu.Lock()
212 delete(s.processes, d.ID)
213 s.processMu.Unlock()
214 }()
215
216 scanner := bufio.NewScanner(stdout)
217
218 ticker := time.NewTicker(10 * time.Second)
219 defer ticker.Stop()
220 done := make(chan struct{})
221 go func() {
222 for {
223 select {
224 case <-ticker.C:
225 logs := s.cache.FlushLogs(d.ID)
226 if logs != "" {
227 s.repo.AppendLogs(d.ID, logs)
228 }
229 case <-done:
230 return
231 }
232 }
233 }()
234
235 for scanner.Scan() {
236 line := scanner.Text()
237 s.cache.AppendLog(d.ID, line)
238 }
239
240 close(done)
241
242 logs := s.cache.FlushLogs(d.ID)
243 if logs != "" {
244 s.repo.AppendLogs(d.ID, logs)
245 }
246
247 if err := cmd.Wait(); err != nil {
248 s.finalizeError(d.ID, err)
249 return false, err
250 }
251
252 if err := s.repo.MarkCompleted(d.ID, "completed"); err != nil {
253 return false, err
254 }
255
256 if err := s.importDownloadedItems(d, tempDownloadDir); err != nil {
257 log.Printf("Download %d completed but import failed: %v", d.ID, err)
258 }
259
260 return true, nil
261}
262
263func (s *DownloadService) finalizeError(id int64, err error) {
264 logs := s.cache.FlushLogs(id)
265 if logs != "" {
266 s.repo.AppendLogs(id, logs)
267 }
268 s.repo.MarkError(id, err.Error())
269}
270
271func (s *DownloadService) importDownloadedItems(d *models.Download, tempDownloadDir string) error {
272 entries, err := os.ReadDir(tempDownloadDir)
273 if err != nil {
274 return err
275 }
276
277 baseLibraryDir := s.cfg.LibraryDir
278 if d.OutputDir.Valid && d.OutputDir.String != "" {
279 cleanDir := filepath.Clean(d.OutputDir.String)
280 fullPath := filepath.Join(baseLibraryDir, cleanDir)
281 resolvedPath, err := filepath.Abs(fullPath)
282 if err != nil {
283 return fmt.Errorf("invalid output directory: %w", err)
284 }
285 resolvedLibraryDir, _ := filepath.Abs(baseLibraryDir)
286 if !strings.HasPrefix(resolvedPath, resolvedLibraryDir+string(filepath.Separator)) && resolvedPath != resolvedLibraryDir {
287 return fmt.Errorf("invalid output directory: path traversal attempt detected")
288 }
289 baseLibraryDir = fullPath
290 }
291 if err := os.MkdirAll(baseLibraryDir, 0755); err != nil {
292 return err
293 }
294
295 var itemDirs []string
296 for _, entry := range entries {
297 if !entry.IsDir() {
298 continue
299 }
300 name := entry.Name()
301 if strings.HasPrefix(name, "item-") {
302 itemDirs = append(itemDirs, filepath.Join(tempDownloadDir, name))
303 }
304 }
305 sort.Strings(itemDirs)
306
307 for _, itemDir := range itemDirs {
308 if err := s.importItemDir(d.URL, itemDir, baseLibraryDir); err != nil {
309 log.Printf("warning: failed to import item %s: %v", itemDir, err)
310 }
311 }
312
313 os.Remove(tempDownloadDir)
314 return nil
315}
316
317func (s *DownloadService) importItemDir(url, itemDir, baseLibraryDir string) error {
318 entries, err := os.ReadDir(itemDir)
319 if err != nil {
320 return err
321 }
322
323 var mediaFiles []os.DirEntry
324 var infoJSONPath string
325 var subtitleFiles []string
326
327 for _, entry := range entries {
328 if entry.IsDir() {
329 continue
330 }
331 name := entry.Name()
332 path := filepath.Join(itemDir, name)
333 ext := strings.ToLower(filepath.Ext(name))
334
335 if name == "info.json" || strings.HasSuffix(name, ".info.json") {
336 infoJSONPath = path
337 continue
338 }
339 if ext == ".vtt" || ext == ".srt" || ext == ".ass" || ext == ".ssa" {
340 subtitleFiles = append(subtitleFiles, path)
341 continue
342 }
343
344 mtype, err := mimetype.DetectFile(path)
345 if err == nil && mtype != nil && (strings.HasPrefix(mtype.String(), "audio/") || strings.HasPrefix(mtype.String(), "video/")) {
346 mediaFiles = append(mediaFiles, entry)
347 }
348 }
349
350 if len(mediaFiles) == 0 {
351 return fmt.Errorf("no media files found in %s", itemDir)
352 }
353
354 name := s.deriveItemName(itemDir, infoJSONPath, mediaFiles)
355 targetDir := s.uniqueDir(baseLibraryDir, name)
356 if err := os.MkdirAll(targetDir, 0755); err != nil {
357 return err
358 }
359
360 if infoJSONPath != "" {
361 if err := os.Rename(infoJSONPath, filepath.Join(targetDir, "info.json")); err != nil {
362 return err
363 }
364 }
365
366 for _, entry := range mediaFiles {
367 if err := os.Rename(filepath.Join(itemDir, entry.Name()), filepath.Join(targetDir, entry.Name())); err != nil {
368 return err
369 }
370 }
371
372 // Probe each media file's duration once, here in the worker (off the request
373 // path), and cache it in the marker so the library never has to probe while
374 // serving pages. Files we can't probe simply get no duration.
375 fileDurations := make(map[string]int)
376 for _, entry := range mediaFiles {
377 if d, ok := probeDuration(filepath.Join(targetDir, entry.Name())); ok {
378 fileDurations[entry.Name()] = d
379 }
380 }
381
382 if len(subtitleFiles) > 0 {
383 subtitlesDir := filepath.Join(targetDir, subtitlesDirName)
384 if err := os.MkdirAll(subtitlesDir, 0755); err != nil {
385 return err
386 }
387 for _, sf := range subtitleFiles {
388 if err := os.Rename(sf, filepath.Join(subtitlesDir, filepath.Base(sf))); err != nil {
389 return err
390 }
391 }
392 }
393
394 metadata := models.ItemMetadata{
395 Name: name,
396 SourceURL: url,
397 FileDurations: fileDurations,
398 }
399
400 markerPath := filepath.Join(targetDir, itemMarkerName)
401 f, err := os.Create(markerPath)
402 if err != nil {
403 return err
404 }
405 defer f.Close()
406 if err := toml.NewEncoder(f).Encode(metadata); err != nil {
407 return err
408 }
409
410 return nil
411}
412
413func (s *DownloadService) deriveItemName(itemDir, infoJSONPath string, mediaFiles []os.DirEntry) string {
414 if infoJSONPath != "" {
415 data, err := os.ReadFile(infoJSONPath)
416 if err == nil {
417 var info struct {
418 Title string `json:"title"`
419 }
420 if err := json.Unmarshal(data, &info); err == nil && info.Title != "" {
421 return sanitizeDirName(info.Title)
422 }
423 }
424 }
425
426 sort.Slice(mediaFiles, func(i, j int) bool {
427 ii, _ := os.Stat(filepath.Join(itemDir, mediaFiles[i].Name()))
428 jj, _ := os.Stat(filepath.Join(itemDir, mediaFiles[j].Name()))
429 if ii == nil || jj == nil {
430 return false
431 }
432 return ii.Size() > jj.Size()
433 })
434
435 base := strings.TrimSuffix(mediaFiles[0].Name(), filepath.Ext(mediaFiles[0].Name()))
436 return sanitizeDirName(base)
437}
438
439func (s *DownloadService) uniqueDir(base, name string) string {
440 dir := filepath.Join(base, name)
441 if _, err := os.Stat(dir); os.IsNotExist(err) {
442 return dir
443 }
444 for i := 1; ; i++ {
445 candidate := fmt.Sprintf("%s-%d", dir, i)
446 if _, err := os.Stat(candidate); os.IsNotExist(err) {
447 return candidate
448 }
449 }
450}
451
452func sanitizeDirName(name string) string {
453 name = strings.TrimSpace(name)
454 replacer := strings.NewReplacer(
455 "/", "-",
456 "\\", "-",
457 ":", "-",
458 "*", "-",
459 "?", "-",
460 "\"", "-",
461 "<", "-",
462 ">", "-",
463 "|", "-",
464 )
465 name = replacer.Replace(name)
466 name = strings.TrimSpace(name)
467 if name == "" {
468 name = "untitled"
469 }
470 return name
471}
472
473// resolutionRe matches a yt-dlp resolution column like "1920x1080".
474var resolutionRe = regexp.MustCompile(`^\d+x\d+$`)
475
476// probeDuration returns the duration of a media file in whole seconds. The bool
477// is false when ffprobe is unavailable or the file has no usable duration.
478func probeDuration(path string) (int, bool) {
479 out, err := exec.Command("ffprobe", "-v", "error",
480 "-show_entries", "format=duration",
481 "-of", "default=nw=1:nk=1", path).Output()
482 if err != nil {
483 return 0, false
484 }
485 f, err := strconv.ParseFloat(strings.TrimSpace(string(out)), 64)
486 if err != nil || f <= 0 {
487 return 0, false
488 }
489 return int(f + 0.5), true
490}
491
492func parseFormatList(output string) []*models.FormatInfo {
493 lines := strings.Split(output, "\n")
494 var formats []*models.FormatInfo
495
496 inFormats := false
497 for _, line := range lines {
498 line = strings.TrimSpace(line)
499 if line == "" {
500 continue
501 }
502
503 if strings.Contains(line, "ID") && strings.Contains(line, "EXT") {
504 inFormats = true
505 continue
506 }
507
508 if !inFormats {
509 continue
510 }
511
512 parts := strings.Fields(line)
513 if len(parts) >= 4 {
514 format := &models.FormatInfo{
515 ID: parts[0],
516 Ext: parts[1],
517 }
518
519 for i, part := range parts {
520 // A resolution token is strictly <digits>x<digits> (e.g. 1920x1080);
521 // matching on a literal "x" anywhere misclassified notes/codecs.
522 if resolutionRe.MatchString(part) {
523 format.Resolution = part
524 if i+1 < len(parts) {
525 format.FPS = parts[i+1]
526 }
527 break
528 }
529 }
530
531 if len(parts) > 3 {
532 format.Note = strings.Join(parts[3:], " ")
533 }
534
535 formats = append(formats, format)
536 }
537 }
538
539 return formats
540}
541
542func sqlNullInt64(v int64) sql.NullInt64 {
543 return sql.NullInt64{Int64: v, Valid: true}
544}
545