package service import ( "bufio" "database/sql" "fmt" "os" "os/exec" "path/filepath" "strings" "sync" "syscall" "time" "github.com/gabriel-vasile/mimetype" "vidarchive/internal/config" "vidarchive/internal/models" "vidarchive/internal/repository" ) type DownloadService struct { repo *repository.DownloadRepository mediaRepo *repository.MediaRepository presetSvc *PresetService settingsSvc *SettingsService cfg *config.Config cache *ProgressCache processMu sync.Mutex processes map[int64]*os.Process } func NewDownloadService(repo *repository.DownloadRepository, mediaRepo *repository.MediaRepository, presetSvc *PresetService, settingsSvc *SettingsService, cfg *config.Config) *DownloadService { return &DownloadService{ repo: repo, mediaRepo: mediaRepo, presetSvc: presetSvc, settingsSvc: settingsSvc, cfg: cfg, cache: NewProgressCache(), processes: make(map[int64]*os.Process), } } func (s *DownloadService) Create(url string, presetID *int64, formatOverride, customFlags, outputDir string) (*models.Download, error) { d := &models.Download{ URL: url, Status: "queued", FormatOverride: formatOverride, CustomFlags: customFlags, OutputDir: sql.NullString{String: outputDir, Valid: outputDir != ""}, } if presetID != nil { d.PresetID = sqlNullInt64(*presetID) } if err := s.repo.Create(d); err != nil { return nil, err } return d, nil } func (s *DownloadService) GetByID(id int64) (*models.Download, error) { // Check cache first for live logs if live, ok := s.cache.Get(id); ok { d, err := s.repo.GetByID(id) if err != nil { return nil, err } logs := live.Logs.String() if logs != "" { d.Logs = sql.NullString{String: logs, Valid: true} } return d, nil } return s.repo.GetByID(id) } func (s *DownloadService) GetAll(status, sortBy string) ([]*models.Download, error) { downloads, err := s.repo.GetAll(status, sortBy) if err != nil { return nil, err } // Overlay live logs from cache for _, d := range downloads { if live, ok := s.cache.Get(d.ID); ok { logs := live.Logs.String() if logs != "" { d.Logs = sql.NullString{String: logs, Valid: true} } } } return downloads, nil } func (s *DownloadService) GetQueued(limit int) ([]*models.Download, error) { return s.repo.GetQueued(limit) } func (s *DownloadService) Delete(id int64) error { s.killProcess(id) s.cache.Delete(id) return s.repo.Delete(id) } func (s *DownloadService) killProcess(id int64) { s.processMu.Lock() proc, ok := s.processes[id] delete(s.processes, id) s.processMu.Unlock() if !ok || proc == nil { return } // Kill the process group (negative PID kills the group on Linux) // The process group was created by Setpgid in ExecuteDownload _ = syscall.Kill(-proc.Pid, syscall.SIGKILL) } func (s *DownloadService) DeleteAll() error { return s.repo.DeleteAll() } func (s *DownloadService) ResetStalledDownloads() error { return s.repo.UpdateStatusWhere("downloading", "queued") } func (s *DownloadService) ListFormats(url string) ([]*models.FormatInfo, error) { cmd := exec.Command(s.cfg.YTDLPPath, "-F", "--no-warnings", url) output, err := cmd.CombinedOutput() if err != nil { return nil, fmt.Errorf("yt-dlp -F failed: %w\nOutput: %s", err, string(output)) } return parseFormatList(string(output)), nil } func (s *DownloadService) ExecuteDownload(d *models.Download) error { if err := s.repo.MarkStarted(d.ID); err != nil { return err } // Initialize in-memory cache for logs s.cache.Set(d.ID, &LiveDownload{ LastUpdate: time.Now(), }) defer s.cache.Delete(d.ID) var preset *models.Preset var err error if d.PresetID.Valid { preset, err = s.presetSvc.GetByID(d.PresetID.Int64) if err != nil { preset, _ = s.presetSvc.GetDefault() } } else { preset, _ = s.presetSvc.GetDefault() } if preset == nil { preset = &models.Preset{ OutputTemplate: "%(title)s.%(ext)s", } } // Build output path with optional subdirectory downloadDir := s.cfg.DownloadDir if d.OutputDir.Valid && d.OutputDir.String != "" { // Sanitize and validate output directory cleanDir := filepath.Clean(d.OutputDir.String) // Prevent path traversal: ensure the cleaned path doesn't escape the download dir fullPath := filepath.Join(downloadDir, cleanDir) resolvedPath, err := filepath.Abs(fullPath) if err != nil { return fmt.Errorf("invalid output directory: %w", err) } resolvedDownloadDir, _ := filepath.Abs(downloadDir) if !strings.HasPrefix(resolvedPath, resolvedDownloadDir+string(filepath.Separator)) && resolvedPath != resolvedDownloadDir { return fmt.Errorf("invalid output directory: path traversal attempt detected") } downloadDir = fullPath // Ensure directory exists if err := os.MkdirAll(downloadDir, 0755); err != nil { return fmt.Errorf("create output directory: %w", err) } } args := s.presetSvc.BuildArgs(preset, d.FormatOverride, d.CustomFlags) // Write cookies to temp file if configured cookies, err := s.settingsSvc.GetCookies() if err == nil && strings.TrimSpace(cookies) != "" { tmpFile, err := os.CreateTemp("", "cookies-*.txt") if err == nil { tmpFile.WriteString(cookies) tmpFile.Close() args = append(args, "--cookies", tmpFile.Name()) defer os.Remove(tmpFile.Name()) } } args = append(args, "-P", downloadDir) args = append(args, d.URL) cmd := exec.Command(s.cfg.YTDLPPath, args...) cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} stdout, err := cmd.StdoutPipe() if err != nil { s.finalizeError(d.ID, err) return err } cmd.Stderr = cmd.Stdout if err := cmd.Start(); err != nil { s.finalizeError(d.ID, err) return err } // Track the process so it can be killed on delete s.processMu.Lock() s.processes[d.ID] = cmd.Process s.processMu.Unlock() defer func() { s.processMu.Lock() delete(s.processes, d.ID) s.processMu.Unlock() }() scanner := bufio.NewScanner(stdout) ticker := time.NewTicker(10 * time.Second) defer ticker.Stop() done := make(chan struct{}) go func() { for { select { case <-ticker.C: logs := s.cache.FlushLogs(d.ID) if logs != "" { s.repo.UpdateLogs(d.ID, logs) } case <-done: return } } }() for scanner.Scan() { line := scanner.Text() s.cache.AppendLog(d.ID, line) } close(done) // Final log flush logs := s.cache.FlushLogs(d.ID) if logs != "" { s.repo.UpdateLogs(d.ID, logs) } if err := cmd.Wait(); err != nil { s.finalizeError(d.ID, err) return err } // Mark as completed if err := s.repo.MarkCompleted(d.ID, "completed"); err != nil { return err } // Scan for new media files s.scanForNewMedia(d.URL, downloadDir) return nil } func (s *DownloadService) finalizeError(id int64, err error) { logs := s.cache.FlushLogs(id) if logs != "" { s.repo.UpdateLogs(id, logs) } s.repo.MarkError(id, err.Error()) } func (s *DownloadService) scanForNewMedia(url string, downloadDir string) { // Give filesystem a moment time.Sleep(100 * time.Millisecond) // Walk download dir and add any new files filepath.Walk(downloadDir, func(path string, info os.FileInfo, err error) error { if err != nil || info.IsDir() { return nil } relPath, _ := filepath.Rel(s.cfg.DownloadDir, path) relPath = filepath.ToSlash(relPath) _, err = s.mediaRepo.GetByRelativePath(relPath) if err == nil { return nil // already exists } // Use content-based MIME detection mtype, err := mimetype.DetectFile(path) if err != nil { return nil } isAudio := mtype != nil && strings.HasPrefix(mtype.String(), "audio/") isVideo := mtype != nil && strings.HasPrefix(mtype.String(), "video/") if !isAudio && !isVideo { return nil } ext := strings.ToLower(filepath.Ext(path)) basePath := strings.TrimSuffix(path, ext) infoJSONPath := basePath + ".info.json" if _, err := os.Stat(infoJSONPath); err != nil { infoJSONPath = "" } media := &models.Media{ URL: url, Filepath: path, RelativePath: relPath, IsAudio: isAudio, HasEmbeddedThumbnail: true, InfoJSONPath: infoJSONPath, Title: strings.TrimSuffix(filepath.Base(path), ext), Duration: extractDuration(infoJSONPath, path), } s.mediaRepo.Create(media) return nil }) } func parseFormatList(output string) []*models.FormatInfo { lines := strings.Split(output, "\n") var formats []*models.FormatInfo inFormats := false for _, line := range lines { line = strings.TrimSpace(line) if line == "" { continue } if strings.Contains(line, "ID") && strings.Contains(line, "EXT") { inFormats = true continue } if !inFormats { continue } parts := strings.Fields(line) if len(parts) >= 4 { format := &models.FormatInfo{ ID: parts[0], Ext: parts[1], } for i, part := range parts { if strings.Contains(part, "x") && !strings.Contains(part, "http") { format.Resolution = part if i+1 < len(parts) { format.FPS = parts[i+1] } break } } if len(parts) > 3 { format.Note = strings.Join(parts[3:], " ") } formats = append(formats, format) } } return formats } func sqlNullInt64(v int64) sql.NullInt64 { return sql.NullInt64{Int64: v, Valid: true} }