diff --git a/v1/topdown/aggregates.go b/v1/topdown/aggregates.go index 03d07668d8..356e7b38b0 100644 --- a/v1/topdown/aggregates.go +++ b/v1/topdown/aggregates.go @@ -45,12 +45,13 @@ func builtinSum(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) err // Non-integer values found, so we need to sum as floats. sum := big.NewFloat(0) + tmp := new(big.Float) err := a.Iter(func(x *ast.Term) error { n, ok := x.Value.(ast.Number) if !ok { return builtins.NewOperandElementErr(1, a, x.Value, "number") } - sum = new(big.Float).Add(sum, builtins.NumberToFloat(n)) + sum = new(big.Float).Add(sum, builtins.NumberToFloatInto(tmp, n)) return nil }) if err != nil { @@ -74,12 +75,13 @@ func builtinSum(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) err } sum := big.NewFloat(0) + tmp := new(big.Float) err := a.Iter(func(x *ast.Term) error { n, ok := x.Value.(ast.Number) if !ok { return builtins.NewOperandElementErr(1, a, x.Value, "number") } - sum = new(big.Float).Add(sum, builtins.NumberToFloat(n)) + sum = new(big.Float).Add(sum, builtins.NumberToFloatInto(tmp, n)) return nil }) if err != nil { @@ -94,12 +96,13 @@ func builtinProduct(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) switch a := operands[0].Value.(type) { case *ast.Array: product := big.NewFloat(1) + tmp := new(big.Float) err := a.Iter(func(x *ast.Term) error { n, ok := x.Value.(ast.Number) if !ok { return builtins.NewOperandElementErr(1, a, x.Value, "number") } - product = new(big.Float).Mul(product, builtins.NumberToFloat(n)) + product = new(big.Float).Mul(product, builtins.NumberToFloatInto(tmp, n)) return nil }) if err != nil { @@ -108,12 +111,13 @@ func builtinProduct(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) return iter(ast.NewTerm(builtins.FloatToNumber(product))) case ast.Set: product := big.NewFloat(1) + tmp := new(big.Float) err := a.Iter(func(x *ast.Term) error { n, ok := x.Value.(ast.Number) if !ok { return builtins.NewOperandElementErr(1, a, x.Value, "number") } - product = new(big.Float).Mul(product, builtins.NumberToFloat(n)) + product = new(big.Float).Mul(product, builtins.NumberToFloatInto(tmp, n)) return nil }) if err != nil { diff --git a/v1/topdown/builtins/builtins.go b/v1/topdown/builtins/builtins.go index 98d80b5f9b..c81b772484 100644 --- a/v1/topdown/builtins/builtins.go +++ b/v1/topdown/builtins/builtins.go @@ -251,11 +251,18 @@ func ArrayOperand(x ast.Value, pos int) (*ast.Array, error) { // NumberToFloat converts n to a big float. func NumberToFloat(n ast.Number) *big.Float { - r, ok := new(big.Float).SetString(string(n)) - if !ok { + return NumberToFloatInto(nil, n) +} + +// NumberToFloatInto converts n to a big float, storing it in dst when provided. +func NumberToFloatInto(dst *big.Float, n ast.Number) *big.Float { + if dst == nil { + dst = new(big.Float) + } + if _, ok := dst.SetString(string(n)); !ok { panic("illegal value") } - return r + return dst } // FloatToNumber converts f to a number. diff --git a/v1/topdown/builtins_test.go b/v1/topdown/builtins_test.go index 5489a28af7..ffaf42f903 100644 --- a/v1/topdown/builtins_test.go +++ b/v1/topdown/builtins_test.go @@ -1,9 +1,11 @@ package topdown import ( + "math/big" "testing" "github.com/open-policy-agent/opa/v1/ast" + "github.com/open-policy-agent/opa/v1/topdown/builtins" "github.com/open-policy-agent/opa/v1/types" ) @@ -42,3 +44,24 @@ func TestCustomBuiltinIterator(t *testing.T) { t.Fatal("Expected x to be 2 but got:", rs[0]) } } + +func TestNumberToFloatInto(t *testing.T) { + t.Parallel() + + first := builtins.NumberToFloatInto(nil, ast.Number("1.5")) + if first == nil { + t.Fatal("expected non-nil float") + } + if first.Cmp(big.NewFloat(1.5)) != 0 { + t.Fatalf("expected 1.5, got %v", first) + } + + reuse := big.NewFloat(0) + second := builtins.NumberToFloatInto(reuse, ast.Number("2.75")) + if second != reuse { + t.Fatal("expected reuse of provided float") + } + if reuse.Cmp(big.NewFloat(2.75)) != 0 { + t.Fatalf("expected 2.75, got %v", reuse) + } +}