summaryrefslogtreecommitdiffstats
path: root/progress.go
blob: ddc810a85e83ae97a680637623874c82efcbc8a4 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
// Package progress implements a simple command line progress bar.
//
// Example:
//
//	import (
//		time"
//
//		"gitlab.com/rhogenson/progress-bar"
//	)
//
//	func main() {
//		b := new(progress.Bar)
//		for i := range 100 {
//			b.Set(float64(i)/100)
//			b.Print()
//			time.Sleep(time.Second)
//		}
//	}
package progress

import (
	"fmt"
	"os"
	"strings"
	"time"

	"gitlab.com/rhogenson/vecdeque"
	"golang.org/x/term"
)

const measurements = 20

type measurement struct {
	t time.Time
	i float64
}

// Bar is a progess bar. The zero value is ready for use.
type Bar struct {
	measurements vecdeque.DQ[measurement]
	cols         int
}

// Set sets the current value to val. val must be between 0 and 1, inclusive.
func (b *Bar) Set(val float64) {
	if !(0 <= val && val <= 1) {
		panic(fmt.Sprintf("progress.Bar.Set: value must be between 0 and 1 (got %f)", val))
	}
	if b.measurements.Len() == measurements {
		b.measurements.PopFront()
	}
	b.measurements.PushBack(measurement{time.Now(), val})
}

// Print shows the progress bar on standard error.
func (b *Bar) Print() {
	if b.measurements.Len() == 0 {
		return
	}

	first := b.measurements.Get(0)
	last := b.measurements.Get(b.measurements.Len() - 1)
	deltaT := last.t.Sub(first.t)
	delta := last.i - first.i
	eta := time.Duration(-1)
	if delta != 0 {
		eta = time.Duration(float64(deltaT) * (1 - last.i) / delta)
	}

	p := int(last.i * float64(b.cols))
	if b.cols == 0 {
		var err error
		if b.cols, _, err = term.GetSize(int(os.Stderr.Fd())); err != nil {
			fmt.Fprintf(os.Stderr, "Warning: unable to determine terminal size: %s\n", err)
		} else {
			fmt.Fprintf(os.Stderr, "%s>\n", strings.Repeat("=", max(p-1, 0)))
		}
	} else {
		fmt.Fprintf(os.Stderr, "\033[2F\033[J%s>\n", strings.Repeat("=", max(p-1, 0)))
	}
	if eta < 0 {
		fmt.Fprintln(os.Stderr, "ETA: calculating...")
	} else {
		fmt.Fprintf(os.Stderr, "ETA: %s\n", eta.Round(time.Second))
	}
}