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