Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
166 changes: 166 additions & 0 deletions numeric_comparison_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
package expr_test

import (
"math"
"math/big"
"testing"

"github.com/expr-lang/expr"
"github.com/expr-lang/expr/types"
)

func TestNumericComparisonUint64(t *testing.T) {
env := map[string]any{"n": uint64(math.MaxUint64), "zero": uint64(0)}
program, err := expr.Compile("n > zero", expr.Env(env))
if err != nil {
t.Fatal(err)
}
got, err := expr.Run(program, env)
if err != nil {
t.Fatal(err)
}
if got != true {
t.Fatalf("MaxUint64 > uint64(0) = %v, want true", got)
}
}

func FuzzNumericComparison(f *testing.F) {
for _, seed := range []struct {
i int64
u uint64
v float64
}{
{-1, math.MaxUint64, -1},
{9007199254740993, 9007199254740993, 9007199254740992},
{math.MaxInt64, math.MaxUint64, 0x1p64},
{math.MinInt64, 0, -0x1p63},
{0, 0, math.SmallestNonzeroFloat64},
{0, 0, -math.SmallestNonzeroFloat64},
{1, 1, 1.5},
{-1, 1, -1.5},
} {
f.Add(seed.i, seed.u, math.Float64bits(seed.v))
}
program, err := expr.Compile("[a == b, a != b, a < b, a <= b, a > b, a >= b]", expr.Env(types.Map{"a": types.Any, "b": types.Any}))
if err != nil {
f.Fatal(err)
}
f.Fuzz(func(t *testing.T, i int64, u, floatBits uint64) {
v := math.Float64frombits(floatBits)
if math.IsNaN(v) || math.IsInf(v, 0) {
return // IEEE unordered/infinite cases are covered by the table test.
}
integer := new(big.Rat).SetInt64(i)
unsigned := new(big.Rat).SetInt(new(big.Int).SetUint64(u))
floating := new(big.Rat).SetFloat64(v)
for _, pair := range []struct {
a, b any
order int
}{
{i, u, integer.Cmp(unsigned)},
{i, v, integer.Cmp(floating)},
{u, v, unsigned.Cmp(floating)},
} {
for reverse := 0; reverse < 2; reverse++ {
a, b, order := pair.a, pair.b, pair.order
if reverse == 1 {
a, b, order = b, a, -order
}
got, err := expr.Run(program, map[string]any{"a": a, "b": b})
if err != nil {
t.Fatal(err)
}
want := []any{order == 0, order != 0, order < 0, order <= 0, order > 0, order >= 0}
for index, value := range got.([]any) {
if value != want[index] {
t.Errorf("%T(%v) versus %T(%v): comparison %d = %v, want %v", a, a, b, b, index, value, want[index])
}
}
}
}
})
}

func TestNumericComparisonSignedUnsigned(t *testing.T) {
for _, test := range []struct {
name string
a, b any
order int
}{
{"negative and unsigned max", int64(-1), uint64(math.MaxUint64), -1},
{"signed and unsigned max", int64(math.MaxInt64), uint64(math.MaxUint64), -1},
{"min signed and unsigned zero", int64(math.MinInt64), uint64(0), -1},
{"equal zero", int8(0), uint8(0), 0},
{"equal positive", int64(42), uint32(42), 0},
{"larger signed positive", int32(43), uint64(42), 1},
} {
t.Run(test.name, func(t *testing.T) {
checkNumericComparisons(t, test.a, test.b, test.order, true)
checkNumericComparisons(t, test.b, test.a, -test.order, true)
})
}
}

func TestNumericComparisonIntegerFloat(t *testing.T) {
for _, test := range []struct {
name string
a, b any
order int
}{
{"signed beyond exact float range", int64(9007199254740993), float64(9007199254740992), 1},
{"unsigned beyond exact float range", uint64(9007199254740993), float64(9007199254740992), 1},
{"signed max rounded up", int64(math.MaxInt64), float64(0x1p63), -1},
{"unsigned max rounded up", uint64(math.MaxUint64), float64(0x1p64), -1},
{"signed min", int64(math.MinInt64), float64(-0x1p63), 0},
{"signed min rounded down", int64(math.MinInt64), math.Nextafter(-0x1p63, math.Inf(-1)), 1},
{"positive fraction", int64(1), float64(1.5), -1},
{"negative fraction", int64(-1), float64(-1.5), 1},
{"unsigned and negative fraction", uint64(0), float64(-0.5), 1},
{"negative zero", int64(0), math.Copysign(0, -1), 0},
{"positive infinity", uint64(math.MaxUint64), math.Inf(1), -1},
{"negative infinity", int64(math.MinInt64), math.Inf(-1), 1},
{"float32 integer", int64(16777217), float32(16777216), 1},
} {
t.Run(test.name, func(t *testing.T) {
checkNumericComparisons(t, test.a, test.b, test.order, true)
checkNumericComparisons(t, test.b, test.a, -test.order, true)
})
}
for _, value := range []any{int64(0), uint64(math.MaxUint64)} {
checkNumericComparisons(t, value, math.NaN(), 0, false)
checkNumericComparisons(t, math.NaN(), value, 0, false)
}
}

