pool.go
⎇
Raw
1package worker
2
3import (
4 "context"
5 "fmt"
6 "log"
7 "time"
8
9 "vidarchive/internal/models"
10 "vidarchive/internal/service"
11)
12
13type Pool struct {
14 downloadSvc *service.DownloadService
15 workers int
16 queue chan *models.Download
17 ctx context.Context
18 cancel context.CancelFunc
19}
20
21func New(downloadSvc *service.DownloadService, workers int) *Pool {
22 ctx, cancel := context.WithCancel(context.Background())
23 return &Pool{
24 downloadSvc: downloadSvc,
25 workers: workers,
26 queue: make(chan *models.Download, 100),
27 ctx: ctx,
28 cancel: cancel,
29 }
30}
31
32func (p *Pool) Start() {
33 // Reset any downloads that were in progress during a previous run
34 if err := p.downloadSvc.ResetStalledDownloads(); err != nil {
35 log.Printf("Warning: failed to reset stalled downloads: %v", err)
36 }
37
38 for i := 0; i < p.workers; i++ {
39 go p.worker(i)
40 }
41
42 // Queue checker - polls DB for queued downloads
43 go p.queueChecker()
44}
45
46func (p *Pool) Stop() {
47 p.cancel()
48}
49
50func (p *Pool) Submit(d *models.Download) {
51 select {
52 case p.queue <- d:
53 log.Printf("Download %d queued", d.ID)
54 case <-p.ctx.Done():
55 }
56}
57
58func (p *Pool) worker(id int) {
59 log.Printf("Worker %d started", id)
60 for {
61 select {
62 case d := <-p.queue:
63 if d == nil {
64 continue
65 }
66 log.Printf("Worker %d processing download %d", id, d.ID)
67 processed, err := p.downloadSvc.ExecuteDownload(d)
68 switch {
69 case err != nil:
70 log.Printf("Worker %d download %d failed: %v", id, d.ID, err)
71 case processed:
72 log.Printf("Worker %d download %d completed", id, d.ID)
73 default:
74 log.Printf("Worker %d download %d already claimed by another worker, skipping", id, d.ID)
75 }
76 case <-p.ctx.Done():
77 log.Printf("Worker %d stopped", id)
78 return
79 }
80 }
81}
82
83func (p *Pool) queueChecker() {
84 ticker := time.NewTicker(2 * time.Second)
85 defer ticker.Stop()
86
87 for {
88 select {
89 case <-ticker.C:
90 p.checkQueue()
91 case <-p.ctx.Done():
92 return
93 }
94 }
95}
96
97func (p *Pool) checkQueue() {
98 // Pull at most a channel's worth of the oldest queued downloads rather than
99 // loading every queued row each tick. Anything beyond the buffer is picked up
100 // on a later tick; duplicates are harmless (MarkStarted claims atomically).
101 downloads, err := p.downloadSvc.GetQueued(cap(p.queue))
102 if err != nil {
103 log.Printf("Queue check error: %v", err)
104 return
105 }
106
107 for _, d := range downloads {
108 // Try to submit - if queue is full, it'll block briefly
109 select {
110 case p.queue <- d:
111 default:
112 // Queue is full, skip for now
113 return
114 }
115 }
116}
117
118func (p *Pool) GetQueueStatus() (queued, active, completed, failed int, err error) {
119 all, err := p.downloadSvc.GetAll("all", "date")
120 if err != nil {
121 return 0, 0, 0, 0, fmt.Errorf("get all downloads: %w", err)
122 }
123
124 for _, d := range all {
125 switch d.Status {
126 case "queued":
127 queued++
128 case "downloading":
129 active++
130 case "completed":
131 completed++
132 case "error":
133 failed++
134 }
135 }
136
137 return queued, active, completed, failed, nil
138}
139