package worker import ( "context" "errors" "log" "sync" "time" "vidarchive/internal/models" "vidarchive/internal/service" ) type Pool struct { downloadSvc *service.DownloadService workers int queue chan *models.Download ctx context.Context cancel context.CancelFunc wg sync.WaitGroup } func New(downloadSvc *service.DownloadService, workers int) *Pool { ctx, cancel := context.WithCancel(context.Background()) return &Pool{ downloadSvc: downloadSvc, workers: workers, queue: make(chan *models.Download, 100), ctx: ctx, cancel: cancel, } } func (p *Pool) Start() { // Reset any downloads that were in progress during a previous run, and clear // the temp dirs they left behind so a re-run doesn't import a duplicate. if err := p.downloadSvc.ResetStalledDownloads(); err != nil { log.Printf("Warning: failed to reset stalled downloads: %v", err) } for i := 0; i < p.workers; i++ { p.wg.Add(1) go p.worker(i) } p.wg.Add(1) go p.queueChecker() } // Stop cancels running downloads and waits for the workers to return, so the // process doesn't exit while a yt-dlp child is still writing to the temp dir. func (p *Pool) Stop() { p.cancel() p.downloadSvc.CancelAll() p.wg.Wait() } // Submit enqueues a download without blocking. When the buffer is full the row // stays "queued" in the database and the queue checker picks it up on a later // tick, so a burst of submissions can't stall the HTTP handler that made them. func (p *Pool) Submit(d *models.Download) { if p.ctx.Err() != nil { return } select { case p.queue <- d: log.Printf("Download %d queued", d.ID) default: log.Printf("Download %d left in database queue (worker buffer full)", d.ID) } } func (p *Pool) worker(id int) { defer p.wg.Done() log.Printf("Worker %d started", id) for { select { case d := <-p.queue: log.Printf("Worker %d processing download %d", id, d.ID) processed, err := p.downloadSvc.ExecuteDownload(p.ctx, d) switch { case errors.Is(err, service.ErrCancelled): log.Printf("Worker %d download %d cancelled", id, d.ID) case err != nil: log.Printf("Worker %d download %d failed: %v", id, d.ID, err) case processed: log.Printf("Worker %d download %d completed", id, d.ID) default: log.Printf("Worker %d download %d already claimed by another worker, skipping", id, d.ID) } case <-p.ctx.Done(): log.Printf("Worker %d stopped", id) return } } } func (p *Pool) queueChecker() { defer p.wg.Done() ticker := time.NewTicker(2 * time.Second) defer ticker.Stop() for { select { case <-ticker.C: p.checkQueue() case <-p.ctx.Done(): return } } } func (p *Pool) checkQueue() { // Pull at most a channel's worth of the oldest queued downloads rather than // loading every queued row each tick. Anything beyond the buffer is picked up // on a later tick; duplicates are harmless (MarkStarted claims atomically). downloads, err := p.downloadSvc.GetQueued(cap(p.queue)) if err != nil { log.Printf("Queue check error: %v", err) return } for _, d := range downloads { select { case p.queue <- d: default: return } } }