func checkNumericComparisons(t *testing.T, a, b any, order int, ordered bool) {
t.Helper()
env := map[string]any{"a": a, "b": b}
for _, mode := range []struct {
name string
shape any
}{
{"typed", env},
{"dynamic", types.Map{"a": types.Any, "b": types.Any}},
} {
for _, operator := range []struct {
name string
want bool
}{
{"==", ordered && order == 0}, {"!=", !ordered || order != 0},
{"<", ordered && order < 0}, {"<=", ordered && order <= 0},
{">", ordered && order > 0}, {">=", ordered && order >= 0},
} {
program, err := expr.Compile("a "+operator.name+" b", expr.Env(mode.shape))
if err != nil {
t.Fatal(err)
}
got, err := expr.Run(program, env)
if err != nil {
t.Fatal(err)
}
if got != operator.want {
t.Errorf("%s: %T(%v) %s %T(%v) = %v, want %v", mode.name, a, a, operator.name, b, b, got, operator.want)
}
}
}
}
56 changes: 56 additions & 0 deletions vm/runtime/comparison.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
package runtime

import "math"

// Compare the integral part before examining the fraction, so converting a
// large integer to float64 cannot round it into equality. The bounds keep
// float-to-integer conversions within their defined range.
func compareIntFloat(x int64, y float64) (int, bool) {
if math.IsNaN(y) {
return 0, false
}
if y >= 0x1p63 {
return -1, true
}
if y < -0x1p63 {
return 1, true
}
whole := int64(y)
if x < whole {
return -1, true
}
if x > whole {
return 1, true
}
return compareFraction(float64(whole), y), true
}

func compareUintFloat(x uint64, y float64) (int, bool) {
if math.IsNaN(y) {
return 0, false
}
if y >= 0x1p64 {
return -1, true
}
if y < 0 {
return 1, true
}
whole := uint64(y)
if x < whole {
return -1, true
}
if x > whole {
return 1, true
}
return compareFraction(float64(whole), y), true
}

func compareFraction(whole, value float64) int {
if whole < value {
return -1
}
if whole > value {
return 1
}
return 0
}
49 changes: 48 additions & 1 deletion vm/runtime/helpers/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,14 +71,31 @@ func cases(op string, xs ...[]string) string {
echo(`case %v:`, a)
echo(`switch y := b.(type) {`)
for _, b := range types {
echo(`case %v:`, b)
if isComparison(op) && !isFloat(a) && !isFloat(b) {
echo(`return %s`, integerComparison(op, a, b))
continue
}
if isComparison(op) && isFloat(a) != isFloat(b) {
integer, x, y, sign := a, "x", "y", ""
if isFloat(a) {
integer, x, y, sign = b, "y", "x", "-"
}
typ, fn := "int64", "compareIntFloat"
if strings.HasPrefix(integer, "uint") {
typ, fn = "uint64", "compareUintFloat"
}
echo(`comparison, ordered := %s(%s(%s), float64(%s))`, fn, typ, x, y)
echo(`return ordered && %scomparison %s 0`, sign, op)
continue
}
t := "int"
if isDuration(a) || isDuration(b) {
t = "time.Duration"
}
if isFloat(a) || isFloat(b) {
t = "float64"
}
echo(`case %v:`, b)
if op == "/" {
echo(`return float64(x) / float64(y)`)
} else {
Expand Down Expand Up @@ -133,6 +150,36 @@ func isFloat(t string) bool {
return strings.HasPrefix(t, "float")
}

func isComparison(op string) bool {
switch op {
case "==", "<", ">", "<=", ">=":
return true
default:
return false
}
}

func integerComparison(op, a, b string) string {
aUnsigned, bUnsigned := strings.HasPrefix(a, "uint"), strings.HasPrefix(b, "uint")
if !aUnsigned && !bUnsigned {
return fmt.Sprintf("int64(x) %s int64(y)", op)
}
comparison := fmt.Sprintf("uint64(x) %s uint64(y)", op)
if aUnsigned && bUnsigned {
return comparison
}
if aUnsigned {
if op == ">" || op == ">=" {
return "y < 0 || " + comparison
}
return "y >= 0 && " + comparison
}
if op == "<" || op == "<=" {
return "x < 0 || " + comparison
}
return "x >= 0 && " + comparison
}

func isDuration(t string) bool {
return t == "time.Duration"
}
Expand Down
Loading