1 files changed,
186 insertions(+),
42 deletions(-)
Author:
Oleksandr Smirnov
olexsmir@gmail.com
Committed at:
2026-09-02 19:12:11 +0300
Authored at:
2026-09-02 18:47:08 +0300
Change ID:
qyvsqmovurkuyvpvnnoskkklrlryuxtq
Parent:
2e3f494
M
internal/decimal/decimal.go
··· 2 2 3 3 import ( 4 4 "fmt" 5 + "math" 5 6 "math/big" 6 7 "strconv" 7 8 "strings" 8 9 ) 9 10 10 11 type Decimal struct { 12 + v int64 11 13 scale int 12 - coeff *big.Int 14 + big *big.Int 13 15 } 14 16 15 17 func FromInt(v int64) Decimal { 16 18 if v == 0 { 17 19 return Decimal{} 18 20 } 19 - return Decimal{coeff: big.NewInt(v)} 21 + return Decimal{v: v} 20 22 } 21 23 22 24 func FromString(s string) (Decimal, error) { ··· 55 57 intPart, fracPart = before, after 56 58 } 57 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 + 58 70 digits := intPart + fracPart 59 71 if neg { 60 72 digits = "-" + digits ··· 63 75 if !ok { 64 76 return Decimal{}, badDecimal(original) 65 77 } 66 - 67 - scale := len(fracPart) - exp 68 78 if scale < 0 { 69 79 coeff.Mul(coeff, pow10(-scale)) 70 80 scale = 0 71 81 } 72 - return Decimal{coeff: coeff, scale: scale}.normalized(), nil 82 + return Decimal{big: coeff, scale: scale}.normalized(), nil 73 83 } 74 84 75 85 func badDecimal(s string) error { return fmt.Errorf("can't convert %s to decimal", s) } 76 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 + 77 115 func (d Decimal) String() string { 78 - if d.coeff == nil || d.coeff.Sign() == 0 { 116 + if d.big != nil { 117 + return bigString(d.big, d.scale) 118 + } 119 + if d.v == 0 { 79 120 return "0" 80 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 +} 81 136 82 - abs := new(big.Int).Set(d.coeff) 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) 83 142 sign := "" 84 143 if abs.Sign() < 0 { 85 144 sign = "-" 86 145 abs.Abs(abs) 87 146 } 88 - 89 147 digits := abs.String() 90 - if d.scale == 0 { 148 + if scale == 0 { 91 149 return sign + digits 92 150 } 93 - 94 - if len(digits) <= d.scale { 95 - digits = strings.Repeat("0", d.scale-len(digits)+1) + digits 151 + if len(digits) <= scale { 152 + digits = strings.Repeat("0", scale-len(digits)+1) + digits 96 153 } 97 - split := len(digits) - d.scale 154 + split := len(digits) - scale 98 155 return sign + digits[:split] + "." + digits[split:] 99 156 } 100 157 ··· 129 186 } 130 187 131 188 var digitBuf [128]byte 132 - digits := d.coeff.Append(digitBuf[:0], 10) 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 + } 133 195 134 196 sign := false 135 197 if len(digits) > 0 && digits[0] == '-' { ··· 187 249 } 188 250 189 251 func (d Decimal) Abs() Decimal { 190 - if d.coeff == nil || d.coeff.Sign() == 0 { 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 { 191 262 return Decimal{} 192 263 } 193 - if d.coeff.Sign() > 0 { 194 - return Decimal{coeff: new(big.Int).Set(d.coeff), scale: d.scale} 264 + if d.v < 0 { 265 + return Decimal{v: -d.v, scale: d.scale} 195 266 } 196 - return Decimal{coeff: new(big.Int).Neg(d.coeff), scale: d.scale} 267 + return Decimal{v: d.v, scale: d.scale} 197 268 } 198 269 199 270 func (d Decimal) Neg() Decimal { 200 - if d.coeff == nil || d.coeff.Sign() == 0 { 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 { 201 278 return Decimal{} 202 279 } 203 - return Decimal{coeff: new(big.Int).Neg(d.coeff), scale: d.scale} 280 + return Decimal{v: -d.v, scale: d.scale} 204 281 } 205 282 206 283 func (d Decimal) Sub(other Decimal) Decimal { return d.Add(other.Neg()) } 207 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 + } 208 290 a, b, scale := align(d, other) 209 291 sum := new(big.Int).Add(a, b) 210 - return Decimal{coeff: sum, scale: scale}.normalized() 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 211 307 } 212 308 213 309 func (d Decimal) Mul(other Decimal) Decimal { 214 310 if d.IsZero() || other.IsZero() { 215 311 return Decimal{} 216 312 } 217 - product := new(big.Int).Mul(d.coeffOrZero(), other.coeffOrZero()) 218 - return Decimal{coeff: product, scale: d.scale + other.scale}.normalized() 313 + product := new(big.Int).Mul(d.coeffBig(), other.coeffBig()) 314 + return Decimal{big: product, scale: d.scale + other.scale}.normalized() 219 315 } 220 316 221 317 func (d Decimal) Div(other Decimal) Decimal { ··· 228 324 229 325 scale := max(d.scale, other.scale) + 10 230 326 231 - dCoeff := d.coeffOrZero() 232 - oCoeff := other.coeffOrZero() 327 + dCoeff := d.coeffBig() 328 + oCoeff := other.coeffBig() 233 329 234 330 shift := scale + other.scale - d.scale 235 331 if shift > 0 { ··· 239 335 } 240 336 241 337 quo := new(big.Int).Quo(dCoeff, oCoeff) 242 - return Decimal{coeff: quo, scale: scale}.normalized() 338 + return Decimal{big: quo, scale: scale}.normalized() 243 339 } 244 340 245 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 + } 246 347 a, b, _ := align(d, other) 247 348 return a.Cmp(b) 248 349 } 249 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 + 250 367 func (d Decimal) Equal(other Decimal) bool { 251 368 return d.Cmp(other) == 0 252 369 } 253 370 254 371 func (d Decimal) IsZero() bool { 255 - return d.coeff == nil || d.coeff.Sign() == 0 372 + if d.big != nil { 373 + return d.big.Sign() == 0 374 + } 375 + return d.v == 0 256 376 } 257 377 258 -func (d Decimal) coeffOrZero() *big.Int { 259 - if d.coeff == nil { 260 - return new(big.Int) 378 +func (d Decimal) coeffBig() *big.Int { 379 + if d.big != nil { 380 + return new(big.Int).Set(d.big) 261 381 } 262 - return new(big.Int).Set(d.coeff) 382 + return big.NewInt(d.v) 263 383 } 264 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. 265 387 func (d Decimal) normalized() Decimal { 266 - if d.coeff == nil || d.coeff.Sign() == 0 { 388 + if d.big == nil || d.big.Sign() == 0 { 267 389 return Decimal{} 268 390 } 269 - if d.scale == 0 { 270 - return Decimal{coeff: new(big.Int).Set(d.coeff)} 271 - } 272 391 273 - sign := d.coeff.Sign() 274 - abs := new(big.Int).Abs(d.coeff) 392 + sign := d.big.Sign() 393 + abs := new(big.Int).Abs(d.big) 275 394 ten := big.NewInt(10) 276 395 rem := new(big.Int) 277 - for d.scale > 0 { 396 + scale := d.scale 397 + for scale > 0 { 278 398 quotient, _ := new(big.Int).QuoRem(abs, ten, rem) 279 399 if rem.Sign() != 0 { 280 400 break 281 401 } 282 402 abs = quotient 283 - d.scale-- 403 + scale-- 284 404 } 285 405 286 406 if sign < 0 { 287 407 abs.Neg(abs) 288 408 } 289 - return Decimal{coeff: abs, scale: d.scale} 409 + if abs.IsInt64() { 410 + return Decimal{v: abs.Int64(), scale: scale} 411 + } 412 + return Decimal{big: abs, scale: scale} 290 413 } 291 414 292 415 func align(a, b Decimal) (aCoeff *big.Int, bCoeff *big.Int, scale int) { 293 416 scale = max(b.scale, a.scale) 294 417 295 - aCoeff = a.coeffOrZero() 296 - bCoeff = b.coeffOrZero() 418 + aCoeff = a.coeffBig() 419 + bCoeff = b.coeffBig() 297 420 if delta := scale - a.scale; delta > 0 { 298 421 aCoeff.Mul(aCoeff, pow10(delta)) 299 422 } ··· 309 432 } 310 433 return new(big.Int).Exp(big.NewInt(10), big.NewInt(int64(n)), nil) 311 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 +}