Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
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
78 changes: 73 additions & 5 deletions arrow/compute/expression.go
Original file line number Diff line number Diff line change
Expand Up @@ -377,7 +377,43 @@ func (c *Call) Equals(other Expression) bool {
if opt, ok := c.options.(FunctionOptionsEqual); ok {
return opt.Equals(rhs.options)
}
return reflect.DeepEqual(c.options, rhs.options)
return equalFunctionOptions(c.options, rhs.options)
}

func equalFunctionOptions(lhs, rhs FunctionOptions) bool {
if left, ok := cumulativeOptions(lhs); ok {
right, ok := cumulativeOptions(rhs)
if !ok {
return false
}
if left == nil || right == nil {
return left == nil && right == nil
}
return left.SkipNulls == right.SkipNulls && equalOptionalScalar(left.Start, right.Start)
}

if lhs == nil || rhs == nil {
return lhs == nil && rhs == nil
}
return reflect.DeepEqual(lhs, rhs)
}

func cumulativeOptions(opts FunctionOptions) (*CumulativeOptions, bool) {
switch opts := opts.(type) {
case CumulativeOptions:
return &opts, true
case *CumulativeOptions:
return opts, true
default:
return nil, false
}
}

func equalOptionalScalar(lhs, rhs scalar.Scalar) bool {
if lhs == nil || rhs == nil {
return lhs == nil && rhs == nil
}
return scalar.Equals(lhs, rhs)
}

func (c *Call) Release() {
Expand Down Expand Up @@ -533,6 +569,7 @@ var (
funcOptsTypes = []FunctionOptions{
SetLookupOptions{}, ArithmeticOptions{}, CastOptions{},
FilterOptions{}, NullOptions{}, StrptimeOptions{}, MakeStructOptions{},
CumulativeOptions{},
}
)

Expand Down Expand Up @@ -565,9 +602,34 @@ func NewFieldRef(field string) Expression {
}

// NewCall constructs an expression that represents a specific function call with
// the given arguments and options.
// the given arguments and options. Cumulative start scalars are retained for
// the lifetime of the expression.
func NewCall(name string, args []Expression, opts FunctionOptions) Expression {
return &Call{funcName: name, args: args, options: opts}
return &Call{funcName: name, args: args, options: cloneExpressionOptions(opts)}
}

func cloneExpressionOptions(opts FunctionOptions) FunctionOptions {
switch opts := opts.(type) {
case CumulativeOptions:
opts.Start = retainExpressionScalar(opts.Start)
return opts
case *CumulativeOptions:
if opts == nil {
return nil
}
cloned := *opts
cloned.Start = retainExpressionScalar(cloned.Start)
return cloned
default:
return opts
}
}

func retainExpressionScalar(value scalar.Scalar) scalar.Scalar {
if releasable, ok := value.(scalar.Releasable); ok {
releasable.Retain()
}
return value
}

// Project is shorthand for `make_struct` to produce a record batch output
Expand Down Expand Up @@ -880,13 +942,19 @@ func DeserializeExpr(mem memory.Allocator, buf *memory.Buffer) (Expression, erro
}

optionsVal := reflect.New(funcOptionsMap[string(typname.(*scalar.Binary).Data())]).Interface()
if err := scalar.FromScalar(optsScalar.(*scalar.Struct), optionsVal); err != nil {
if err := scalar.FromScalarWithAllocator(optsScalar.(*scalar.Struct), optionsVal, mem); err != nil {
return nil, err
}
opts = optionsVal.(FunctionOptions)
}
index += 2
return NewCall(val, args, opts), nil
expr := NewCall(val, args, opts)
if _, ok := cumulativeOptions(opts); ok {
if r, ok := opts.(releasable); ok {
r.Release()
}
}
return expr, nil
}

arg, err := getone()
Expand Down
143 changes: 143 additions & 0 deletions arrow/compute/expression_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,12 @@ import (
"github.com/stretchr/testify/assert"
)

type privateFunctionOptions struct {
value int
}

func (privateFunctionOptions) TypeName() string { return "privateFunctionOptions" }

func TestExpressionToString(t *testing.T) {
ts, _ := scalar.MakeScalar("1990-10-23 10:23:33.123456").CastTo(arrow.FixedWidthTypes.Timestamp_ns)

Expand Down Expand Up @@ -115,6 +121,143 @@ func TestExpressionEquality(t *testing.T) {
}
}

