download.go
⎇
Raw
1package service
2
3import (
4 "bufio"
5 "database/sql"
6 "fmt"
7 "os"
8 "os/exec"
9 "path/filepath"
10 "strings"
11 "sync"
12 "syscall"
13 "time"
14
15 "github.com/gabriel-vasile/mimetype"
16
17 "vidarchive/internal/config"
18 "vidarchive/internal/models"
19 "vidarchive/internal/repository"
20)
21
22type DownloadService struct {
23 repo *repository.DownloadRepository
24 mediaRepo *repository.MediaRepository
25 presetSvc *PresetService
26 settingsSvc *SettingsService
27 cfg *config.Config
28 cache *ProgressCache
29 processMu sync.Mutex
30 processes map[int64]*os.Process
31}
32
33func NewDownloadService(repo *repository.DownloadRepository, mediaRepo *repository.MediaRepository, presetSvc *PresetService, settingsSvc *SettingsService, cfg *config.Config) *DownloadService {
34 return &DownloadService{
35 repo: repo,
36 mediaRepo: mediaRepo,
37 presetSvc: presetSvc,
38 settingsSvc: settingsSvc,
39 cfg: cfg,
40 cache: NewProgressCache(),
41 processes: make(map[int64]*os.Process),
42 }
43}
44
45func (s *DownloadService) Create(url string, presetID *int64, formatOverride, customFlags, outputDir string) (*models.Download, error) {
46 d := &models.Download{
47 URL: url,
48 Status: "queued",
49 FormatOverride: formatOverride,
50 CustomFlags: customFlags,
51 OutputDir: sql.NullString{String: outputDir, Valid: outputDir != ""},
52 }
53
54 if presetID != nil {
55 d.PresetID = sqlNullInt64(*presetID)
56 }
57
58 if err := s.repo.Create(d); err != nil {
59 return nil, err
60 }
61 return d, nil
62}
63
64func (s *DownloadService) GetByID(id int64) (*models.Download, error) {
65 // Check cache first for live logs
66 if live, ok := s.cache.Get(id); ok {
67 d, err := s.repo.GetByID(id)
68 if err != nil {
69 return nil, err
70 }
71 logs := live.Logs.String()
72 if logs != "" {
73 d.Logs = sql.NullString{String: logs, Valid: true}
74 }
75 return d, nil
76 }
77 return s.repo.GetByID(id)
78}
79
80func (s *DownloadService) GetAll(status, sortBy string) ([]*models.Download, error) {
81 downloads, err := s.repo.GetAll(status, sortBy)
82 if err != nil {
83 return nil, err
84 }
85
86 // Overlay live logs from cache
87 for _, d := range downloads {
88 if live, ok := s.cache.Get(d.ID); ok {
89 logs := live.Logs.String()
90 if logs != "" {
91 d.Logs = sql.NullString{String: logs, Valid: true}
92 }
93 }
94 }
95
96 return downloads, nil
97}
98
99func (s *DownloadService) GetQueued(limit int) ([]*models.Download, error) {
100 return s.repo.GetQueued(limit)
101}
102
103func (s *DownloadService) Delete(id int64) error {
104 s.killProcess(id)
105 s.cache.Delete(id)
106 return s.repo.Delete(id)
107}
108
109func (s *DownloadService) killProcess(id int64) {
110 s.processMu.Lock()
111 proc, ok := s.processes[id]
112 delete(s.processes, id)
113 s.processMu.Unlock()
114
115 if !ok || proc == nil {
116 return
117 }
118
119 // Kill the process group (negative PID kills the group on Linux)
120 // The process group was created by Setpgid in ExecuteDownload
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 // Initialize in-memory cache for logs
148 s.cache.Set(d.ID, &LiveDownload{
149 LastUpdate: time.Now(),
150 })
151 defer s.cache.Delete(d.ID)
152
153 var preset *models.Preset
154 var err error
155
156 if d.PresetID.Valid {
157 preset, err = s.presetSvc.GetByID(d.PresetID.Int64)
158 if err != nil {
159 preset, _ = s.presetSvc.GetDefault()
160 }
161 } else {
162 preset, _ = s.presetSvc.GetDefault()
163 }
164
165 if preset == nil {
166 preset = &models.Preset{
167 OutputTemplate: "%(title)s.%(ext)s",
168 }
169 }
170
171 // Build output path with optional subdirectory
172 downloadDir := s.cfg.DownloadDir
173 if d.OutputDir.Valid && d.OutputDir.String != "" {
174 // Sanitize and validate output directory
175 cleanDir := filepath.Clean(d.OutputDir.String)
176 // Prevent path traversal: ensure the cleaned path doesn't escape the download dir
177 fullPath := filepath.Join(downloadDir, cleanDir)
178 resolvedPath, err := filepath.Abs(fullPath)
179 if err != nil {
180 return fmt.Errorf("invalid output directory: %w", err)
181 }
182 resolvedDownloadDir, _ := filepath.Abs(downloadDir)
183 if !strings.HasPrefix(resolvedPath, resolvedDownloadDir+string(filepath.Separator)) && resolvedPath != resolvedDownloadDir {
184 return fmt.Errorf("invalid output directory: path traversal attempt detected")
185 }
186 downloadDir = fullPath
187 // Ensure directory exists
188 if err := os.MkdirAll(downloadDir, 0755); err != nil {
189 return fmt.Errorf("create output directory: %w", err)
190 }
191 }
192
193 args := s.presetSvc.BuildArgs(preset, d.FormatOverride, d.CustomFlags)
194
195 // Write cookies to temp file if configured
196 cookies, err := s.settingsSvc.GetCookies()
197 if err == nil && strings.TrimSpace(cookies) != "" {
198 tmpFile, err := os.CreateTemp("", "cookies-*.txt")
199 if err == nil {
200 tmpFile.WriteString(cookies)
201 tmpFile.Close()
202 args = append(args, "--cookies", tmpFile.Name())
203 defer os.Remove(tmpFile.Name())
204 }
205 }
206
207 args = append(args, "-P", downloadDir)
208 args = append(args, d.URL)
209
210 cmd := exec.Command(s.cfg.YTDLPPath, args...)
211 cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
212
213 stdout, err := cmd.StdoutPipe()
214 if err != nil {
215 s.finalizeError(d.ID, err)
216 return err
217 }
218 cmd.Stderr = cmd.Stdout
219
220 if err := cmd.Start(); err != nil {
221 s.finalizeError(d.ID, err)
222 return err
223 }
224
225 // Track the process so it can be killed on delete
226 s.processMu.Lock()
227 s.processes[d.ID] = cmd.Process
228 s.processMu.Unlock()
229 defer func() {
230 s.processMu.Lock()
231 delete(s.processes, d.ID)
232 s.processMu.Unlock()
233 }()
234
235 scanner := bufio.NewScanner(stdout)
236
237 ticker := time.NewTicker(10 * time.Second)
238 defer ticker.Stop()
239 done := make(chan struct{})
240 go func() {
241 for {
242 select {
243 case <-ticker.C:
244 logs := s.cache.FlushLogs(d.ID)
245 if logs != "" {
246 s.repo.UpdateLogs(d.ID, logs)
247 }
248 case <-done:
249 return
250 }
251 }
252 }()
253
254 for scanner.Scan() {
255 line := scanner.Text()
256 s.cache.AppendLog(d.ID, line)
257 }
258
259 close(done)
260
261 // Final log flush
262 logs := s.cache.FlushLogs(d.ID)
263 if logs != "" {
264 s.repo.UpdateLogs(d.ID, logs)
265 }
266
267 if err := cmd.Wait(); err != nil {
268 s.finalizeError(d.ID, err)
269 return err
270 }
271
272 // Mark as completed
273 if err := s.repo.MarkCompleted(d.ID, "completed"); err != nil {
274 return err
275 }
276
277 // Scan for new media files
278 s.scanForNewMedia(d.URL, downloadDir)
279
280 return nil
281}
282
283func (s *DownloadService) finalizeError(id int64, err error) {
284 logs := s.cache.FlushLogs(id)
285 if logs != "" {
286 s.repo.UpdateLogs(id, logs)
287 }
288 s.repo.MarkError(id, err.Error())
289}
290
291func (s *DownloadService) scanForNewMedia(url string, downloadDir string) {
292 // Give filesystem a moment
293 time.Sleep(100 * time.Millisecond)
294
295 // Walk download dir and add any new files
296 filepath.Walk(downloadDir, func(path string, info os.FileInfo, err error) error {
297 if err != nil || info.IsDir() {
298 return nil
299 }
300
301 relPath, _ := filepath.Rel(s.cfg.DownloadDir, path)
302 relPath = filepath.ToSlash(relPath)
303
304 _, err = s.mediaRepo.GetByRelativePath(relPath)
305 if err == nil {
306 return nil // already exists
307 }
308
309 // Use content-based MIME detection
310 mtype, err := mimetype.DetectFile(path)
311 if err != nil {
312 return nil
313 }
314
315 isAudio := mtype != nil && strings.HasPrefix(mtype.String(), "audio/")
316 isVideo := mtype != nil && strings.HasPrefix(mtype.String(), "video/")
317
318 if !isAudio && !isVideo {
319 return nil
320 }
321
322 ext := strings.ToLower(filepath.Ext(path))
323 basePath := strings.TrimSuffix(path, ext)
324 infoJSONPath := basePath + ".info.json"
325 if _, err := os.Stat(infoJSONPath); err != nil {
326 infoJSONPath = ""
327 }
328
329 media := &models.Media{
330 URL: url,
331 Filepath: path,
332 RelativePath: relPath,
333 IsAudio: isAudio,
334 HasEmbeddedThumbnail: true,
335 InfoJSONPath: infoJSONPath,
336 Title: strings.TrimSuffix(filepath.Base(path), ext),
337 Duration: extractDuration(infoJSONPath, path),
338 }
339
340 s.mediaRepo.Create(media)
341 return nil
342 })
343}
344
345func parseFormatList(output string) []*models.FormatInfo {
346 lines := strings.Split(output, "\n")
347 var formats []*models.FormatInfo
348
349 inFormats := false
350 for _, line := range lines {
351 line = strings.TrimSpace(line)
352 if line == "" {
353 continue
354 }
355
356 if strings.Contains(line, "ID") && strings.Contains(line, "EXT") {
357 inFormats = true
358 continue
359 }
360
361 if !inFormats {
362 continue
363 }
364
365 parts := strings.Fields(line)
366 if len(parts) >= 4 {
367 format := &models.FormatInfo{
368 ID: parts[0],
369 Ext: parts[1],
370 }
371
372 for i, part := range parts {
373 if strings.Contains(part, "x") && !strings.Contains(part, "http") {
374 format.Resolution = part
375 if i+1 < len(parts) {
376 format.FPS = parts[i+1]
377 }
378 break
379 }
380 }
381
382 if len(parts) > 3 {
383 format.Note = strings.Join(parts[3:], " ")
384 }
385
386 formats = append(formats, format)
387 }
388 }
389
390 return formats
391}
392
393func sqlNullInt64(v int64) sql.NullInt64 {
394 return sql.NullInt64{Int64: v, Valid: true}
395}
396