all repos

clerk @ f2dd916162e462998dbd1b2fa9d0cc8ce9529226

missing tooling for ledger/hledger

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

Oleksandr Smirnov Oleksandr Smirnov
olexsmir@gmail.com
decimal: int64 fast lane for small amounts, big.Int fallback, 1 month ago
1
package decimal
2
3
import (
4
	"fmt"
5
	"math"
6
	"math/big"
7
	"strconv"
8
	"strings"
9
)
10
11
type Decimal struct {
12
	v     int64
13
	scale int
14
	big   *big.Int
15
}
16
17
func FromInt(v int64) Decimal {
18
	if v == 0 {
19
		return Decimal{}
20
	}
21
	return Decimal{v: v}
22
}
23
24
func FromString(s string) (Decimal, error) {
25
	original := s
26
	if s == "" {
27
		return Decimal{}, badDecimal(original)
28
	}
29
30
	neg := false
31
	if s[0] == '+' || s[0] == '-' {
32
		neg = s[0] == '-'
33
		s = s[1:]
34
	}
35
36
	exp := 0
37
	if e := strings.IndexAny(s, "eE"); e >= 0 {
38
		n, err := strconv.Atoi(s[e+1:])
39
		if err != nil {
40
			return Decimal{}, badDecimal(original)
41
		}
42
		// bounded so pow10 can't be asked to materialize an astronomically
43
		// large number (e.g. 1E1000000000000)
44
		const maxExp = 10000
45
		if n > maxExp || n < -maxExp {
46
			return Decimal{}, badDecimal(original)
47
		}
48
		exp = n
49
		s = s[:e]
50
	}
51
52
	intPart, fracPart := s, ""
53
	if before, after, ok := strings.Cut(s, "."); ok {
54
		if strings.IndexByte(after, '.') >= 0 {
55
			return Decimal{}, badDecimal(original)
56
		}
57
		intPart, fracPart = before, after
58
	}
59
60
	scale := len(fracPart) - exp
61
	if scale >= 0 {
62
		if v, ok := parseDigits(intPart + fracPart); ok {
63
			if neg {
64
				v = -v
65
			}
66
			return Decimal{v: v, scale: scale}.fastNormalized(), nil
67
		}
68
	}
69
70
	digits := intPart + fracPart
71
	if neg {
72
		digits = "-" + digits
73
	}
74
	coeff, ok := new(big.Int).SetString(digits, 10)
75
	if !ok {
76
		return Decimal{}, badDecimal(original)
77
	}
78
	if scale < 0 {
79
		coeff.Mul(coeff, pow10(-scale))
80
		scale = 0
81
	}
82
	return Decimal{big: coeff, scale: scale}.normalized(), nil
83
}
84
85
func badDecimal(s string) error { return fmt.Errorf("can't convert %s to decimal", s) }
86
87
func parseDigits(s string) (int64, bool) {
88
	var v int64
89
	for i := 0; i < len(s); i++ {
90
		c := s[i]
91
		if c < '0' || c > '9' {
92
			return 0, false
93
		}
94
		d := int64(c - '0')
95
		if v > (math.MaxInt64-d)/10 {
96
			return 0, false
97
		}
98
		v = v*10 + d
99
	}
100
	return v, len(s) > 0
101
}
102
103
func (d Decimal) fastNormalized() Decimal {
104
	v, scale := d.v, d.scale
105
	for v != 0 && scale > 0 && v%10 == 0 {
106
		v /= 10
107
		scale--
108
	}
109
	if v == 0 {
110
		return Decimal{}
111
	}
112
	return Decimal{v: v, scale: scale}
113
}
114
115
func (d Decimal) String() string {
116
	if d.big != nil {
117
		return bigString(d.big, d.scale)
118
	}
119
	if d.v == 0 {
120
		return "0"
121
	}
122
	digits := strconv.FormatInt(d.v, 10)
123
	sign := ""
124
	if digits[0] == '-' {
125
		sign, digits = "-", digits[1:]
126
	}
127
	if d.scale == 0 {
128
		return sign + digits
129
	}
130
	if len(digits) <= d.scale {
131
		digits = strings.Repeat("0", d.scale-len(digits)+1) + digits
132
	}
133
	split := len(digits) - d.scale
134
	return sign + digits[:split] + "." + digits[split:]
135
}
136
137
func bigString(coeff *big.Int, scale int) string {
138
	if coeff.Sign() == 0 {
139
		return "0"
140
	}
141
	abs := new(big.Int).Set(coeff)
142
	sign := ""
143
	if abs.Sign() < 0 {
144
		sign = "-"
145
		abs.Abs(abs)
146
	}
147
	digits := abs.String()
148
	if scale == 0 {
149
		return sign + digits
150
	}
151
	if len(digits) <= scale {
152
		digits = strings.Repeat("0", scale-len(digits)+1) + digits
153
	}
154
	split := len(digits) - scale
155
	return sign + digits[:split] + "." + digits[split:]
156
}
157
158
// StringFixed returns a string representation with exactly places digits
159
// after the decimal point. Pads with zeros or truncates as needed.
160
// decSep and thousandsSep control formatting; zero values mean no custom separator.
161
func (d Decimal) StringFixed(places int, decSep, thousandsSep byte) string {
162
	var sb strings.Builder
163
	sb.Grow(32)
164
	d.WriteFixed(&sb, places, decSep, thousandsSep)
165
	return sb.String()
166
}
167
168
// WriteFixed writes a string representation with exactly places digits
169
// after the decimal point directly into sb. Pads with zeros or truncates
170
// as needed. decSep and thousandsSep control formatting; zero values mean
171
// no custom separator.
172
func (d Decimal) WriteFixed(sb *strings.Builder, places int, decSep, thousandsSep byte) {
173
	if d.IsZero() {
174
		sb.WriteByte('0')
175
		if places > 0 {
176
			if decSep != 0 {
177
				sb.WriteByte(decSep)
178
			} else {
179
				sb.WriteByte('.')
180
			}
181
			for range places {
182
				sb.WriteByte('0')
183
			}
184
		}
185
		return
186
	}
187
188
	var digitBuf [128]byte
189
	var digits []byte
190
	if d.big != nil {
191
		digits = d.big.Append(digitBuf[:0], 10)
192
	} else {
193
		digits = strconv.AppendInt(digitBuf[:0], d.v, 10)
194
	}
195
196
	sign := false
197
	if len(digits) > 0 && digits[0] == '-' {
198
		sign = true
199
		digits = digits[1:]
200
	}
201
202
	intLen := len(digits) - d.scale
203
204
	if sign {
205
		sb.WriteByte('-')
206
	}
207
208
	// Integer part
209
	if intLen <= 0 {
210
		sb.WriteByte('0')
211
	} else {
212
		for i := range intLen {
213
			if thousandsSep != 0 && i > 0 && (intLen-i)%3 == 0 {
214
				sb.WriteByte(thousandsSep)
215
			}
216
			sb.WriteByte(digits[i])
217
		}
218
	}
219
220
	// Fractional part
221
	if places > 0 {
222
		if decSep != 0 {
223
			sb.WriteByte(decSep)
224
		} else {
225
			sb.WriteByte('.')
226
		}
227
228
		written := 0
229
		if intLen > 0 {
230
			n := min(d.scale, places)
231
			for i := range n {
232
				sb.WriteByte(digits[intLen+i])
233
				written++
234
			}
235
		} else {
236
			leadingZeros := -intLen
237
			for ; written < leadingZeros && written < places; written++ {
238
				sb.WriteByte('0')
239
			}
240
			for i := 0; i < len(digits) && written < places; i++ {
241
				sb.WriteByte(digits[i])
242
				written++
243
			}
244
		}
245
		for ; written < places; written++ {
246
			sb.WriteByte('0')
247
		}
248
	}
249
}
250
251
func (d Decimal) Abs() Decimal {
252
	if d.big != nil {
253
		if d.big.Sign() == 0 {
254
			return Decimal{}
255
		}
256
		if d.big.Sign() > 0 {
257
			return Decimal{big: new(big.Int).Set(d.big), scale: d.scale}
258
		}
259
		return Decimal{big: new(big.Int).Neg(d.big), scale: d.scale}
260
	}
261
	if d.v == 0 {
262
		return Decimal{}
263
	}
264
	if d.v < 0 {
265
		return Decimal{v: -d.v, scale: d.scale}
266
	}
267
	return Decimal{v: d.v, scale: d.scale}
268
}
269
270
func (d Decimal) Neg() Decimal {
271
	if d.big != nil {
272
		if d.big.Sign() == 0 {
273
			return Decimal{}
274
		}
275
		return Decimal{big: new(big.Int).Neg(d.big), scale: d.scale}
276
	}
277
	if d.v == 0 {
278
		return Decimal{}
279
	}
280
	return Decimal{v: -d.v, scale: d.scale}
281
}
282
283
func (d Decimal) Sub(other Decimal) Decimal { return d.Add(other.Neg()) }
284
func (d Decimal) Add(other Decimal) Decimal {
285
	if d.big == nil && other.big == nil {
286
		if r, ok := fastAdd(d, other); ok {
287
			return r
288
		}
289
	}
290
	a, b, scale := align(d, other)
291
	sum := new(big.Int).Add(a, b)
292
	return Decimal{big: sum, scale: scale}.normalized()
293
}
294
295
func fastAdd(a, b Decimal) (Decimal, bool) {
296
	scale := max(a.scale, b.scale)
297
	av, ok1 := scaleBy(a.v, scale-a.scale)
298
	bv, ok2 := scaleBy(b.v, scale-b.scale)
299
	if !ok1 || !ok2 {
300
		return Decimal{}, false
301
	}
302
	sum := av + bv
303
	if (av > 0 && bv > 0 && sum < 0) || (av < 0 && bv < 0 && sum >= 0) {
304
		return Decimal{}, false
305
	}
306
	return Decimal{v: sum, scale: scale}.fastNormalized(), true
307
}
308
309
func (d Decimal) Mul(other Decimal) Decimal {
310
	if d.IsZero() || other.IsZero() {
311
		return Decimal{}
312
	}
313
	product := new(big.Int).Mul(d.coeffBig(), other.coeffBig())
314
	return Decimal{big: product, scale: d.scale + other.scale}.normalized()
315
}
316
317
func (d Decimal) Div(other Decimal) Decimal {
318
	if other.IsZero() {
319
		panic("decimal: division by zero")
320
	}
321
	if d.IsZero() {
322
		return Decimal{}
323
	}
324
325
	scale := max(d.scale, other.scale) + 10
326
327
	dCoeff := d.coeffBig()
328
	oCoeff := other.coeffBig()
329
330
	shift := scale + other.scale - d.scale
331
	if shift > 0 {
332
		dCoeff = new(big.Int).Mul(dCoeff, pow10(shift))
333
	} else if shift < 0 {
334
		oCoeff = new(big.Int).Mul(oCoeff, pow10(-shift))
335
	}
336
337
	quo := new(big.Int).Quo(dCoeff, oCoeff)
338
	return Decimal{big: quo, scale: scale}.normalized()
339
}
340
341
func (d Decimal) Cmp(other Decimal) int {
342
	if d.big == nil && other.big == nil {
343
		if r, ok := fastCmp(d, other); ok {
344
			return r
345
		}
346
	}
347
	a, b, _ := align(d, other)
348
	return a.Cmp(b)
349
}
350
351
func fastCmp(a, b Decimal) (int, bool) {
352
	scale := max(a.scale, b.scale)
353
	av, ok1 := scaleBy(a.v, scale-a.scale)
354
	bv, ok2 := scaleBy(b.v, scale-b.scale)
355
	if !ok1 || !ok2 {
356
		return 0, false
357
	}
358
	switch {
359
	case av < bv:
360
		return -1, true
361
	case av > bv:
362
		return 1, true
363
	}
364
	return 0, true
365
}
366
367
func (d Decimal) Equal(other Decimal) bool {
368
	return d.Cmp(other) == 0
369
}
370
371
func (d Decimal) IsZero() bool {
372
	if d.big != nil {
373
		return d.big.Sign() == 0
374
	}
375
	return d.v == 0
376
}
377
378
func (d Decimal) coeffBig() *big.Int {
379
	if d.big != nil {
380
		return new(big.Int).Set(d.big)
381
	}
382
	return big.NewInt(d.v)
383
}
384
385
// normalized strips trailing zero digits from a big coefficient and demotes
386
// to the fast lane when the result fits in an int64.
387
func (d Decimal) normalized() Decimal {
388
	if d.big == nil || d.big.Sign() == 0 {
389
		return Decimal{}
390
	}
391
392
	sign := d.big.Sign()
393
	abs := new(big.Int).Abs(d.big)
394
	ten := big.NewInt(10)
395
	rem := new(big.Int)
396
	scale := d.scale
397
	for scale > 0 {
398
		quotient, _ := new(big.Int).QuoRem(abs, ten, rem)
399
		if rem.Sign() != 0 {
400
			break
401
		}
402
		abs = quotient
403
		scale--
404
	}
405
406
	if sign < 0 {
407
		abs.Neg(abs)
408
	}
409
	if abs.IsInt64() {
410
		return Decimal{v: abs.Int64(), scale: scale}
411
	}
412
	return Decimal{big: abs, scale: scale}
413
}
414
415
func align(a, b Decimal) (aCoeff *big.Int, bCoeff *big.Int, scale int) {
416
	scale = max(b.scale, a.scale)
417
418
	aCoeff = a.coeffBig()
419
	bCoeff = b.coeffBig()
420
	if delta := scale - a.scale; delta > 0 {
421
		aCoeff.Mul(aCoeff, pow10(delta))
422
	}
423
	if delta := scale - b.scale; delta > 0 {
424
		bCoeff.Mul(bCoeff, pow10(delta))
425
	}
426
	return aCoeff, bCoeff, scale
427
}
428
429
func pow10(n int) *big.Int {
430
	if n <= 0 {
431
		return big.NewInt(1)
432
	}
433
	return new(big.Int).Exp(big.NewInt(10), big.NewInt(int64(n)), nil)
434
}
435
436
// pow10i64 holds 10^n for n up to 18 (10^18 fits in an int64).
437
var pow10i64 = [...]int64{
438
	1, 10, 100, 1_000, 10_000, 100_000, 1_000_000, 10_000_000, 100_000_000,
439
	1_000_000_000, 10_000_000_000, 100_000_000_000, 1_000_000_000_000,
440
	10_000_000_000_000, 100_000_000_000_000, 1_000_000_000_000_000,
441
	10_000_000_000_000_000, 100_000_000_000_000_000, 1_000_000_000_000_000_000,
442
}
443
444
// scaleBy returns v·10^n; ok is false when the product overflows int64.
445
func scaleBy(v int64, n int) (int64, bool) {
446
	if n <= 0 {
447
		return v, true
448
	}
449
	if n >= len(pow10i64) {
450
		return 0, false
451
	}
452
	p := pow10i64[n]
453
	r := v * p
454
	return r, v == 0 || r/v == p
455
}