diff --git a/README.md b/README.md index 9742a72..a87a6ee 100644 --- a/README.md +++ b/README.md @@ -4,9 +4,12 @@ A comprehensive mathematical utilities library for the Lux ecosystem. ## Features -- **Math utilities**: Safe math operations, averagers, and meters +- **Big integer utilities**: `HexOrDecimal256`, `ParseBig256`, `U256`, `BigPow` etc. +- **Safe arithmetic**: Overflow-checked `SafeAdd`, `SafeSub`, `SafeMul` +- **Bit operations**: XOR, AND, compression utilities - **Set operations**: Efficient set implementations including bit sets -- **Numerical computing**: Integration with gonum for advanced mathematical operations +- **Data structures**: Linked lists, hash maps, heaps +- **Averagers**: Time-windowed averaging utilities ## Installation @@ -14,15 +17,60 @@ A comprehensive mathematical utilities library for the Lux ecosystem. go get github.com/luxfi/math ``` +## Package Structure + +| Package | Description | +|---------|-------------| +| `github.com/luxfi/math` | Root package with re-exports for backwards compatibility | +| `github.com/luxfi/math/big` | Big integer utilities (HexOrDecimal256, U256, parsing) | +| `github.com/luxfi/math/safe` | Overflow-safe arithmetic operations | +| `github.com/luxfi/math/bit` | Bit manipulation utilities | +| `github.com/luxfi/math/set` | Set data structures | +| `github.com/luxfi/math/linked` | Linked data structures | +| `github.com/luxfi/math/heap` | Heap implementations | + ## Usage ```go +// Import root package (re-exports from subpackages) +import "github.com/luxfi/math" + +// Or import specific subpackages directly import ( - "github.com/luxfi/math/math" + "github.com/luxfi/math/big" + "github.com/luxfi/math/safe" "github.com/luxfi/math/set" ) ``` +### Big Integer Operations + +```go +import "github.com/luxfi/math/big" + +// Parse hex or decimal +val, _ := big.ParseBig256("0x1234") + +// 256-bit unsigned operations +result := big.U256(someInt) + +// Power operation +power := big.BigPow(2, 256) +``` + +### Safe Arithmetic + +```go +import "github.com/luxfi/math/safe" + +// Returns (result, overflow bool) +sum, overflow := safe.SafeAdd(a, b) +product, overflow := safe.SafeMul(x, y) + +// Returns (result, error) +sum, err := safe.Add64(a, b) +``` + ## License See the LICENSE file for licensing terms. \ No newline at end of file diff --git a/math/averager.go b/averager.go similarity index 100% rename from math/averager.go rename to averager.go diff --git a/math/averager_heap.go b/averager_heap.go similarity index 100% rename from math/averager_heap.go rename to averager_heap.go diff --git a/math/averager_heap_test.go b/averager_heap_test.go similarity index 100% rename from math/averager_heap_test.go rename to averager_heap_test.go diff --git a/big.go b/big.go new file mode 100644 index 0000000..97b2fd7 --- /dev/null +++ b/big.go @@ -0,0 +1,39 @@ +// Copyright (C) 2019-2025, Lux Industries Inc. All rights reserved. +// See the file LICENSE for licensing terms. + +// Package math re-exports big.Int utilities for backwards compatibility. +// New code should import github.com/luxfi/math/big directly. +package math + +import ( + "math/big" + + luxbig "github.com/luxfi/math/big" +) + +// Type aliases for backwards compatibility. +type ( + HexOrDecimal256 = luxbig.HexOrDecimal256 + Decimal256 = luxbig.Decimal256 + HexOrDecimal64 = luxbig.HexOrDecimal64 +) + +// Function aliases for backwards compatibility. +var ( + NewHexOrDecimal256 = luxbig.NewHexOrDecimal256 + NewDecimal256 = luxbig.NewDecimal256 + ParseBig256 = luxbig.ParseBig256 + MustParseBig256 = luxbig.MustParseBig256 + ParseUint64 = luxbig.ParseUint64 + MustParseUint64 = luxbig.MustParseUint64 + BigPow = luxbig.BigPow + BigMax = luxbig.BigMax + BigMin = luxbig.BigMin + PaddedBigBytes = luxbig.PaddedBigBytes + ReadBits = luxbig.ReadBits + U256 = luxbig.U256 + U256Bytes = luxbig.U256Bytes +) + +// MaxBig256 is the maximum value for a 256-bit unsigned integer. +var MaxBig256 = new(big.Int).Set(luxbig.MaxBig256) diff --git a/math/big.go b/big/big.go similarity index 74% rename from math/big.go rename to big/big.go index 27032f9..52ea6e0 100644 --- a/math/big.go +++ b/big/big.go @@ -1,11 +1,13 @@ // Copyright (C) 2019-2025, Lux Industries Inc. All rights reserved. // See the file LICENSE for licensing terms. -package math +// Package big provides big.Int utilities and parsing functions. +package big import ( "fmt" "math/big" + "strconv" ) // Various big integer limit values. @@ -16,9 +18,7 @@ var ( ) const ( - // number of bits in a big.Word - wordBits = 32 << (uint64(^big.Word(0)) >> 63) - // number of bytes in a big.Word + wordBits = 32 << (uint64(^big.Word(0)) >> 63) wordBytes = wordBits / 8 ) @@ -32,7 +32,6 @@ func NewHexOrDecimal256(x int64) *HexOrDecimal256 { return &h } -// UnmarshalJSON implements json.Unmarshaler. func (i *HexOrDecimal256) UnmarshalJSON(input []byte) error { if len(input) > 1 && input[0] == '"' { input = input[1 : len(input)-1] @@ -40,7 +39,6 @@ func (i *HexOrDecimal256) UnmarshalJSON(input []byte) error { return i.UnmarshalText(input) } -// UnmarshalText implements encoding.TextUnmarshaler. func (i *HexOrDecimal256) UnmarshalText(input []byte) error { bigint, ok := ParseBig256(string(input)) if !ok { @@ -50,7 +48,6 @@ func (i *HexOrDecimal256) UnmarshalText(input []byte) error { return nil } -// MarshalText implements encoding.TextMarshaler. func (i *HexOrDecimal256) MarshalText() ([]byte, error) { if i == nil { return []byte("0x0"), nil @@ -61,14 +58,12 @@ func (i *HexOrDecimal256) MarshalText() ([]byte, error) { // Decimal256 unmarshals big.Int as a decimal string. type Decimal256 big.Int -// NewDecimal256 creates a new Decimal256 func NewDecimal256(x int64) *Decimal256 { b := big.NewInt(x) d := Decimal256(*b) return &d } -// UnmarshalText implements encoding.TextUnmarshaler. func (i *Decimal256) UnmarshalText(input []byte) error { bigint, ok := ParseBig256(string(input)) if !ok { @@ -78,12 +73,10 @@ func (i *Decimal256) UnmarshalText(input []byte) error { return nil } -// MarshalText implements encoding.TextMarshaler. func (i *Decimal256) MarshalText() ([]byte, error) { return []byte(i.String()), nil } -// String implements Stringer. func (i *Decimal256) String() string { if i == nil { return "0" @@ -91,6 +84,29 @@ func (i *Decimal256) String() string { return fmt.Sprintf("%#d", (*big.Int)(i)) } +// HexOrDecimal64 marshals uint64 as hex or decimal. +type HexOrDecimal64 uint64 + +func (i *HexOrDecimal64) UnmarshalJSON(input []byte) error { + if len(input) > 1 && input[0] == '"' { + input = input[1 : len(input)-1] + } + return i.UnmarshalText(input) +} + +func (i *HexOrDecimal64) UnmarshalText(input []byte) error { + n, ok := ParseUint64(string(input)) + if !ok { + return fmt.Errorf("invalid hex or decimal integer %q", input) + } + *i = HexOrDecimal64(n) + return nil +} + +func (i HexOrDecimal64) MarshalText() ([]byte, error) { + return fmt.Appendf(nil, "%#x", uint64(i)), nil +} + // ParseBig256 parses s as a 256 bit integer in decimal or hexadecimal syntax. func ParseBig256(s string) (*big.Int, bool) { if s == "" { @@ -118,6 +134,28 @@ func MustParseBig256(s string) *big.Int { return v } +// ParseUint64 parses s as an integer in decimal or hexadecimal syntax. +func ParseUint64(s string) (uint64, bool) { + if s == "" { + return 0, true + } + if len(s) >= 2 && (s[:2] == "0x" || s[:2] == "0X") { + v, err := strconv.ParseUint(s[2:], 16, 64) + return v, err == nil + } + v, err := strconv.ParseUint(s, 10, 64) + return v, err == nil +} + +// MustParseUint64 parses s as an integer and panics if invalid. +func MustParseUint64(s string) uint64 { + v, ok := ParseUint64(s) + if !ok { + panic("invalid unsigned 64 bit integer: " + s) + } + return v +} + // BigPow returns a ** b as a big integer. func BigPow(a, b int64) *big.Int { r := big.NewInt(a) diff --git a/big/go.mod b/big/go.mod new file mode 100644 index 0000000..468123a --- /dev/null +++ b/big/go.mod @@ -0,0 +1,3 @@ +module github.com/luxfi/math/big + +go 1.24.0 diff --git a/math/continuous_averager.go b/continuous_averager.go similarity index 100% rename from math/continuous_averager.go rename to continuous_averager.go diff --git a/math/continuous_averager_benchmark_test.go b/continuous_averager_benchmark_test.go similarity index 100% rename from math/continuous_averager_benchmark_test.go rename to continuous_averager_benchmark_test.go diff --git a/math/continuous_averager_test.go b/continuous_averager_test.go similarity index 100% rename from math/continuous_averager_test.go rename to continuous_averager_test.go diff --git a/go.mod b/go.mod index a9bed46..a3b69d5 100644 --- a/go.mod +++ b/go.mod @@ -2,4 +2,10 @@ module github.com/luxfi/math go 1.25.5 -require github.com/luxfi/sampler v1.0.0 +require ( + github.com/luxfi/math/big v1.0.0 + github.com/luxfi/sampler v1.0.0 + golang.org/x/exp v0.0.0-20241217172543-b2144cdd0a67 +) + +replace github.com/luxfi/math/big => ./big diff --git a/heap/map.go b/heap/map.go index 43281e1..527d108 100644 --- a/heap/map.go +++ b/heap/map.go @@ -3,11 +3,13 @@ package heap -import ( - "container/heap" +import "container/heap" - luxmath "github.com/luxfi/math" -) +// zero returns the zero value of any type T. +func zero[T any]() T { + var z T + return z +} var _ heap.Interface = (*indexedQueue[int, int])(nil) @@ -47,12 +49,12 @@ func (m *Map[K, V]) Push(k K, v V) (V, bool) { } heap.Push(m.queue, entry[K, V]{k: k, v: v}) - return luxmath.Zero[V](), false + return zero[V](), false } func (m *Map[K, V]) Pop() (K, V, bool) { if m.Len() == 0 { - return luxmath.Zero[K](), luxmath.Zero[V](), false + return zero[K](), zero[V](), false } popped := heap.Pop(m.queue).(entry[K, V]) @@ -61,7 +63,7 @@ func (m *Map[K, V]) Pop() (K, V, bool) { func (m *Map[K, V]) Peek() (K, V, bool) { if m.Len() == 0 { - return luxmath.Zero[K](), luxmath.Zero[V](), false + return zero[K](), zero[V](), false } entry := m.queue.entries[0] @@ -77,7 +79,7 @@ func (m *Map[K, V]) Remove(k K) (V, bool) { removed := heap.Remove(m.queue, i).(entry[K, V]) return removed.v, true } - return luxmath.Zero[V](), false + return zero[V](), false } func (m *Map[K, V]) Contains(k K) bool { @@ -90,7 +92,7 @@ func (m *Map[K, V]) Get(k K) (V, bool) { got := m.queue.entries[i] return got.v, true } - return luxmath.Zero[V](), false + return zero[V](), false } func (m *Map[K, V]) Fix(k K) { diff --git a/heap/queue.go b/heap/queue.go index 0a54a3f..24656e8 100644 --- a/heap/queue.go +++ b/heap/queue.go @@ -3,11 +3,7 @@ package heap -import ( - "container/heap" - - luxmath "github.com/luxfi/math" -) +import "container/heap" var _ heap.Interface = (*queue[int])(nil) @@ -44,7 +40,7 @@ func (q *Queue[T]) Push(t T) { func (q *Queue[T]) Pop() (T, bool) { if q.Len() == 0 { - return luxmath.Zero[T](), false + return zero[T](), false } return heap.Pop(q.queue).(T), true @@ -52,7 +48,7 @@ func (q *Queue[T]) Pop() (T, bool) { func (q *Queue[T]) Peek() (T, bool) { if q.Len() == 0 { - return luxmath.Zero[T](), false + return zero[T](), false } return q.queue.entries[0], true @@ -87,7 +83,7 @@ func (q *queue[T]) Pop() any { end := len(q.entries) - 1 popped := q.entries[end] - q.entries[end] = luxmath.Zero[T]() + q.entries[end] = zero[T]() q.entries = q.entries[:end] return popped diff --git a/math/integer.go b/math/integer.go deleted file mode 100644 index d0bedec..0000000 --- a/math/integer.go +++ /dev/null @@ -1,76 +0,0 @@ -// Copyright (C) 2019-2025, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package math - -import ( - "fmt" - "math/bits" - "strconv" -) - -// HexOrDecimal64 marshals uint64 as hex or decimal. -type HexOrDecimal64 uint64 - -// UnmarshalJSON implements json.Unmarshaler. -func (i *HexOrDecimal64) UnmarshalJSON(input []byte) error { - if len(input) > 1 && input[0] == '"' { - input = input[1 : len(input)-1] - } - return i.UnmarshalText(input) -} - -// UnmarshalText implements encoding.TextUnmarshaler. -func (i *HexOrDecimal64) UnmarshalText(input []byte) error { - n, ok := ParseUint64(string(input)) - if !ok { - return fmt.Errorf("invalid hex or decimal integer %q", input) - } - *i = HexOrDecimal64(n) - return nil -} - -// MarshalText implements encoding.TextMarshaler. -func (i HexOrDecimal64) MarshalText() ([]byte, error) { - return fmt.Appendf(nil, "%#x", uint64(i)), nil -} - -// ParseUint64 parses s as an integer in decimal or hexadecimal syntax. -func ParseUint64(s string) (uint64, bool) { - if s == "" { - return 0, true - } - if len(s) >= 2 && (s[:2] == "0x" || s[:2] == "0X") { - v, err := strconv.ParseUint(s[2:], 16, 64) - return v, err == nil - } - v, err := strconv.ParseUint(s, 10, 64) - return v, err == nil -} - -// MustParseUint64 parses s as an integer and panics if invalid. -func MustParseUint64(s string) uint64 { - v, ok := ParseUint64(s) - if !ok { - panic("invalid unsigned 64 bit integer: " + s) - } - return v -} - -// SafeSub returns x-y and checks for overflow. -func SafeSub(x, y uint64) (uint64, bool) { - diff, borrowOut := bits.Sub64(x, y, 0) - return diff, borrowOut != 0 -} - -// SafeAdd returns x+y and checks for overflow. -func SafeAdd(x, y uint64) (uint64, bool) { - sum, carryOut := bits.Add64(x, y, 0) - return sum, carryOut != 0 -} - -// SafeMul returns x*y and checks for overflow. -func SafeMul(x, y uint64) (uint64, bool) { - hi, lo := bits.Mul64(x, y) - return lo, hi != 0 -} diff --git a/math/meter/continuous_meter.go b/math/meter/continuous_meter.go deleted file mode 100644 index 34a3958..0000000 --- a/math/meter/continuous_meter.go +++ /dev/null @@ -1,79 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package meter - -import ( - "math" - "time" -) - -var ( - convertEToBase2 = math.Log(2) - - _ Factory = (*ContinuousFactory)(nil) - _ Meter = (*continuousMeter)(nil) -) - -// ContinuousFactory implements the Factory interface by returning a continuous -// time meter. -type ContinuousFactory struct{} - -func (ContinuousFactory) New(halflife time.Duration) Meter { - return NewMeter(halflife) -} - -type continuousMeter struct { - halflife float64 - value float64 - - numCoresRunning float64 - lastUpdated time.Time -} - -// NewMeter returns a new Meter with the provided halflife -func NewMeter(halflife time.Duration) Meter { - return &continuousMeter{ - halflife: float64(halflife) / convertEToBase2, - } -} - -func (a *continuousMeter) Inc(now time.Time, numCores float64) { - a.Read(now) - a.numCoresRunning += numCores -} - -func (a *continuousMeter) Dec(now time.Time, numCores float64) { - a.Read(now) - a.numCoresRunning -= numCores -} - -func (a *continuousMeter) Read(now time.Time) float64 { - timeSincePreviousUpdate := a.lastUpdated.Sub(now) - if timeSincePreviousUpdate >= 0 { - return a.value - } - a.lastUpdated = now - - factor := math.Exp(float64(timeSincePreviousUpdate) / a.halflife) - a.value *= factor - a.value += a.numCoresRunning * (1 - factor) - return a.value -} - -func (a *continuousMeter) TimeUntil(now time.Time, value float64) time.Duration { - currentValue := a.Read(now) - if currentValue <= value { - return time.Duration(0) - } - // Note that [factor] >= 1 - factor := currentValue / value - // Note that [numHalfLives] >= 0 - numHalflives := math.Log(factor) - duration := numHalflives * a.halflife - // Overflow protection - if duration > math.MaxInt64 { - return time.Duration(math.MaxInt64) - } - return time.Duration(duration) -} diff --git a/math/meter/factory.go b/math/meter/factory.go deleted file mode 100644 index c29c6e9..0000000 --- a/math/meter/factory.go +++ /dev/null @@ -1,12 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package meter - -import "time" - -// Factory returns new meters. -type Factory interface { - // New returns a new meter with the provided halflife. - New(halflife time.Duration) Meter -} diff --git a/math/meter/meter.go b/math/meter/meter.go deleted file mode 100644 index 360890b..0000000 --- a/math/meter/meter.go +++ /dev/null @@ -1,28 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package meter - -import "time" - -// Meter tracks a continuous exponential moving average of the % of time this -// meter has been running. -type Meter interface { - // Inc the meter, the read value will be monotonically increasing while - // the meter is running. - Inc(time.Time, float64) - - // Dec the meter, the read value will be exponentially decreasing while the - // meter is off. - Dec(time.Time, float64) - - // Read the current value of the meter, this can be used to approximate the - // percent of time the meter has been running recently. The definition of - // recently depends on the halflife of the decay function. - Read(time.Time) float64 - - // Returns the duration between [now] and when the value of this meter - // reaches [value], assuming that the number of cores running is always 0. - // If the value of this meter is already <= [value], returns the zero duration. - TimeUntil(now time.Time, value float64) time.Duration -} diff --git a/math/meter/meter_benchmark_test.go b/math/meter/meter_benchmark_test.go deleted file mode 100644 index 55f41b7..0000000 --- a/math/meter/meter_benchmark_test.go +++ /dev/null @@ -1,39 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package meter - -import ( - "fmt" - "testing" - "time" -) - -func BenchmarkMeters(b *testing.B) { - for _, meterDef := range meters { - period := time.Second + 500*time.Millisecond - name := fmt.Sprintf("%s-%s", meterDef.name, period) - b.Run(name, func(b *testing.B) { - m := meterDef.factory.New(halflife) - MeterBenchmark(b, m, period) - }) - - period = time.Millisecond - name = fmt.Sprintf("%s-%s", meterDef.name, period) - b.Run(name, func(b *testing.B) { - m := meterDef.factory.New(halflife) - MeterBenchmark(b, m, period) - }) - } -} - -func MeterBenchmark(b *testing.B, m Meter, period time.Duration) { - currentTime := time.Now() - m.Inc(currentTime, 1) - - b.ResetTimer() - for i := 0; i < b.N; i++ { - currentTime = currentTime.Add(period) - m.Read(currentTime) - } -} diff --git a/math/meter/meter_test.go b/math/meter/meter_test.go deleted file mode 100644 index 1790a29..0000000 --- a/math/meter/meter_test.go +++ /dev/null @@ -1,166 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package meter - -import ( - "fmt" - "testing" - "time" - - "github.com/stretchr/testify/require" -) - -var ( - halflife = time.Second - meters = []struct { - name string - factory Factory - }{ - { - name: "continuous", - factory: ContinuousFactory{}, - }, - } - - meterTests = []struct { - name string - test func(*testing.T, Factory) - }{ - { - name: "new", - test: NewTest, - }, - { - name: "standard usage", - test: StandardUsageTest, - }, - { - name: "time travel", - test: TimeTravelTest, - }, - } -) - -func TestMeters(t *testing.T) { - for _, s := range meters { - for _, test := range meterTests { - t.Run(fmt.Sprintf("meter %s test %s", s.name, test.name), func(t *testing.T) { - test.test(t, s.factory) - }) - } - } -} - -func NewTest(t *testing.T, factory Factory) { - require.NotNil(t, factory.New(halflife)) -} - -func TimeTravelTest(t *testing.T, factory Factory) { - require := require.New(t) - - m := factory.New(halflife) - - now := time.Date(1, 2, 3, 4, 5, 6, 7, time.UTC) - m.Inc(now, 1) - - now = now.Add(halflife - 1) - delta := 0.0001 - require.InDelta(.5, m.Read(now), delta) - - m.Dec(now, 1) - - now = now.Add(-halflife) - require.InDelta(.5, m.Read(now), delta) - - m.Inc(now, 1) - - now = now.Add(halflife / 2) - require.InDelta(.5, m.Read(now), delta) -} - -func StandardUsageTest(t *testing.T, factory Factory) { - require := require.New(t) - - m := factory.New(halflife) - - now := time.Date(1, 2, 3, 4, 5, 6, 7, time.UTC) - m.Inc(now, 1) - - now = now.Add(halflife - 1) - delta := 0.0001 - require.InDelta(.5, m.Read(now), delta) - - m.Inc(now, 1) - require.InDelta(.5, m.Read(now), delta) - - m.Dec(now, 1) - require.InDelta(.5, m.Read(now), delta) - - m.Dec(now, 1) - - require.InDelta(.5, m.Read(now), delta) - - now = now.Add(halflife) - require.InDelta(.25, m.Read(now), delta) - - m.Inc(now, 1) - - now = now.Add(halflife) - require.InDelta(.625, m.Read(now), delta) - - now = now.Add(34 * halflife) - require.InDelta(1, m.Read(now), delta) - - m.Dec(now, 1) - - now = now.Add(34 * halflife) - require.InDelta(0, m.Read(now), delta) - - m.Inc(now, 1) - - now = now.Add(2 * halflife) - require.InDelta(.75, m.Read(now), delta) - - // Second start - m.Inc(now, 1) - - now = now.Add(34 * halflife) - require.InDelta(2, m.Read(now), delta) - - // Stop the second CPU - m.Dec(now, 1) - - now = now.Add(34 * halflife) - require.InDelta(1, m.Read(now), delta) -} - -func TestTimeUntil(t *testing.T) { - require := require.New(t) - - halflife := 5 * time.Second - f := ContinuousFactory{} - m := f.New(halflife) - now := time.Now() - // Start the meter - m.Inc(now, 1) - // One halflife passes; stop the meter - now = now.Add(halflife) - m.Dec(now, 1) - // Read the current value - currentVal := m.Read(now) - // Suppose we want to wait for the value to be - // a third of its current value - desiredVal := currentVal / 3 - // See when that should happen - timeUntilDesiredVal := m.TimeUntil(now, desiredVal) - // Get the actual value at that time - now = now.Add(timeUntilDesiredVal) - actualVal := m.Read(now) - // Make sure the actual/expected are close - require.InDelta(desiredVal, actualVal, .00001) - // Make sure TimeUntil returns the zero duration if - // the value provided >= the current value - require.Zero(m.TimeUntil(now, actualVal)) - require.Zero(m.TimeUntil(now, actualVal+.1)) -} diff --git a/math/safe_math.go b/math/safe_math.go deleted file mode 100644 index 51e2104..0000000 --- a/math/safe_math.go +++ /dev/null @@ -1,76 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package math - -import ( - "errors" - "math" - - "golang.org/x/exp/constraints" - - luxmath "github.com/luxfi/math" -) - -var ( - ErrOverflow = errors.New("overflow") - ErrUnderflow = errors.New("underflow") -) - -// Add64 returns: -// 1) a + b -// 2) If there is overflow, an error -// -// Note that we don't have a generic Add function because checking for -// an overflow requires knowing the max size of a given type, which we -// don't know if we're adding generic types. -func Add64(a, b uint64) (uint64, error) { - if a > math.MaxUint64-b { - return 0, ErrOverflow - } - return a + b, nil -} - -// Sub returns: -// 1) a - b -// 2) If there is underflow, an error -func Sub[T constraints.Unsigned](a, b T) (T, error) { - if a < b { - return luxmath.Zero[T](), ErrUnderflow - } - return a - b, nil -} - -// Mul64 returns: -// 1) a * b -// 2) If there is overflow, an error -// -// Note that we don't have a generic Mul function because checking for -// an overflow requires knowing the max size of a given type, which we -// don't know if we're adding generic types. -func Mul64(a, b uint64) (uint64, error) { - if b != 0 && a > math.MaxUint64/b { - return 0, ErrOverflow - } - return a * b, nil -} - -func AbsDiff[T constraints.Unsigned](a, b T) T { - return max(a, b) - min(a, b) -} - -// Min returns the minimum of two uint64 values -func Min(a, b uint64) uint64 { - if a < b { - return a - } - return b -} - -// Max returns the maximum of two uint64 values -func Max(a, b uint64) uint64 { - if a > b { - return a - } - return b -} diff --git a/math/safe_math_test.go b/math/safe_math_test.go deleted file mode 100644 index 6344122..0000000 --- a/math/safe_math_test.go +++ /dev/null @@ -1,106 +0,0 @@ -// Copyright (C) 2019-2024, Lux Industries Inc. All rights reserved. -// See the file LICENSE for licensing terms. - -package math - -import ( - "math" - "testing" - - "github.com/stretchr/testify/require" -) - -const maxUint64 uint64 = math.MaxUint64 - -func TestAdd64(t *testing.T) { - require := require.New(t) - - sum, err := Add64(0, maxUint64) - require.NoError(err) - require.Equal(maxUint64, sum) - - sum, err = Add64(maxUint64, 0) - require.NoError(err) - require.Equal(maxUint64, sum) - - sum, err = Add64(uint64(1<<62), uint64(1<<62)) - require.NoError(err) - require.Equal(uint64(1<<63), sum) - - _, err = Add64(1, maxUint64) - require.ErrorIs(err, ErrOverflow) - - _, err = Add64(maxUint64, 1) - require.ErrorIs(err, ErrOverflow) - - _, err = Add64(maxUint64, maxUint64) - require.ErrorIs(err, ErrOverflow) -} - -func TestSub(t *testing.T) { - require := require.New(t) - - got, err := Sub(uint64(2), uint64(1)) - require.NoError(err) - require.Equal(uint64(1), got) - - got, err = Sub(uint64(2), uint64(2)) - require.NoError(err) - require.Zero(got) - - got, err = Sub(maxUint64, maxUint64) - require.NoError(err) - require.Zero(got) - - got, err = Sub(uint64(3), uint64(2)) - require.NoError(err) - require.Equal(uint64(1), got) - - _, err = Sub(uint64(1), uint64(2)) - require.ErrorIs(err, ErrUnderflow) - - _, err = Sub(maxUint64-1, maxUint64) - require.ErrorIs(err, ErrUnderflow) -} - -func TestMul64(t *testing.T) { - require := require.New(t) - - got, err := Mul64(0, maxUint64) - require.NoError(err) - require.Zero(got) - - got, err = Mul64(maxUint64, 0) - require.NoError(err) - require.Zero(got) - - got, err = Mul64(uint64(1), uint64(3)) - require.NoError(err) - require.Equal(uint64(3), got) - - got, err = Mul64(uint64(3), uint64(1)) - require.NoError(err) - require.Equal(uint64(3), got) - - got, err = Mul64(uint64(2), uint64(3)) - require.NoError(err) - require.Equal(uint64(6), got) - - got, err = Mul64(maxUint64, 0) - require.NoError(err) - require.Zero(got) - - _, err = Mul64(maxUint64-1, 2) - require.ErrorIs(err, ErrOverflow) -} - -func TestAbsDiff(t *testing.T) { - require := require.New(t) - - require.Equal(maxUint64, AbsDiff(0, maxUint64)) - require.Equal(maxUint64, AbsDiff(maxUint64, 0)) - require.Equal(uint64(2), AbsDiff(uint64(3), uint64(1))) - require.Equal(uint64(2), AbsDiff(uint64(1), uint64(3))) - require.Zero(AbsDiff(uint64(1), uint64(1))) - require.Zero(AbsDiff(uint64(0), uint64(0))) -} diff --git a/safe.go b/safe.go index 93b2f4d..c833712 100644 --- a/safe.go +++ b/safe.go @@ -59,3 +59,23 @@ func Mul[T constraints.Unsigned](a, b T) (T, error) { func AbsDiff[T constraints.Unsigned](a, b T) T { return max(a, b) - min(a, b) } + +// SafeAdd returns x+y and whether overflow occurred. +func SafeAdd(x, y uint64) (uint64, bool) { + sum := x + y + return sum, sum < x +} + +// SafeSub returns x-y and whether underflow occurred. +func SafeSub(x, y uint64) (uint64, bool) { + return x - y, x < y +} + +// SafeMul returns x*y and whether overflow occurred. +func SafeMul(x, y uint64) (uint64, bool) { + if x == 0 || y == 0 { + return 0, false + } + result := x * y + return result, result/y != x +} diff --git a/safe/go.mod b/safe/go.mod new file mode 100644 index 0000000..80f09a5 --- /dev/null +++ b/safe/go.mod @@ -0,0 +1,5 @@ +module github.com/luxfi/math/safe + +go 1.24.0 + +require golang.org/x/exp v0.0.0-20241217172543-b2144cdd0a67 diff --git a/safe/safe.go b/safe/safe.go new file mode 100644 index 0000000..7a392eb --- /dev/null +++ b/safe/safe.go @@ -0,0 +1,81 @@ +// Copyright (C) 2019-2025, Lux Industries Inc. All rights reserved. +// See the file LICENSE for licensing terms. + +// Package safe provides overflow-safe arithmetic operations. +package safe + +import ( + "errors" + "math" + "math/bits" + + "golang.org/x/exp/constraints" +) + +var ( + ErrOverflow = errors.New("overflow") + ErrUnderflow = errors.New("underflow") +) + +// Add64 returns a + b, or error if overflow. +func Add64(a, b uint64) (uint64, error) { + if a > math.MaxUint64-b { + return 0, ErrOverflow + } + return a + b, nil +} + +// Sub returns a - b, or error if underflow. +func Sub[T constraints.Unsigned](a, b T) (T, error) { + if a < b { + return 0, ErrUnderflow + } + return a - b, nil +} + +// Mul64 returns a * b, or error if overflow. +func Mul64(a, b uint64) (uint64, error) { + if b != 0 && a > math.MaxUint64/b { + return 0, ErrOverflow + } + return a * b, nil +} + +// SafeAdd returns x+y and whether overflow occurred. +func SafeAdd(x, y uint64) (uint64, bool) { + sum, carryOut := bits.Add64(x, y, 0) + return sum, carryOut != 0 +} + +// SafeSub returns x-y and whether underflow occurred. +func SafeSub(x, y uint64) (uint64, bool) { + diff, borrowOut := bits.Sub64(x, y, 0) + return diff, borrowOut != 0 +} + +// SafeMul returns x*y and whether overflow occurred. +func SafeMul(x, y uint64) (uint64, bool) { + hi, lo := bits.Mul64(x, y) + return lo, hi != 0 +} + +// AbsDiff returns |a - b|. +func AbsDiff[T constraints.Unsigned](a, b T) T { + return max(a, b) - min(a, b) +} + +// Min returns the minimum of two uint64 values. +func Min(a, b uint64) uint64 { + if a < b { + return a + } + return b +} + +// Max returns the maximum of two uint64 values. +func Max(a, b uint64) uint64 { + if a > b { + return a + } + return b +} diff --git a/math/sync_averager.go b/sync_averager.go similarity index 100% rename from math/sync_averager.go rename to sync_averager.go