diff --git a/arrow/compute/expression.go b/arrow/compute/expression.go index 2f1dc482..0d7c23ad 100644 --- a/arrow/compute/expression.go +++ b/arrow/compute/expression.go @@ -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() { @@ -533,6 +569,7 @@ var ( funcOptsTypes = []FunctionOptions{ SetLookupOptions{}, ArithmeticOptions{}, CastOptions{}, FilterOptions{}, NullOptions{}, StrptimeOptions{}, MakeStructOptions{}, + CumulativeOptions{}, } ) @@ -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 @@ -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() diff --git a/arrow/compute/expression_test.go b/arrow/compute/expression_test.go index 42f64394..9f4d720e 100644 --- a/arrow/compute/expression_test.go +++ b/arrow/compute/expression_test.go @@ -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) @@ -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) diff --git a/arrow/compute/internal/kernels/vector_cumulative.go b/arrow/compute/internal/kernels/vector_cumulative.go new file mode 100644 index 00000000..006a1261 --- /dev/null +++ b/arrow/compute/internal/kernels/vector_cumulative.go @@ -0,0 +1,411 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.18 + +package kernels + +import ( + "fmt" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/bitutil" + "github.com/apache/arrow-go/v18/arrow/compute/exec" + "github.com/apache/arrow-go/v18/arrow/scalar" +) + +// CumulativeOptions controls cumulative operations. +type CumulativeOptions struct { + // Start is the initial value. A nil value uses the zero value for the + // input type. A non-nil null scalar is invalid. + Start scalar.Scalar `compute:"start"` + // SkipNulls controls whether a null input stops the cumulative operation. + // Null positions are still null in the output when this is true. + SkipNulls bool `compute:"skip_nulls"` +} + +func (CumulativeOptions) TypeName() string { return "CumulativeOptions" } + +func (opts CumulativeOptions) Release() { + if opts.Start == nil { + return + } + + if releasable, ok := opts.Start.(interface{ Release() }); ok { + releasable.Release() + } +} + +type cumulativeSumState[T arrow.NumericType] struct { + current T + skipNulls bool + encounteredNull bool + checked bool + add func(T, T) (T, error) +} + +type ScalarCastFn func(*exec.KernelCtx, scalar.Scalar, arrow.DataType) (scalar.Scalar, error) + +func safeNumericCastScalar(ctx *exec.KernelCtx, cast ScalarCastFn, start scalar.Scalar, typ arrow.DataType) (scalar.Scalar, error) { + targetID := typ.ID() + if !arrow.IsInteger(targetID) && !arrow.IsFloating(targetID) { + return nil, fmt.Errorf("%w: cumulative sum input type must be numeric, got %s", arrow.ErrType, typ) + } + if arrow.TypeEqual(start.DataType(), typ) { + return start, nil + } + + if cast == nil { + return nil, fmt.Errorf("%w: cumulative sum start value caster is not configured", arrow.ErrInvalid) + } + + casted, err := cast(ctx, start, typ) + if err != nil { + return nil, fmt.Errorf("%w: cannot cast cumulative sum start value to %s: %v", arrow.ErrInvalid, typ, err) + } + return casted, nil +} + +func cumulativeStartValue[T arrow.NumericType](ctx *exec.KernelCtx, cast ScalarCastFn, start scalar.Scalar, typ arrow.DataType) (T, error) { + var zero T + if start == nil { + return zero, nil + } + if !start.IsValid() { + return zero, fmt.Errorf("%w: cumulative sum start value must be valid", arrow.ErrInvalid) + } + + casted, err := safeNumericCastScalar(ctx, cast, start, typ) + if err != nil { + return zero, err + } + if releasable, ok := casted.(scalar.Releasable); ok { + defer releasable.Release() + } + + primitive, ok := casted.(scalar.PrimitiveScalar) + if !ok { + return zero, fmt.Errorf("%w: cumulative sum start value is not primitive", arrow.ErrInvalid) + } + + span := &exec.ArraySpan{Type: typ, Len: 1} + span.Buffers[1].Buf = primitive.Data() + return exec.GetSpanValues[T](span, 1)[0], nil +} + +func initCumulativeSum[T arrow.NumericType](checked bool, cast ScalarCastFn) exec.KernelInitFn { + return func(ctx *exec.KernelCtx, args exec.KernelInitArgs) (exec.KernelState, error) { + opts := CumulativeOptions{} + switch value := args.Options.(type) { + case nil: + case CumulativeOptions: + opts = value + case *CumulativeOptions: + if value != nil { + opts = *value + } + default: + return nil, fmt.Errorf("%w: attempted to initialize cumulative sum from invalid function options", arrow.ErrInvalid) + } + + start, err := cumulativeStartValue[T](ctx, cast, opts.Start, args.Inputs[0]) + if err != nil { + return nil, err + } + + return &cumulativeSumState[T]{ + current: start, + skipNulls: opts.SkipNulls, + checked: checked, + add: checkedAdder[T](), + }, nil + } +} + +func checkedAddSigned[T arrow.IntType](left, right T) (T, error) { + if (right > 0 && left > MaxOf[T]()-right) || (right < 0 && left < MinOf[T]()-right) { + return 0, errOverflow + } + return left + right, nil +} + +func checkedAddUnsigned[T arrow.UintType](left, right T) (T, error) { + if left > MaxOf[T]()-right { + return 0, errOverflow + } + return left + right, nil +} + +func checkedAdder[T arrow.NumericType]() func(T, T) (T, error) { + var zero T + switch any(zero).(type) { + case int8: + return func(left, right T) (T, error) { + value, err := checkedAddSigned(int8(left), int8(right)) + return T(value), err + } + case int16: + return func(left, right T) (T, error) { + value, err := checkedAddSigned(int16(left), int16(right)) + return T(value), err + } + case int32: + return func(left, right T) (T, error) { + value, err := checkedAddSigned(int32(left), int32(right)) + return T(value), err + } + case int64: + return func(left, right T) (T, error) { + value, err := checkedAddSigned(int64(left), int64(right)) + return T(value), err + } + case uint8: + return func(left, right T) (T, error) { + value, err := checkedAddUnsigned(uint8(left), uint8(right)) + return T(value), err + } + case uint16: + return func(left, right T) (T, error) { + value, err := checkedAddUnsigned(uint16(left), uint16(right)) + return T(value), err + } + case uint32: + return func(left, right T) (T, error) { + value, err := checkedAddUnsigned(uint32(left), uint32(right)) + return T(value), err + } + case uint64: + return func(left, right T) (T, error) { + value, err := checkedAddUnsigned(uint64(left), uint64(right)) + return T(value), err + } + default: + return func(left, right T) (T, error) { return left + right, nil } + } +} + +func prepareCumulativeOutput[T arrow.NumericType](ctx *exec.KernelCtx, out *exec.ExecResult, needsValidity bool) { + if out.Len == 0 { + return + } + + data := ctx.Allocate(int(out.Len) * arrow.GetDataType[T]().(arrow.FixedWidthDataType).Bytes()) + out.Buffers[1].WrapBuffer(data) + + if needsValidity { + validity := ctx.AllocateBitmap(out.Len) + validityBytes := validity.Bytes() + for i := range validityBytes { + validityBytes[i] = 0xFF + } + out.Buffers[0].WrapBuffer(validity) + } +} + +func cumulativeSumNoNulls[T arrow.NumericType](state *cumulativeSumState[T], inputs []*exec.ArraySpan, values []T) { + var outputOffset int64 + current := state.current + for _, input := range inputs { + inputValues := exec.GetSpanValues[T](input, 1) + for i := int64(0); i < input.Len; i++ { + current += inputValues[i] + values[outputOffset+i] = current + } + outputOffset += input.Len + } + state.current = current +} + +func cumulativeSumNoNullsChecked[T arrow.NumericType](state *cumulativeSumState[T], inputs []*exec.ArraySpan, values []T) error { + var outputOffset int64 + current := state.current + for _, input := range inputs { + inputValues := exec.GetSpanValues[T](input, 1) + for i := int64(0); i < input.Len; i++ { + var err error + current, err = state.add(current, inputValues[i]) + if err != nil { + return err + } + values[outputOffset+i] = current + } + outputOffset += input.Len + } + state.current = current + return nil +} + +func cumulativeSumWithNulls[T arrow.NumericType](state *cumulativeSumState[T], inputs []*exec.ArraySpan, values []T, validity []byte) int64 { + var ( + nulls int64 + outputOffset int64 + ) + current := state.current + for _, input := range inputs { + inputValues := exec.GetSpanValues[T](input, 1) + for i := int64(0); i < input.Len; i++ { + valid := len(input.Buffers[0].Buf) == 0 || bitutil.BitIsSet(input.Buffers[0].Buf, int(input.Offset+i)) + outputIndex := outputOffset + i + if !valid || state.encounteredNull { + bitutil.ClearBit(validity, int(outputIndex)) + nulls++ + if !valid && !state.skipNulls { + state.encounteredNull = true + } + continue + } + + current += inputValues[i] + values[outputIndex] = current + } + outputOffset += input.Len + } + state.current = current + return nulls +} + +func cumulativeSumWithNullsChecked[T arrow.NumericType](state *cumulativeSumState[T], inputs []*exec.ArraySpan, values []T, validity []byte) (int64, error) { + var ( + nulls int64 + outputOffset int64 + ) + current := state.current + for _, input := range inputs { + inputValues := exec.GetSpanValues[T](input, 1) + for i := int64(0); i < input.Len; i++ { + valid := len(input.Buffers[0].Buf) == 0 || bitutil.BitIsSet(input.Buffers[0].Buf, int(input.Offset+i)) + outputIndex := outputOffset + i + if !valid || state.encounteredNull { + bitutil.ClearBit(validity, int(outputIndex)) + nulls++ + if !valid && !state.skipNulls { + state.encounteredNull = true + } + continue + } + + var err error + current, err = state.add(current, inputValues[i]) + if err != nil { + return nulls, err + } + values[outputIndex] = current + } + outputOffset += input.Len + } + state.current = current + return nulls, nil +} + +func cumulativeSumSpans[T arrow.NumericType](ctx *exec.KernelCtx, state *cumulativeSumState[T], inputs []*exec.ArraySpan, out *exec.ExecResult, needsValidity bool) error { + prepareCumulativeOutput[T](ctx, out, needsValidity) + values := exec.GetSpanValues[T](out, 1) + + if !needsValidity && !state.encounteredNull { + if state.checked { + if err := cumulativeSumNoNullsChecked(state, inputs, values); err != nil { + out.Release() + return err + } + } else { + cumulativeSumNoNulls(state, inputs, values) + } + return nil + } + + var ( + nulls int64 + err error + ) + if state.checked { + nulls, err = cumulativeSumWithNullsChecked(state, inputs, values, out.Buffers[0].Buf) + } else { + nulls = cumulativeSumWithNulls(state, inputs, values, out.Buffers[0].Buf) + } + if err != nil { + out.Release() + return err + } + out.Nulls = nulls + return nil +} + +func cumulativeSumExec[T arrow.NumericType](ctx *exec.KernelCtx, batch *exec.ExecSpan, out *exec.ExecResult) error { + state := ctx.State.(*cumulativeSumState[T]) + input := &batch.Values[0].Array + + out.Len = input.Len + if input.Len == 0 { + return nil + } + + return cumulativeSumSpans(ctx, state, []*exec.ArraySpan{input}, out, state.encounteredNull || input.MayHaveNulls()) +} + +func cumulativeSumExecChunked[T arrow.NumericType](ctx *exec.KernelCtx, batch []*arrow.Chunked, out *exec.ExecResult) ([]*exec.ExecResult, error) { + state := ctx.State.(*cumulativeSumState[T]) + input := batch[0] + out.Len = int64(input.Len()) + if out.Len == 0 { + return []*exec.ExecResult{out}, nil + } + + chunks := input.Chunks() + spans := make([]exec.ArraySpan, len(chunks)) + inputs := make([]*exec.ArraySpan, len(chunks)) + needsValidity := state.encounteredNull || input.NullN() != 0 + for i, chunk := range chunks { + spans[i].SetMembers(chunk.Data()) + inputs[i] = &spans[i] + needsValidity = needsValidity || spans[i].MayHaveNulls() + } + + if err := cumulativeSumSpans(ctx, state, inputs, out, needsValidity); err != nil { + return nil, err + } + return []*exec.ExecResult{out}, nil +} + +func newCumulativeSumKernel[T arrow.NumericType](typ arrow.DataType, checked bool, cast ScalarCastFn) exec.VectorKernel { + kernel := exec.NewVectorKernel( + []exec.InputType{exec.NewExactInput(typ)}, + exec.NewOutputType(typ), + cumulativeSumExec[T], + initCumulativeSum[T](checked, cast)) + kernel.Parallelizable = false + kernel.CanExecuteChunkWise = false + kernel.ExecChunked = cumulativeSumExecChunked[T] + return kernel +} + +func cumulativeSumKernels(checked bool, cast ScalarCastFn) []exec.VectorKernel { + return []exec.VectorKernel{ + newCumulativeSumKernel[int8](arrow.PrimitiveTypes.Int8, checked, cast), + newCumulativeSumKernel[int16](arrow.PrimitiveTypes.Int16, checked, cast), + newCumulativeSumKernel[int32](arrow.PrimitiveTypes.Int32, checked, cast), + newCumulativeSumKernel[int64](arrow.PrimitiveTypes.Int64, checked, cast), + newCumulativeSumKernel[uint8](arrow.PrimitiveTypes.Uint8, checked, cast), + newCumulativeSumKernel[uint16](arrow.PrimitiveTypes.Uint16, checked, cast), + newCumulativeSumKernel[uint32](arrow.PrimitiveTypes.Uint32, checked, cast), + newCumulativeSumKernel[uint64](arrow.PrimitiveTypes.Uint64, checked, cast), + newCumulativeSumKernel[float32](arrow.PrimitiveTypes.Float32, checked, cast), + newCumulativeSumKernel[float64](arrow.PrimitiveTypes.Float64, checked, cast), + } +} + +func GetVectorCumulativeKernels(cast ScalarCastFn) (sum, checked []exec.VectorKernel) { + return cumulativeSumKernels(false, cast), cumulativeSumKernels(true, cast) +} diff --git a/arrow/compute/registry.go b/arrow/compute/registry.go index f1be3b91..bea37025 100644 --- a/arrow/compute/registry.go +++ b/arrow/compute/registry.go @@ -54,6 +54,7 @@ func GetFunctionRegistry() FunctionRegistry { RegisterScalarArithmetic(registry) RegisterScalarComparisons(registry) RegisterVectorHash(registry) + RegisterVectorCumulative(registry) RegisterVectorRunEndFuncs(registry) RegisterScalarSetLookup(registry) }) diff --git a/arrow/compute/vector_cumulative.go b/arrow/compute/vector_cumulative.go new file mode 100644 index 00000000..cc7d2853 --- /dev/null +++ b/arrow/compute/vector_cumulative.go @@ -0,0 +1,101 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.18 + +package compute + +import ( + "context" + "fmt" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/compute/exec" + "github.com/apache/arrow-go/v18/arrow/compute/internal/kernels" + "github.com/apache/arrow-go/v18/arrow/scalar" +) + +var ( + cumulativeSumDoc = FunctionDoc{ + Summary: "Compute the cumulative sum of numeric input", + Description: `Return the cumulative sum of the input array. Integer values +wrap on overflow; nulls stop the remaining output unless SkipNulls is enabled. +A nil Start uses zero. For chunked input, accumulation continues across all +chunks and the result is returned as a chunked array.`, + ArgNames: []string{"values"}, + OptionsType: "CumulativeOptions", + } + cumulativeSumCheckedDoc = FunctionDoc{ + Summary: "Compute cumulative sum of numeric input with overflow checking", + Description: `Return the cumulative sum of the input array and report +integer overflow. Null handling and Start follow CumulativeOptions. For +chunked input, accumulation continues across all chunks and the result is +returned as a chunked array.`, + ArgNames: []string{"values"}, + OptionsType: "CumulativeOptions", + } +) + +type CumulativeOptions = kernels.CumulativeOptions + +func safeCastScalar(ctx *exec.KernelCtx, start scalar.Scalar, typ arrow.DataType) (scalar.Scalar, error) { + input := NewDatumWithoutOwning(start) + result, err := CastDatum(ctx.Ctx, input, SafeCastOptions(typ)) + if err != nil { + return nil, err + } + + casted, ok := result.(*ScalarDatum) + if !ok { + result.Release() + return nil, fmt.Errorf("%w: safe cast of cumulative sum start value returned %T", arrow.ErrInvalid, result) + } + + value := casted.Value + casted.Value = nil + result.Release() + return value, nil +} + +func RegisterVectorCumulative(reg FunctionRegistry) { + sum, checked := kernels.GetVectorCumulativeKernels(safeCastScalar) + + sumFn := NewVectorFunction("cumulative_sum", Unary(), cumulativeSumDoc) + sumFn.SetDefaultOptions(&CumulativeOptions{}) + for _, k := range sum { + if err := sumFn.AddKernel(k); err != nil { + panic(err) + } + } + reg.AddFunction(sumFn, false) + + checkedFn := NewVectorFunction("cumulative_sum_checked", Unary(), cumulativeSumCheckedDoc) + checkedFn.SetDefaultOptions(&CumulativeOptions{}) + for _, k := range checked { + if err := checkedFn.AddKernel(k); err != nil { + panic(err) + } + } + reg.AddFunction(checkedFn, false) +} + +func CumulativeSum(ctx context.Context, opts CumulativeOptions, values Datum) (Datum, error) { + return CallFunction(ctx, "cumulative_sum", &opts, values) +} + +func CumulativeSumChecked(ctx context.Context, opts CumulativeOptions, values Datum) (Datum, error) { + return CallFunction(ctx, "cumulative_sum_checked", &opts, values) +} diff --git a/arrow/compute/vector_cumulative_bench_test.go b/arrow/compute/vector_cumulative_bench_test.go new file mode 100644 index 00000000..9e8f00cd --- /dev/null +++ b/arrow/compute/vector_cumulative_bench_test.go @@ -0,0 +1,124 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.18 + +package compute_test + +import ( + "context" + "testing" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/compute" + "github.com/apache/arrow-go/v18/arrow/memory" +) + +func benchmarkInt64Array(b *testing.B, length int, withNulls bool) arrow.Array { + b.Helper() + builder := array.NewInt64Builder(memory.DefaultAllocator) + builder.Reserve(length) + for i := 0; i < length; i++ { + if withNulls && i%100 == 0 { + builder.AppendNull() + } else { + builder.Append(1) + } + } + result := builder.NewArray() + builder.Release() + return result +} + +func benchmarkFloat64Array(b *testing.B, length int) arrow.Array { + b.Helper() + builder := array.NewFloat64Builder(memory.DefaultAllocator) + builder.Reserve(length) + for i := 0; i < length; i++ { + builder.Append(1) + } + result := builder.NewArray() + builder.Release() + return result +} + +func benchmarkInt64Chunked(b *testing.B, chunks, length int) *arrow.Chunked { + b.Helper() + values := make([]arrow.Array, chunks) + for i := range values { + values[i] = benchmarkInt64Array(b, length, false) + } + result := arrow.NewChunked(arrow.PrimitiveTypes.Int64, values) + for _, value := range values { + value.Release() + } + return result +} + +func BenchmarkCumulativeSum(b *testing.B) { + const length = 10_000_000 + + intInput := benchmarkInt64Array(b, length, false) + defer intInput.Release() + nullInput := benchmarkInt64Array(b, length, true) + defer nullInput.Release() + floatInput := benchmarkFloat64Array(b, length) + defer floatInput.Release() + chunkedInput := benchmarkInt64Chunked(b, 100, length/100) + defer chunkedInput.Release() + + ctx := context.Background() + tests := []struct { + name string + input compute.Datum + opts compute.CumulativeOptions + checked bool + }{ + {name: "int64", input: &compute.ArrayDatum{Value: intInput.Data()}}, + {name: "int64_checked", input: &compute.ArrayDatum{Value: intInput.Data()}, checked: true}, + {name: "int64_1pct_nulls_skip", input: &compute.ArrayDatum{Value: nullInput.Data()}, opts: compute.CumulativeOptions{SkipNulls: true}}, + {name: "float64", input: &compute.ArrayDatum{Value: floatInput.Data()}}, + {name: "int64_chunked", input: &compute.ChunkedDatum{Value: chunkedInput}}, + } + + for _, tc := range tests { + b.Run(tc.name, func(b *testing.B) { + if input, ok := tc.input.(compute.ArrayLikeDatum); ok { + if typ, ok := input.Type().(arrow.FixedWidthDataType); ok { + b.SetBytes(int64(input.Len()) * int64(typ.Bytes())) + } + } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + var ( + result compute.Datum + err error + ) + if tc.checked { + result, err = compute.CumulativeSumChecked(ctx, tc.opts, tc.input) + } else { + result, err = compute.CumulativeSum(ctx, tc.opts, tc.input) + } + if err != nil { + b.Fatal(err) + } + result.Release() + } + }) + } +} diff --git a/arrow/compute/vector_cumulative_test.go b/arrow/compute/vector_cumulative_test.go new file mode 100644 index 00000000..8bbc853e --- /dev/null +++ b/arrow/compute/vector_cumulative_test.go @@ -0,0 +1,823 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.18 + +package compute_test + +import ( + "context" + "strings" + "testing" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/compute" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/arrow/scalar" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func cumulativeInput(t *testing.T, mem memory.Allocator, typ arrow.DataType, values string) arrow.Array { + arr, _, err := array.FromJSON(mem, typ, strings.NewReader(values)) + require.NoError(t, err) + return arr +} + +func TestCumulativeSum(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 2, 3, 4]`) + defer input.Release() + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 3, 6, 10]`) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + +} + +func TestCumulativeSumValueOptions(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 2, 3]`) + defer input.Release() + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 3, 6]`) + defer expected.Release() + + result, err := compute.CallFunction( + ctx, + "cumulative_sum", + compute.CumulativeOptions{}, + &compute.ArrayDatum{Value: input.Data()}, + ) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) +} + +func TestCumulativeSumAdditionalInputs(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + tests := []struct { + name string + typ arrow.DataType + in string + want string + }{ + {name: "empty", typ: arrow.PrimitiveTypes.Int32, in: `[]`, want: `[]`}, + {name: "all null", typ: arrow.PrimitiveTypes.Int32, in: `[null, null]`, want: `[null, null]`}, + {name: "uint8", typ: arrow.PrimitiveTypes.Uint8, in: `[1, 2, 3]`, want: `[1, 3, 6]`}, + {name: "float32", typ: arrow.PrimitiveTypes.Float32, in: `[1.5, 2.5]`, want: `[1.5, 4]`}, + {name: "float64", typ: arrow.PrimitiveTypes.Float64, in: `[1.5, 2.5]`, want: `[1.5, 4]`}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + input := cumulativeInput(t, mem, tc.typ, tc.in) + defer input.Release() + expected := cumulativeInput(t, mem, tc.typ, tc.want) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + } + + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[3]`) + defer expected.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, compute.NewDatum(int32(3))) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + + fullInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[0, 1, 2, 3]`) + defer fullInput.Release() + slicedInput := array.NewSlice(fullInput, 1, 3) + defer slicedInput.Release() + slicedExpected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 3]`) + defer slicedExpected.Release() + result, err = compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: slicedInput.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: slicedExpected.Data()}, result, nil, nil) + + t.Run("sliced nulls", func(t *testing.T) { + fullInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[9, null, 2, 3, 99]`) + defer fullInput.Release() + slicedInput := array.NewSlice(fullInput, 1, 4) + defer slicedInput.Release() + + tests := []struct { + name string + opts compute.CumulativeOptions + want string + }{ + {name: "propagate nulls", want: `[null, null, null]`}, + {name: "skip nulls", opts: compute.CumulativeOptions{SkipNulls: true}, want: `[null, 2, 5]`}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, tc.want) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, tc.opts, &compute.ArrayDatum{Value: slicedInput.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + } + }) +} + +func TestCumulativeSumNullScalarInput(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + types := []arrow.DataType{ + arrow.PrimitiveTypes.Int8, + arrow.PrimitiveTypes.Int16, + arrow.PrimitiveTypes.Int32, + arrow.PrimitiveTypes.Int64, + arrow.PrimitiveTypes.Uint8, + arrow.PrimitiveTypes.Uint16, + arrow.PrimitiveTypes.Uint32, + arrow.PrimitiveTypes.Uint64, + arrow.PrimitiveTypes.Float32, + arrow.PrimitiveTypes.Float64, + } + functions := []struct { + name string + run func(context.Context, compute.CumulativeOptions, compute.Datum) (compute.Datum, error) + }{ + {name: "unchecked", run: compute.CumulativeSum}, + {name: "checked", run: compute.CumulativeSumChecked}, + } + options := []struct { + name string + start bool + skip bool + }{ + {name: "no_start_no_skip"}, + {name: "no_start_skip", skip: true}, + {name: "start_no_skip", start: true}, + {name: "start_skip", start: true, skip: true}, + } + + for _, typ := range types { + start, err := scalar.ParseScalar(typ, "10") + require.NoError(t, err) + if releasable, ok := start.(scalar.Releasable); ok { + defer releasable.Release() + } + + for _, fn := range functions { + for _, tc := range options { + t.Run(typ.Name()+"/"+fn.name+"/"+tc.name, func(t *testing.T) { + input := compute.NewDatum(scalar.MakeNullScalar(typ)) + defer input.Release() + + opts := compute.CumulativeOptions{SkipNulls: tc.skip} + if tc.start { + opts.Start = start + } + + result, err := fn.run(ctx, opts, input) + require.NoError(t, err) + defer result.Release() + + expected := cumulativeInput(t, mem, typ, `[null]`) + defer expected.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + } + } + } +} + +func TestCumulativeSumNullsAndStart(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, null, 2, null, 3]`) + defer input.Release() + + t.Run("propagate nulls", func(t *testing.T) { + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, null, null, null, null]`) + defer expected.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + + t.Run("skip nulls", func(t *testing.T) { + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, null, 3, null, 6]`) + defer expected.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{SkipNulls: true}, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + + t.Run("start value", func(t *testing.T) { + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[11, null, 13, null, 16]`) + defer expected.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: scalar.NewInt64Scalar(10), + SkipNulls: true, + }, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + +} + +func TestCumulativeSumRejectsTypedNullStart(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1]`) + defer input.Release() + + for _, tc := range []struct { + name string + run func(context.Context, compute.CumulativeOptions, compute.Datum) (compute.Datum, error) + }{ + {name: "unchecked", run: compute.CumulativeSum}, + {name: "checked", run: compute.CumulativeSumChecked}, + } { + t.Run(tc.name, func(t *testing.T) { + result, err := tc.run(ctx, compute.CumulativeOptions{ + Start: scalar.MakeNullScalar(arrow.PrimitiveTypes.Int32), + }, &compute.ArrayDatum{Value: input.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) + }) + } +} + +func TestCumulativeSumStartSafeCast(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + tests := []struct { + name string + typ arrow.DataType + start scalar.Scalar + }{ + {name: "signed integer overflow", typ: arrow.PrimitiveTypes.Int8, start: scalar.NewInt64Scalar(128)}, + {name: "signed integer underflow", typ: arrow.PrimitiveTypes.Int8, start: scalar.NewInt64Scalar(-129)}, + {name: "unsigned integer underflow", typ: arrow.PrimitiveTypes.Uint8, start: scalar.NewInt64Scalar(-1)}, + {name: "float truncation", typ: arrow.PrimitiveTypes.Int32, start: scalar.NewFloat64Scalar(1.5)}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + input := cumulativeInput(t, mem, tc.typ, `[0]`) + defer input.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: tc.start, + }, &compute.ArrayDatum{Value: input.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) + }) + } + + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[0]`) + defer input.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: scalar.NewInt64Scalar(127), + }, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + actual := result.(*compute.ArrayDatum).MakeArray() + defer actual.Release() + assert.Equal(t, int8(127), actual.(*array.Int8).Value(0)) +} + +func TestCumulativeSumStartSafeCastDictionary(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + dictValues := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int64, `[10]`) + defer dictValues.Release() + dictIndices := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[0]`) + defer dictIndices.Release() + dictType := &arrow.DictionaryType{ + IndexType: arrow.PrimitiveTypes.Int8, + ValueType: arrow.PrimitiveTypes.Int64, + } + dict := array.NewDictionaryArray(dictType, dictIndices, dictValues) + defer dict.Release() + + start, err := scalar.GetScalar(dict, 0) + require.NoError(t, err) + if releasable, ok := start.(scalar.Releasable); ok { + defer releasable.Release() + } + + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1]`) + defer input.Release() + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[11]`) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{Start: start}, + &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) +} + +func TestCumulativeSumStartScalarConversions(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1]`) + defer input.Release() + + tests := []struct { + name string + start scalar.Scalar + want string + }{ + {name: "string", start: scalar.NewStringScalar("10"), want: `[11]`}, + {name: "boolean", start: scalar.NewBooleanScalar(true), want: `[2]`}, + } + binaryBuffer := memory.NewBufferBytes([]byte("10")) + tests = append(tests, struct { + name string + start scalar.Scalar + want string + }{name: "binary", start: scalar.NewBinaryScalar(binaryBuffer, arrow.BinaryTypes.Binary), want: `[11]`}) + binaryBuffer.Release() + for _, tc := range tests { + if releasable, ok := tc.start.(scalar.Releasable); ok { + defer releasable.Release() + } + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + expected := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, tc.want) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: tc.start, + }, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + } + + invalidStart := scalar.NewStringScalar("not a number") + defer invalidStart.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: invalidStart, + }, &compute.ArrayDatum{Value: input.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) +} + +func TestCumulativeSumDoesNotTakeStartOwnership(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + data := mem.Allocate(2) + copy(data, "10") + buffer := memory.NewBufferWithAllocator(data, mem) + start := scalar.NewBinaryScalar(buffer, arrow.BinaryTypes.Binary) + buffer.Release() + + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1]`) + defer input.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{Start: start}, + &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + result.Release() + + assert.Equal(t, "10", string(start.Data())) + start.Release() +} + +func TestCumulativeSumDecimalStarts(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + decimalType := &arrow.Decimal128Type{Precision: 10, Scale: 0} + decimalStart, err := scalar.ParseScalar(decimalType, "10") + require.NoError(t, err) + + intInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1]`) + defer intInput.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: decimalStart, + }, &compute.ArrayDatum{Value: intInput.Data()}) + require.NoError(t, err) + defer result.Release() + actual := result.(*compute.ArrayDatum).MakeArray() + defer actual.Release() + assert.Equal(t, int32(11), actual.(*array.Int32).Value(0)) + + floatInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Float64, `[1.5]`) + defer floatInput.Release() + result, err = compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: decimalStart, + }, &compute.ArrayDatum{Value: floatInput.Data()}) + require.NoError(t, err) + defer result.Release() + actual = result.(*compute.ArrayDatum).MakeArray() + defer actual.Release() + assert.Equal(t, 11.5, actual.(*array.Float64).Value(0)) + + fractionalStart, err := scalar.ParseScalar(&arrow.Decimal128Type{Precision: 10, Scale: 1}, "1.5") + require.NoError(t, err) + result, err = compute.CumulativeSum(ctx, compute.CumulativeOptions{ + Start: fractionalStart, + }, &compute.ArrayDatum{Value: intInput.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) +} + +func TestCumulativeOptionsSerialization(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + binaryBuffer := memory.NewBufferBytes([]byte("10")) + binaryStart := scalar.NewBinaryScalar(binaryBuffer, arrow.BinaryTypes.Binary) + binaryBuffer.Release() + + tests := []struct { + name string + start scalar.Scalar + }{ + {name: "nil", start: nil}, + {name: "int32", start: scalar.NewInt32Scalar(10)}, + {name: "string", start: scalar.NewStringScalar("10")}, + {name: "binary", start: binaryStart}, + {name: "typed null", start: scalar.MakeNullScalar(arrow.PrimitiveTypes.Int32)}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + expr := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")}, + &compute.CumulativeOptions{Start: tc.start, SkipNulls: true}) + + serialized, err := compute.SerializeExpr(expr, mem) + require.NoError(t, err) + roundTripped, err := compute.DeserializeExpr(mem, serialized) + serialized.Release() + require.NoError(t, err) + + assert.True(t, expr.Equals(roundTripped)) + assert.NotEmpty(t, roundTripped.String()) + roundTripped.Release() + expr.Release() + if releasable, ok := tc.start.(scalar.Releasable); ok { + releasable.Release() + } + }) + } +} + +func TestCumulativeOptionsDictionarySerialization(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + dictValues := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int64, `[10]`) + defer dictValues.Release() + dictIndices := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[0]`) + defer dictIndices.Release() + dictType := &arrow.DictionaryType{ + IndexType: arrow.PrimitiveTypes.Int8, + ValueType: arrow.PrimitiveTypes.Int64, + } + dict := array.NewDictionaryArray(dictType, dictIndices, dictValues) + defer dict.Release() + + start, err := scalar.GetScalar(dict, 0) + require.NoError(t, err) + defer start.(scalar.Releasable).Release() + expr := compute.NewCall("cumulative_sum", []compute.Expression{compute.NewFieldRef("values")}, + &compute.CumulativeOptions{Start: start}) + defer expr.Release() + + serialized, err := compute.SerializeExpr(expr, mem) + require.NoError(t, err) + defer serialized.Release() + + roundTripped, err := compute.DeserializeExpr(mem, serialized) + require.NoError(t, err) + defer roundTripped.Release() + + assert.True(t, expr.Equals(roundTripped)) +} + +func TestCumulativeSumChunked(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + first := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 2]`) + second := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[3, 4]`) + input := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{first, second}) + defer input.Release() + defer first.Release() + defer second.Release() + + expectedArray := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 3, 6, 10]`) + defer expectedArray.Release() + expected := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{expectedArray}) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ChunkedDatum{Value: input}) + require.NoError(t, err) + defer result.Release() + require.Equal(t, compute.KindChunked, result.Kind()) + assertDatumsEqual(t, &compute.ChunkedDatum{Value: expected}, result, nil, nil) + +} + +func TestCumulativeSumEmptyChunkedInput(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + input := arrow.NewChunked(arrow.PrimitiveTypes.Int32, nil) + defer input.Release() + + result, err := compute.CumulativeSum( + context.Background(), + compute.CumulativeOptions{}, + &compute.ChunkedDatum{Value: input}, + ) + require.NoError(t, err) + defer result.Release() + + chunked, ok := result.(*compute.ChunkedDatum) + require.True(t, ok) + assert.Empty(t, chunked.Value.Chunks()) + assert.Equal(t, int64(0), result.Len()) +} + +func TestCumulativeSumChunkedOutputIgnoresChunkSizeAndEmptyChunks(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + emptyBefore := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[]`) + values := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 2, 3, 4]`) + emptyAfter := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[]`) + input := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{emptyBefore, values, emptyAfter}) + defer input.Release() + defer emptyBefore.Release() + defer values.Release() + defer emptyAfter.Release() + + execCtx := compute.DefaultExecCtx() + execCtx.ChunkSize = 2 + ctx := compute.SetExecCtx(context.Background(), execCtx) + ctx = compute.WithAllocator(ctx, mem) + + expectedArray := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, 3, 6, 10]`) + defer expectedArray.Release() + expected := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{expectedArray}) + defer expected.Release() + result, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ChunkedDatum{Value: input}) + require.NoError(t, err) + defer result.Release() + require.Equal(t, compute.KindChunked, result.Kind()) + assertDatumsEqual(t, &compute.ChunkedDatum{Value: expected}, result, nil, nil) +} + +func TestCumulativeSumStateAcrossChunks(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + first := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, null]`) + second := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[2, 3]`) + input := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{first, second}) + defer input.Release() + defer first.Release() + defer second.Release() + + for _, tc := range []struct { + name string + opts compute.CumulativeOptions + expected string + }{ + {name: "propagate nulls", expected: `[1, null, null, null]`}, + {name: "skip nulls", opts: compute.CumulativeOptions{SkipNulls: true}, expected: `[1, null, 3, 6]`}, + } { + t.Run(tc.name, func(t *testing.T) { + expectedArray := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, tc.expected) + defer expectedArray.Release() + expected := arrow.NewChunked(arrow.PrimitiveTypes.Int32, []arrow.Array{expectedArray}) + defer expected.Release() + + result, err := compute.CumulativeSum(ctx, tc.opts, &compute.ChunkedDatum{Value: input}) + require.NoError(t, err) + defer result.Release() + require.Equal(t, compute.KindChunked, result.Kind()) + assertDatumsEqual(t, &compute.ChunkedDatum{Value: expected}, result, nil, nil) + }) + } +} + +func TestCumulativeSumCheckedChunkedOverflow(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + first := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[127]`) + second := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[1]`) + input := arrow.NewChunked(arrow.PrimitiveTypes.Int8, []arrow.Array{first, second}) + defer input.Release() + defer first.Release() + defer second.Release() + + ctx := compute.WithAllocator(context.Background(), mem) + result, err := compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{}, &compute.ChunkedDatum{Value: input}) + assert.Nil(t, result) + assert.ErrorIs(t, err, arrow.ErrInvalid) +} + +func TestCumulativeSumIgnoresExecutorChunkSize(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, `[1, null, 2, 3]`) + defer input.Release() + + execCtx := compute.DefaultExecCtx() + execCtx.ChunkSize = 1 + ctx := compute.WithAllocator(context.Background(), mem) + ctx = compute.SetExecCtx(ctx, execCtx) + + tests := []struct { + name string + opts compute.CumulativeOptions + expected string + }{ + {name: "propagate nulls", expected: `[1, null, null, null]`}, + {name: "skip nulls", opts: compute.CumulativeOptions{SkipNulls: true}, expected: `[1, null, 3, 6]`}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + expectedArray := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int32, tc.expected) + defer expectedArray.Release() + + result, err := compute.CumulativeSum(ctx, tc.opts, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + require.Equal(t, compute.KindArray, result.Kind()) + assertDatumsEqual(t, &compute.ArrayDatum{Value: expectedArray.Data()}, result, nil, nil) + }) + } + + overflowInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[127, 1]`) + defer overflowInput.Release() + result, err := compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: overflowInput.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) + + startOverflowInput := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[1]`) + defer startOverflowInput.Release() + result, err = compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{ + Start: scalar.NewInt8Scalar(127), + }, &compute.ArrayDatum{Value: startOverflowInput.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) +} + +func TestCumulativeSumChecked(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + input := cumulativeInput(t, mem, arrow.PrimitiveTypes.Int8, `[127, 1]`) + defer input.Release() + + unchecked, err := compute.CumulativeSum(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer unchecked.Release() + uncheckedArray := unchecked.(*compute.ArrayDatum).MakeArray() + defer uncheckedArray.Release() + assert.Equal(t, int8(-128), uncheckedArray.(*array.Int8).Value(1)) + + _, err = compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{}, &compute.ArrayDatum{Value: input.Data()}) + assert.ErrorIs(t, err, arrow.ErrInvalid) + +} + +func TestCumulativeSumCheckedIntegerOverflow(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + tests := []struct { + name string + typ arrow.DataType + values string + }{ + {name: "int8 positive", typ: arrow.PrimitiveTypes.Int8, values: `[127, 1]`}, + {name: "int8 negative", typ: arrow.PrimitiveTypes.Int8, values: `[-128, -1]`}, + {name: "int16 positive", typ: arrow.PrimitiveTypes.Int16, values: `[32767, 1]`}, + {name: "int16 negative", typ: arrow.PrimitiveTypes.Int16, values: `[-32768, -1]`}, + {name: "int32 positive", typ: arrow.PrimitiveTypes.Int32, values: `[2147483647, 1]`}, + {name: "int32 negative", typ: arrow.PrimitiveTypes.Int32, values: `[-2147483648, -1]`}, + {name: "int64 positive", typ: arrow.PrimitiveTypes.Int64, values: `[9223372036854775807, 1]`}, + {name: "int64 negative", typ: arrow.PrimitiveTypes.Int64, values: `[-9223372036854775808, -1]`}, + {name: "uint8", typ: arrow.PrimitiveTypes.Uint8, values: `[255, 1]`}, + {name: "uint16", typ: arrow.PrimitiveTypes.Uint16, values: `[65535, 1]`}, + {name: "uint32", typ: arrow.PrimitiveTypes.Uint32, values: `[4294967295, 1]`}, + {name: "uint64", typ: arrow.PrimitiveTypes.Uint64, values: `[18446744073709551615, 1]`}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + input := cumulativeInput(t, mem, tc.typ, tc.values) + defer input.Release() + + result, err := compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{}, + &compute.ArrayDatum{Value: input.Data()}) + if result != nil { + result.Release() + } + assert.ErrorIs(t, err, arrow.ErrInvalid) + }) + } +} + +func TestCumulativeSumCheckedIntegerBoundaries(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := compute.WithAllocator(context.Background(), mem) + + tests := []struct { + name string + typ arrow.DataType + input string + expected string + }{ + {name: "int8 positive", typ: arrow.PrimitiveTypes.Int8, input: `[126, 1]`, expected: `[126, 127]`}, + {name: "int8 negative", typ: arrow.PrimitiveTypes.Int8, input: `[-127, -1]`, expected: `[-127, -128]`}, + {name: "int64 negative", typ: arrow.PrimitiveTypes.Int64, input: `[-9223372036854775807, -1]`, expected: `[-9223372036854775807, -9223372036854775808]`}, + {name: "uint8", typ: arrow.PrimitiveTypes.Uint8, input: `[254, 1]`, expected: `[254, 255]`}, + {name: "uint64", typ: arrow.PrimitiveTypes.Uint64, input: `[18446744073709551614, 1]`, expected: `[18446744073709551614, 18446744073709551615]`}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + input := cumulativeInput(t, mem, tc.typ, tc.input) + defer input.Release() + expected := cumulativeInput(t, mem, tc.typ, tc.expected) + defer expected.Release() + + result, err := compute.CumulativeSumChecked(ctx, compute.CumulativeOptions{}, + &compute.ArrayDatum{Value: input.Data()}) + require.NoError(t, err) + defer result.Release() + assertDatumsEqual(t, &compute.ArrayDatum{Value: expected.Data()}, result, nil, nil) + }) + } +} diff --git a/arrow/ipc/file_reader.go b/arrow/ipc/file_reader.go index 3128a150..a7602e88 100644 --- a/arrow/ipc/file_reader.go +++ b/arrow/ipc/file_reader.go @@ -239,6 +239,7 @@ func NewMappedFileReader(data []byte, opts ...Option) (*FileReader, error) { ) if err := f.init(cfg); err != nil { + _ = f.Close() return nil, err } return &f, nil @@ -260,6 +261,7 @@ func NewFileReader(r ReadAtSeeker, opts ...Option) (*FileReader, error) { ) if err := f.init(cfg); err != nil { + _ = f.Close() return nil, err } return &f, nil @@ -330,8 +332,10 @@ func (f *FileReader) readSchema(ensureNativeEndian bool) error { if err != nil { return err } - - kind, err = readDictionary(&f.memo, msg.meta, msg.body, f.swapEndianness, f.mem) + kind, err = func() (dictutils.Kind, error) { + defer msg.Release() + return readDictionary(&f.memo, msg.meta, msg.body, f.swapEndianness, f.mem) + }() if err != nil { return err } @@ -408,6 +412,7 @@ func (f *FileReader) Close() error { f.record.Release() f.record = nil } + f.memo.Clear() return nil } diff --git a/arrow/ipc/file_reader_internal_test.go b/arrow/ipc/file_reader_internal_test.go index 93ab1c39..35ff72d4 100644 --- a/arrow/ipc/file_reader_internal_test.go +++ b/arrow/ipc/file_reader_internal_test.go @@ -18,14 +18,46 @@ package ipc import ( + "bytes" "testing" "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" "github.com/apache/arrow-go/v18/arrow/internal/dictutils" "github.com/apache/arrow-go/v18/arrow/memory" "github.com/stretchr/testify/require" ) +func TestFileReaderInitFailureReleasesDictionaries(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.NewGoAllocator()) + defer mem.AssertSize(t, 0) + + schema := arrow.NewSchema([]arrow.Field{{ + Name: "value", + Type: &arrow.DictionaryType{ + IndexType: arrow.PrimitiveTypes.Int8, + ValueType: arrow.BinaryTypes.String, + }, + }}, nil) + builder := array.NewRecordBuilder(mem, schema) + column := builder.Field(0).(*array.BinaryDictionaryBuilder) + column.Append([]byte("value")) + record := builder.NewRecordBatch() + defer record.Release() + defer builder.Release() + + var buf bytes.Buffer + writer, err := NewFileWriter(&buf, WithAllocator(mem), WithSchema(schema)) + require.NoError(t, err) + require.NoError(t, writer.Write(record)) + require.NoError(t, writer.Close()) + + wrongSchema := arrow.NewSchema([]arrow.Field{{Name: "value", Type: arrow.PrimitiveTypes.Int32}}, nil) + reader, err := NewFileReader(bytes.NewReader(buf.Bytes()), WithAllocator(mem), WithSchema(wrongSchema)) + require.Error(t, err) + require.Nil(t, reader) +} + func TestLoadRecordBatchReturnsMalformedMetadataErrors(t *testing.T) { meta := memory.NewBufferBytes([]byte{0}) defer meta.Release() diff --git a/arrow/ipc/file_writer.go b/arrow/ipc/file_writer.go index 0d970ba5..9bbd8556 100644 --- a/arrow/ipc/file_writer.go +++ b/arrow/ipc/file_writer.go @@ -285,6 +285,8 @@ func NewFileWriter(w io.Writer, opts ...Option) (*FileWriter, error) { } func (f *FileWriter) Close() error { + defer f.releaseDictionaries() + err := f.checkStarted() if err != nil { return fmt.Errorf("arrow/ipc: could not write empty file: %w", err) @@ -303,6 +305,13 @@ func (f *FileWriter) Close() error { return nil } +func (f *FileWriter) releaseDictionaries() { + for _, d := range f.lastWrittenDicts { + d.Release() + } + f.lastWrittenDicts = nil +} + func (f *FileWriter) Write(rec arrow.RecordBatch) error { schema := rec.Schema() if schema == nil || !schema.Equal(f.schema) { diff --git a/arrow/scalar/parse.go b/arrow/scalar/parse.go index 6d1e6a33..248eeb0a 100644 --- a/arrow/scalar/parse.go +++ b/arrow/scalar/parse.go @@ -48,9 +48,27 @@ type hasTypename interface { var ( hasTypenameType = reflect.TypeOf((*hasTypename)(nil)).Elem() dataTypeType = reflect.TypeOf((*arrow.DataType)(nil)).Elem() + scalarType = reflect.TypeOf((*Scalar)(nil)).Elem() ) func FromScalar(sc *Struct, val interface{}) error { + return FromScalarWithAllocator(sc, val, memory.DefaultAllocator) +} + +// FromScalarWithAllocator populates val from a Struct scalar, allocating any +// cloned scalar fields with mem. +func FromScalarWithAllocator(sc *Struct, val interface{}, mem memory.Allocator) error { + var rollbacks []func() + err := fromScalarWithAllocator(sc, val, mem, &rollbacks) + if err != nil { + for i := len(rollbacks) - 1; i >= 0; i-- { + rollbacks[i]() + } + } + return err +} + +func fromScalarWithAllocator(sc *Struct, val interface{}, mem memory.Allocator, rollbacks *[]func()) error { if sc == nil || len(sc.Value) == 0 { return nil } @@ -76,7 +94,7 @@ func FromScalar(sc *Struct, val interface{}) error { if err != nil { return err } - if err := setFromScalar(fldVal, value.Field(i)); err != nil { + if err := setFromScalar(fldVal, value.Field(i), mem, rollbacks); err != nil { return err } } @@ -84,7 +102,27 @@ func FromScalar(sc *Struct, val interface{}) error { return nil } -func setFromScalar(s Scalar, v reflect.Value) error { +func setFromScalar(s Scalar, v reflect.Value, mem memory.Allocator, rollbacks *[]func()) error { + if v.Type() == scalarType { + if !s.IsValid() && s.DataType().ID() == arrow.NULL { + v.Set(reflect.Zero(v.Type())) + return nil + } + + clone, err := cloneScalar(s, mem) + if err != nil { + return err + } + v.Set(reflect.ValueOf(clone)) + *rollbacks = append(*rollbacks, func() { + if releasable, ok := clone.(interface{ Release() }); ok { + releasable.Release() + } + v.Set(reflect.Zero(v.Type())) + }) + return nil + } + if v.Type() == dataTypeType { v.Set(reflect.ValueOf(s.DataType())) return nil @@ -110,7 +148,7 @@ func setFromScalar(s Scalar, v reflect.Value) error { case ListScalar: return fromListScalar(s, v) case *Struct: - return FromScalar(s, v.Interface()) + return fromScalarWithAllocator(s, v.Interface(), mem, rollbacks) default: if v.Type() == reflect.TypeOf(arrow.TimeUnit(0)) { v.Set(reflect.ValueOf(arrow.TimeUnit(s.value().(uint32)))) @@ -122,11 +160,17 @@ func setFromScalar(s Scalar, v reflect.Value) error { } func ToScalar(val interface{}, mem memory.Allocator) (Scalar, error) { + if val == nil { + return ScalarNull, nil + } + switch v := val.(type) { case arrow.DataType: return MakeScalar(v), nil case TypeToScalar: return v.ToScalar() + case Scalar: + return cloneScalar(v, mem) } v := reflect.Indirect(reflect.ValueOf(val)) @@ -134,6 +178,13 @@ func ToScalar(val interface{}, mem memory.Allocator) (Scalar, error) { case reflect.Struct: scalars := make([]Scalar, 0, v.Type().NumField()) fields := make([]string, 0, v.Type().NumField()) + releaseScalars := func() { + for _, s := range scalars { + if releasable, ok := s.(Releasable); ok { + releasable.Release() + } + } + } for i := 0; i < v.Type().NumField(); i++ { fld := v.Type().Field(i) tag := fld.Tag.Get("compute") @@ -144,6 +195,7 @@ func ToScalar(val interface{}, mem memory.Allocator) (Scalar, error) { fldVal := v.Field(i) s, err := ToScalar(fldVal.Interface(), mem) if err != nil { + releaseScalars() return nil, err } scalars = append(scalars, s) @@ -156,7 +208,12 @@ func ToScalar(val interface{}, mem memory.Allocator) (Scalar, error) { fields = append(fields, "_type_name") } - return NewStructScalarWithNames(scalars, fields) + result, err := NewStructScalarWithNames(scalars, fields) + if err != nil { + releaseScalars() + return nil, err + } + return result, nil case reflect.Slice: return createListScalar(v, mem) default: @@ -164,6 +221,39 @@ func ToScalar(val interface{}, mem memory.Allocator) (Scalar, error) { } } +func cloneScalar(val Scalar, mem memory.Allocator) (Scalar, error) { + if !val.IsValid() { + return MakeNullScalar(val.DataType()), nil + } + + if binary, ok := val.(BinaryScalar); ok { + data := mem.Allocate(len(binary.Data())) + copy(data, binary.Data()) + buf := memory.NewBufferWithAllocator(data, mem) + defer buf.Release() + + switch val.(type) { + case *String: + return NewStringScalarFromBuffer(buf), nil + case *LargeString: + return NewLargeStringScalarFromBuffer(buf), nil + case *LargeBinary: + return NewLargeBinaryScalar(buf), nil + case *FixedSizeBinary: + return NewFixedSizeBinaryScalar(buf, val.DataType()), nil + default: + return NewBinaryScalar(buf, val.DataType()), nil + } + } + + arr, err := MakeArrayFromScalar(val, 1, mem) + if err != nil { + return nil, err + } + defer arr.Release() + return GetScalar(arr, 0) +} + func createListScalar(sliceval reflect.Value, mem memory.Allocator) (Scalar, error) { if sliceval.Kind() != reflect.Slice { return nil, fmt.Errorf("createListScalar only works for slices, not %s", sliceval.Kind()) diff --git a/arrow/scalar/scalar.go b/arrow/scalar/scalar.go index 8cb31764..51493997 100644 --- a/arrow/scalar/scalar.go +++ b/arrow/scalar/scalar.go @@ -589,11 +589,11 @@ func GetScalar(arr arrow.Array, idx int) (Scalar, error) { switch arr := arr.(type) { case *array.Binary: - buf := memory.NewBufferBytes(arr.Value(idx)) + buf := scalarValueBuffer(arr.Data().Buffers()[2], arr.ValueOffset(idx), arr.ValueLen(idx)) defer buf.Release() return NewBinaryScalar(buf, arr.DataType()), nil case *array.LargeBinary: - buf := memory.NewBufferBytes(arr.Value(idx)) + buf := scalarValueBuffer(arr.Data().Buffers()[2], int(arr.ValueOffset(idx)), arr.ValueLen(idx)) defer buf.Release() return NewLargeBinaryScalar(buf), nil case *array.Boolean: @@ -617,7 +617,8 @@ func GetScalar(arr arrow.Array, idx int) (Scalar, error) { } return NewExtensionScalar(storage, arr.DataType()), nil case *array.FixedSizeBinary: - buf := memory.NewBufferBytes(arr.Value(idx)) + width := arr.DataType().(*arrow.FixedSizeBinaryType).ByteWidth + buf := scalarValueBuffer(arr.Data().Buffers()[1], (arr.Data().Offset()+idx)*width, width) defer buf.Release() return NewFixedSizeBinaryScalar(buf, arr.DataType()), nil case *array.FixedSizeList: @@ -752,6 +753,13 @@ func GetScalar(arr arrow.Array, idx int) (Scalar, error) { return nil, fmt.Errorf("cannot create scalar from array of type %s", arr.DataType()) } +func scalarValueBuffer(values *memory.Buffer, offset, length int) *memory.Buffer { + if values == nil { + return memory.NewBufferBytes(nil) + } + return memory.SliceBuffer(values, offset, length) +} + // MakeArrayOfNull creates an array of size length which is all null of the given data type. // // Deprecated: Use array.MakeArrayOfNull @@ -864,6 +872,26 @@ func MakeArrayFromScalar(sc Scalar, length int, mem memory.Allocator) (arrow.Arr data := finishFixedWidth(arrow.Decimal256Traits.CastToBytes([]decimal256.Num{s.Value})) defer data.Release() return array.MakeFromData(data), nil + case *Dictionary: + if err := s.Validate(); err != nil { + return nil, err + } + + indices, err := MakeArrayFromScalar(s.Value.Index, length, mem) + if err != nil { + return nil, err + } + defer indices.Release() + + // Copy the dictionary values as well as the indices so the resulting + // scalar does not retain memory owned by the source dictionary. + dict, err := array.Concatenate([]arrow.Array{s.Value.Dict}, mem) + if err != nil { + return nil, err + } + defer dict.Release() + + return array.NewDictionaryArray(s.DataType(), indices, dict), nil case PrimitiveScalar: data := finishFixedWidth(s.Data()) defer data.Release() diff --git a/arrow/scalar/scalar_test.go b/arrow/scalar/scalar_test.go index 5634864f..612fa369 100644 --- a/arrow/scalar/scalar_test.go +++ b/arrow/scalar/scalar_test.go @@ -1253,6 +1253,43 @@ func TestMakeArrayFromScalar(t *testing.T) { } } +func TestMakeArrayFromDictionaryScalar(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + dictValues, _, err := array.FromJSON(mem, arrow.PrimitiveTypes.Int64, strings.NewReader(`[10, 20]`)) + require.NoError(t, err) + defer dictValues.Release() + dictIndices, _, err := array.FromJSON(mem, arrow.PrimitiveTypes.Int8, strings.NewReader(`[0, 1]`)) + require.NoError(t, err) + defer dictIndices.Release() + dictType := &arrow.DictionaryType{ + IndexType: arrow.PrimitiveTypes.Int8, + ValueType: arrow.PrimitiveTypes.Int64, + } + dictArray := array.NewDictionaryArray(dictType, dictIndices, dictValues) + defer dictArray.Release() + + dictScalar, err := scalar.GetScalar(dictArray, 0) + require.NoError(t, err) + + result, err := scalar.MakeArrayFromScalar(dictScalar, 2, mem) + require.NoError(t, err) + defer result.Release() + + actual, err := scalar.GetScalar(result, 1) + require.NoError(t, err) + defer actual.(scalar.Releasable).Release() + assert.True(t, scalar.Equals(dictScalar, actual)) + + structScalar, err := scalar.NewStructScalarWithNames([]scalar.Scalar{dictScalar}, []string{"value"}) + require.NoError(t, err) + defer structScalar.Release() + structArray, err := scalar.MakeArrayFromScalar(structScalar, 2, mem) + require.NoError(t, err) + defer structArray.Release() +} + func TestMakeArrayFromScalarRejectsNegativeLength(t *testing.T) { mem := memory.NewCheckedAllocator(memory.NewGoAllocator()) defer mem.AssertSize(t, 0) @@ -1291,6 +1328,165 @@ type OptionValTest struct { func (OptionValTest) TypeName() string { return "OptionValTest" } +type scalarFieldOption struct { + Value scalar.Scalar `compute:"value"` +} + +func (scalarFieldOption) TypeName() string { return "scalarFieldOption" } + +type scalarFieldWithUnsupportedSlice struct { + Value scalar.Scalar `compute:"value"` + Unsupported []float64 `compute:"unsupported"` +} + +type scalarFieldWithLaterListFailure struct { + Value scalar.Scalar `compute:"value"` + Invalid bool `compute:"invalid"` +} + +type zeroingAllocator struct{} + +func (*zeroingAllocator) Allocate(size int) []byte { return make([]byte, size) } + +func (*zeroingAllocator) Reallocate(size int, old []byte) []byte { + next := make([]byte, size) + copy(next, old) + clear(old) + return next +} + +func (*zeroingAllocator) Free(buf []byte) { clear(buf) } + +func TestScalarFieldCloneOwnsBinaryValue(t *testing.T) { + mem := memory.NewCheckedAllocator(&zeroingAllocator{}) + defer mem.AssertSize(t, 0) + + original := scalar.NewStringScalar("10") + encoded, err := scalar.ToScalar(scalarFieldOption{Value: original}, mem) + require.NoError(t, err) + original.Release() + + var decoded scalarFieldOption + require.NoError(t, scalar.FromScalarWithAllocator(encoded.(*scalar.Struct), &decoded, mem)) + encoded.(*scalar.Struct).Release() + + value, ok := decoded.Value.(scalar.BinaryScalar) + require.True(t, ok) + assert.Equal(t, "10", string(value.Data())) + value.Release() +} + +func TestToScalarReleasesFieldsWhenLaterFieldFails(t *testing.T) { + mem := memory.NewCheckedAllocator(&zeroingAllocator{}) + defer mem.AssertSize(t, 0) + + data := mem.Allocate(2) + copy(data, []byte("10")) + buffer := memory.NewBufferWithAllocator(data, mem) + original := scalar.NewBinaryScalar(buffer, arrow.BinaryTypes.Binary) + buffer.Release() + + _, err := scalar.ToScalar(scalarFieldWithUnsupportedSlice{ + Value: original, + Unsupported: []float64{1}, + }, mem) + require.Error(t, err) + original.Release() +} + +func TestFromScalarWithAllocatorReleasesFieldsWhenLaterFieldFails(t *testing.T) { + mem := memory.NewCheckedAllocator(&zeroingAllocator{}) + defer mem.AssertSize(t, 0) + + data := mem.Allocate(2) + copy(data, []byte("10")) + buffer := memory.NewBufferWithAllocator(data, mem) + value := scalar.NewBinaryScalar(buffer, arrow.BinaryTypes.Binary) + buffer.Release() + + listValues, _, err := array.FromJSON(mem, arrow.PrimitiveTypes.Int32, strings.NewReader(`[1]`)) + require.NoError(t, err) + defer listValues.Release() + invalid := scalar.NewListScalar(listValues) + encoded, err := scalar.NewStructScalarWithNames( + []scalar.Scalar{value, invalid}, []string{"value", "invalid"}) + require.NoError(t, err) + defer encoded.Release() + + var decoded scalarFieldWithLaterListFailure + err = scalar.FromScalarWithAllocator(encoded, &decoded, mem) + require.Error(t, err) + assert.Nil(t, decoded.Value) +} + +func TestGetScalarBinaryValueOwnsArrayBytes(t *testing.T) { + mem := memory.NewCheckedAllocator(&zeroingAllocator{}) + defer mem.AssertSize(t, 0) + + bldr := array.NewBinaryBuilder(mem, arrow.BinaryTypes.Binary) + bldr.Append([]byte("10")) + arr := bldr.NewArray() + bldr.Release() + + value, err := scalar.GetScalar(arr, 0) + require.NoError(t, err) + arr.Release() + + binaryValue, ok := value.(scalar.BinaryScalar) + require.True(t, ok) + assert.Equal(t, "10", string(binaryValue.Data())) + binaryValue.Release() +} + +func TestGetScalarBinaryValueOwnsAllBinaryArrayBytes(t *testing.T) { + mem := memory.NewCheckedAllocator(&zeroingAllocator{}) + defer mem.AssertSize(t, 0) + + tests := []struct { + name string + build func() arrow.Array + want string + }{ + { + name: "large binary", + build: func() arrow.Array { + bldr := array.NewBinaryBuilder(mem, arrow.BinaryTypes.LargeBinary) + bldr.Append([]byte("large")) + arr := bldr.NewArray() + bldr.Release() + return arr + }, + want: "large", + }, + { + name: "fixed size binary", + build: func() arrow.Array { + typ := &arrow.FixedSizeBinaryType{ByteWidth: 5} + bldr := array.NewFixedSizeBinaryBuilder(mem, typ) + bldr.Append([]byte("fixed")) + arr := bldr.NewArray() + bldr.Release() + return arr + }, + want: "fixed", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + arr := tc.build() + value, err := scalar.GetScalar(arr, 0) + require.NoError(t, err) + arr.Release() + + binaryValue, ok := value.(scalar.BinaryScalar) + require.True(t, ok) + assert.Equal(t, tc.want, string(binaryValue.Data())) + binaryValue.Release() + }) + } +} + func TestToScalar(t *testing.T) { ot := &OptionValTest{ToType: arrow.BinaryTypes.String, Allow: true} sc, err := scalar.ToScalar(ot, memory.DefaultAllocator)