all repos

clerk @ 2f348cb

missing tooling for ledger/hledger

clerk/internal/decimal/decimal.go (view raw)

Oleksandr Smirnov Oleksandr Smirnov
olexsmir@gmail.com
syntax: respect commodity format; =*; ===; :=, 2 months ago
1
package decimal
2
3
import (
4
	"fmt"
5
	"math/big"
6
	"strconv"
7
	"strings"
8
)
9
10
type Decimal struct {
11
	scale int
12
	coeff *big.Int
13
}
14
15
func FromInt(v int64) Decimal {
16
	if v == 0 {
17
		return Decimal{}
18
	}
19
	return Decimal{coeff: big.NewInt(v)}
20
}
21
22
func FromString(s string) (Decimal, error) {
23
	original := s
24
	if s == "" {
25
		return Decimal{}, badDecimal(original)
26
	}
27
28
	neg := false
29
	if s[0] == '+' || s[0] == '-' {
30
		neg = s[0] == '-'
31
		s = s[1:]
32
	}
33
34
	exp := 0
35
	if e := strings.IndexAny(s, "eE"); e >= 0 {
36
		n, err := strconv.Atoi(s[e+1:])
37
		if err != nil {
38
			return Decimal{}, badDecimal(original)
39
		}
40
		// bounded so pow10 can't be asked to materialize an astronomically
41
		// large number (e.g. 1E1000000000000)
42
		const maxExp = 10000
43
		if n > maxExp || n < -maxExp {
44
			return Decimal{}, badDecimal(original)
45
		}
46
		exp = n
47
		s = s[:e]
48
	}
49
50
	intPart, fracPart := s, ""
51
	if before, after, ok := strings.Cut(s, "."); ok {
52
		if strings.IndexByte(after, '.') >= 0 {
53
			return Decimal{}, badDecimal(original)
54
		}
55
		intPart, fracPart = before, after
56
	}
57
58
	digits := intPart + fracPart
59
	if neg {
60
		digits = "-" + digits
61
	}
62
	coeff, ok := new(big.Int).SetString(digits, 10)
63
	if !ok {
64
		return Decimal{}, badDecimal(original)
65
	}
66
67
	scale := len(fracPart) - exp
68
	if scale < 0 {
69
		coeff.Mul(coeff, pow10(-scale))
70
		scale = 0
71
	}
72
	return Decimal{coeff: coeff, scale: scale}.normalized(), nil
73
}
74
75
func badDecimal(s string) error { return fmt.Errorf("can't convert %s to decimal", s) }
76
77
func (d Decimal) String() string {
78
	if d.coeff == nil || d.coeff.Sign() == 0 {
79
		return "0"
80
	}
81
82
	abs := new(big.Int).Set(d.coeff)
83
	sign := ""
84
	if abs.Sign() < 0 {
85
		sign = "-"
86
		abs.Abs(abs)
87
	}
88
89
	digits := abs.String()
90
	if d.scale == 0 {
91
		return sign + digits
92
	}
93
94
	if len(digits) <= d.scale {
95
		digits = strings.Repeat("0", d.scale-len(digits)+1) + digits
96
	}
97
	split := len(digits) - d.scale
98
	return sign + digits[:split] + "." + digits[split:]
99
}
100
101
// StringFixed returns a string representation with exactly places digits
102
// after the decimal point. Pads with zeros or truncates as needed.
103
// decSep and thousandsSep control formatting; zero values mean no custom separator.
104
func (d Decimal) StringFixed(places int, decSep, thousandsSep byte) string {
105
	var sb strings.Builder
106
	sb.Grow(32)
107
	d.WriteFixed(&sb, places, decSep, thousandsSep)
108
	return sb.String()
109
}
110
111
// WriteFixed writes a string representation with exactly places digits
112
// after the decimal point directly into sb. Pads with zeros or truncates
113
// as needed. decSep and thousandsSep control formatting; zero values mean
114
// no custom separator.
115
func (d Decimal) WriteFixed(sb *strings.Builder, places int, decSep, thousandsSep byte) {
116
	if d.IsZero() {
117
		sb.WriteByte('0')
118
		if places > 0 {
119
			if decSep != 0 {
120
				sb.WriteByte(decSep)
121
			} else {
122
				sb.WriteByte('.')
123
			}
124
			for range places {
125
				sb.WriteByte('0')
126
			}
127
		}
128
		return
129
	}
130
131
	var digitBuf [128]byte
132
	digits := d.coeff.Append(digitBuf[:0], 10)
133
134
	sign := false
135
	if len(digits) > 0 && digits[0] == '-' {
136
		sign = true
137
		digits = digits[1:]
138
	}
139
140
	intLen := len(digits) - d.scale
141
142
	if sign {
143
		sb.WriteByte('-')
144
	}
145
146
	// Integer part
147
	if intLen <= 0 {
148
		sb.WriteByte('0')
149
	} else {
150
		for i := range intLen {
151
			if thousandsSep != 0 && i > 0 && (intLen-i)%3 == 0 {
152
				sb.WriteByte(thousandsSep)
153
			}
154
			sb.WriteByte(digits[i])
155
		}
156
	}
157
158
	// Fractional part
159
	if places > 0 {
160
		if decSep != 0 {
161
			sb.WriteByte(decSep)
162
		} else {
163
			sb.WriteByte('.')
164
		}
165
166
		written := 0
167
		if intLen > 0 {
168
			n := min(d.scale, places)
169
			for i := range n {
170
				sb.WriteByte(digits[intLen+i])
171
				written++
172
			}
173
		} else {
174
			leadingZeros := -intLen
175
			for ; written < leadingZeros && written < places; written++ {
176
				sb.WriteByte('0')
177
			}
178
			for i := 0; i < len(digits) && written < places; i++ {
179
				sb.WriteByte(digits[i])
180
				written++
181
			}
182
		}
183
		for ; written < places; written++ {
184
			sb.WriteByte('0')
185
		}
186
	}
187
}
188
189
func (d Decimal) Abs() Decimal {
190
	if d.coeff == nil || d.coeff.Sign() == 0 {
191
		return Decimal{}
192
	}
193
	if d.coeff.Sign() > 0 {
194
		return Decimal{coeff: new(big.Int).Set(d.coeff), scale: d.scale}
195
	}
196
	return Decimal{coeff: new(big.Int).Neg(d.coeff), scale: d.scale}
197
}
198
199
func (d Decimal) Neg() Decimal {
200
	if d.coeff == nil || d.coeff.Sign() == 0 {
201
		return Decimal{}
202
	}
203
	return Decimal{coeff: new(big.Int).Neg(d.coeff), scale: d.scale}
204
}
205
206
func (d Decimal) Sub(other Decimal) Decimal { return d.Add(other.Neg()) }
207
func (d Decimal) Add(other Decimal) Decimal {
208
	a, b, scale := align(d, other)
209
	sum := new(big.Int).Add(a, b)
210
	return Decimal{coeff: sum, scale: scale}.normalized()
211
}
212
213
func (d Decimal) Mul(other Decimal) Decimal {
214
	if d.IsZero() || other.IsZero() {
215
		return Decimal{}
216
	}
217
	product := new(big.Int).Mul(d.coeffOrZero(), other.coeffOrZero())
218
	return Decimal{coeff: product, scale: d.scale + other.scale}.normalized()
219
}
220
221
func (d Decimal) Div(other Decimal) Decimal {
222
	if other.IsZero() {
223
		panic("decimal: division by zero")
224
	}
225
	if d.IsZero() {
226
		return Decimal{}
227
	}
228
229
	scale := max(d.scale, other.scale) + 10
230
231
	dCoeff := d.coeffOrZero()
232
	oCoeff := other.coeffOrZero()
233
234
	shift := scale + other.scale - d.scale
235
	if shift > 0 {
236
		dCoeff = new(big.Int).Mul(dCoeff, pow10(shift))
237
	} else if shift < 0 {
238
		oCoeff = new(big.Int).Mul(oCoeff, pow10(-shift))
239
	}
240
241
	quo := new(big.Int).Quo(dCoeff, oCoeff)
242
	return Decimal{coeff: quo, scale: scale}.normalized()
243
}
244
245
func (d Decimal) Cmp(other Decimal) int {
246
	a, b, _ := align(d, other)
247
	return a.Cmp(b)
248
}
249
250
func (d Decimal) Equal(other Decimal) bool {
251
	return d.Cmp(other) == 0
252
}
253
254
func (d Decimal) IsZero() bool {
255
	return d.coeff == nil || d.coeff.Sign() == 0
256
}
257
258
func (d Decimal) coeffOrZero() *big.Int {
259
	if d.coeff == nil {
260
		return new(big.Int)
261
	}
262
	return new(big.Int).Set(d.coeff)
263
}
264
265
func (d Decimal) normalized() Decimal {
266
	if d.coeff == nil || d.coeff.Sign() == 0 {
267
		return Decimal{}
268
	}
269
	if d.scale == 0 {
270
		return Decimal{coeff: new(big.Int).Set(d.coeff)}
271
	}
272
273
	sign := d.coeff.Sign()
274
	abs := new(big.Int).Abs(d.coeff)
275
	ten := big.NewInt(10)
276
	rem := new(big.Int)
277
	for d.scale > 0 {
278
		quotient, _ := new(big.Int).QuoRem(abs, ten, rem)
279
		if rem.Sign() != 0 {
280
			break
281
		}
282
		abs = quotient
283
		d.scale--
284
	}
285
286
	if sign < 0 {
287
		abs.Neg(abs)
288
	}
289
	return Decimal{coeff: abs, scale: d.scale}
290
}
291
292
func align(a, b Decimal) (aCoeff *big.Int, bCoeff *big.Int, scale int) {
293
	scale = max(b.scale, a.scale)
294
295
	aCoeff = a.coeffOrZero()
296
	bCoeff = b.coeffOrZero()
297
	if delta := scale - a.scale; delta > 0 {
298
		aCoeff.Mul(aCoeff, pow10(delta))
299
	}
300
	if delta := scale - b.scale; delta > 0 {
301
		bCoeff.Mul(bCoeff, pow10(delta))
302
	}
303
	return aCoeff, bCoeff, scale
304
}
305
306
func pow10(n int) *big.Int {
307
	if n <= 0 {
308
		return big.NewInt(1)
309
	}
310
	return new(big.Int).Exp(big.NewInt(10), big.NewInt(int64(n)), nil)
311
}