func TestExpressionEqualityWithPrivateFunctionOptions(t *testing.T) {
left := compute.NewCall("test", nil, privateFunctionOptions{value: 1})
right := compute.NewCall("test", nil, privateFunctionOptions{value: 1})
different := compute.NewCall("test", nil, privateFunctionOptions{value: 2})
defer left.Release()
defer right.Release()
defer different.Release()

assert.NotPanics(t, func() {
assert.True(t, left.Equals(right))
assert.False(t, left.Equals(different))
})
}

func TestCumulativeOptionsEquality(t *testing.T) {
newBinaryStart := func() scalar.Scalar {
buf := memory.NewBufferBytes([]byte("10"))
defer buf.Release()
return scalar.NewBinaryScalar(buf, arrow.BinaryTypes.Binary)
}

tests := []struct {
name string
leftStart, rightStart func() scalar.Scalar
leftSkip, rightSkip bool
want bool
}{
{
name: "both nil",
leftStart: func() scalar.Scalar { return nil },
rightStart: func() scalar.Scalar { return nil },
want: true,
},
{
name: "one nil",
leftStart: func() scalar.Scalar { return nil },
rightStart: func() scalar.Scalar { return scalar.NewInt32Scalar(10) },
want: false,
},
{
name: "equal numeric scalars",
leftStart: func() scalar.Scalar { return scalar.NewInt32Scalar(10) },
rightStart: func() scalar.Scalar { return scalar.NewInt32Scalar(10) },
want: true,
},
{
name: "equal string scalars",
leftStart: func() scalar.Scalar { return scalar.NewStringScalar("10") },
rightStart: func() scalar.Scalar { return scalar.NewStringScalar("10") },
want: true,
},
{
name: "equal binary scalars",
leftStart: newBinaryStart,
rightStart: newBinaryStart,
want: true,
},
{
name: "different scalar values",
leftStart: func() scalar.Scalar { return scalar.NewInt32Scalar(10) },
rightStart: func() scalar.Scalar { return scalar.NewInt32Scalar(11) },
want: false,
},
{
name: "different scalar types",
leftStart: func() scalar.Scalar { return scalar.NewInt32Scalar(10) },
rightStart: func() scalar.Scalar { return scalar.NewInt64Scalar(10) },
want: false,
},
{
name: "different skip nulls",
leftStart: func() scalar.Scalar { return nil },
rightStart: func() scalar.Scalar { return nil },
leftSkip: false,
rightSkip: true,
want: false,
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
left := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")},
&compute.CumulativeOptions{Start: tc.leftStart(), SkipNulls: tc.leftSkip})
right := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")},
&compute.CumulativeOptions{Start: tc.rightStart(), SkipNulls: tc.rightSkip})
defer left.Release()
defer right.Release()

assert.Equal(t, tc.want, left.Equals(right))
})
}

}

func TestCumulativeOptionsValueAndPointerEquality(t *testing.T) {
value := compute.CumulativeOptions{Start: scalar.NewInt32Scalar(10)}
pointer := &compute.CumulativeOptions{Start: scalar.NewInt32Scalar(10)}

left := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")}, value)
right := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")}, pointer)
defer left.Release()
defer right.Release()

assert.True(t, left.Equals(right))
}

func TestCumulativeOptionsRelease(t *testing.T) {
mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
defer mem.AssertSize(t, 0)

newStart := func() scalar.Scalar {
data := mem.Allocate(2)
copy(data, []byte("10"))
buffer := memory.NewBufferWithAllocator(data, mem)
start := scalar.NewBinaryScalar(buffer, arrow.BinaryTypes.Binary)
buffer.Release()
return start
}

t.Run("pointer options", func(t *testing.T) {
start := newStart()
expr := compute.NewCall("cumulative_sum", nil,
&compute.CumulativeOptions{Start: start})
expr.Release()
assert.Equal(t, "10", string(start.(scalar.BinaryScalar).Data()))
start.(scalar.Releasable).Release()
})
t.Run("value options", func(t *testing.T) {
start := newStart()
expr := compute.NewCall("cumulative_sum", nil,
compute.CumulativeOptions{Start: start})
expr.Release()
assert.Equal(t, "10", string(start.(scalar.BinaryScalar).Data()))
start.(scalar.Releasable).Release()
})
}

func TestExpressionHashing(t *testing.T) {
set := make(map[uint64]compute.Expression)

Expand Down
Loading