This is an automated email from the ASF dual-hosted git repository.
zeroshade pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-go.git
The following commit(s) were added to refs/heads/main by this push:
new 41e7fe57 feat(compute): add cumulative_sum and cumulative_sum_checked
(#1139)
41e7fe57 is described below
commit 41e7fe576ff0a9075c0bb018731888f90a5b98fa
Author: Minh Vu <[email protected]>
AuthorDate: Fri Aug 14 19:25:42 2026 +0200
feat(compute): add cumulative_sum and cumulative_sum_checked (#1139)
### Rationale for this change
Arrow-Go has array and chunked-array compute kernels, but no cumulative
sum operation for running totals. Adding this as a vector kernel also
gives chunked inputs the same stateful behavior as a single contiguous
array.
### What changes are included in this PR?
- Add cumulative_sum for signed integers, unsigned integers, and
floating-point arrays.
- Add cumulative_sum_checked, which reports integer overflow instead of
returning an overflowing result.
- Add CumulativeOptions with a start value and SkipNulls. By default, a
null stops the cumulative result; with SkipNulls, null positions remain
null while later valid values continue from the previous sum.
- Validate cross-type start values with the same safe numeric conversion
rules used by compute casts.
- Carry the running state across chunks and keep the output chunked.
### Are these changes tested?
- go test ./arrow/compute/... -count=1
- Added coverage for normal sums, null propagation, skipped nulls, start
values, rejected overflow and truncation in start values, chunked input,
and checked overflow at every integer width.
### Are there any user-facing changes?
Yes. This adds the public compute.CumulativeSum and
compute.CumulativeSumChecked functions plus CumulativeOptions. Existing
functions and APIs are unchanged.
---
arrow/compute/expression.go | 91 ++-
arrow/compute/expression_test.go | 154 ++++
.../compute/internal/kernels/vector_cumulative.go | 421 +++++++++++
arrow/compute/registry.go | 1 +
arrow/compute/vector_cumulative.go | 101 +++
arrow/compute/vector_cumulative_bench_test.go | 124 +++
arrow/compute/vector_cumulative_test.go | 840 +++++++++++++++++++++
arrow/ipc/file_reader.go | 9 +-
arrow/ipc/file_reader_internal_test.go | 32 +
arrow/ipc/file_writer.go | 25 +-
arrow/ipc/ipc.go | 1 +
arrow/ipc/ipc_test.go | 23 +
arrow/ipc/writer_test.go | 18 +
arrow/scalar/parse.go | 83 +-
arrow/scalar/scalar.go | 34 +-
arrow/scalar/scalar_test.go | 196 +++++
16 files changed, 2138 insertions(+), 15 deletions(-)
diff --git a/arrow/compute/expression.go b/arrow/compute/expression.go
index 2f1dc482..dcf572e7 100644
--- a/arrow/compute/expression.go
+++ b/arrow/compute/expression.go
@@ -377,7 +377,52 @@ 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 isNilScalar(lhs) || isNilScalar(rhs) {
+ return isNilScalar(lhs) && isNilScalar(rhs)
+ }
+ return scalar.Equals(lhs, rhs)
+}
+
+func isNilScalar(value scalar.Scalar) bool {
+ if value == nil {
+ return true
+ }
+
+ reflected := reflect.ValueOf(value)
+ return reflected.Kind() == reflect.Ptr && reflected.IsNil()
}
func (c *Call) Release() {
@@ -533,6 +578,7 @@ var (
funcOptsTypes = []FunctionOptions{
SetLookupOptions{}, ArithmeticOptions{}, CastOptions{},
FilterOptions{}, NullOptions{}, StrptimeOptions{},
MakeStructOptions{},
+ CumulativeOptions{},
}
)
@@ -565,9 +611,38 @@ 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 isNilScalar(value) {
+ return nil
+ }
+
+ 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 +955,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..37c58ad6 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,154 @@ 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 TestCumulativeOptionsTypedNilStart(t *testing.T) {
+ var start *scalar.Binary
+ opts := compute.CumulativeOptions{Start: start}
+
+ assert.NotPanics(t, func() { opts.Release() })
+ assert.NotPanics(t, func() {
+ expr := compute.NewCall("cumulative_sum", nil, &opts)
+ expr.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..ae2a73b7
--- /dev/null
+++ b/arrow/compute/internal/kernels/vector_cumulative.go
@@ -0,0 +1,421 @@
+// 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"
+ "reflect"
+
+ "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 isNilScalar(opts.Start) {
+ 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 isNilScalar(value scalar.Scalar) bool {
+ if value == nil {
+ return true
+ }
+
+ reflected := reflect.ValueOf(value)
+ return reflected.Kind() == reflect.Ptr && reflected.IsNil()
+}
+
+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 isNilScalar(start) {
+ 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..6616af21
--- /dev/null
+++ b/arrow/compute/vector_cumulative_test.go
@@ -0,0 +1,840 @@
+// 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 TestCumulativeSumTypedNilStart(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()
+
+ var start *scalar.Int32
+ 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 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 c9e1f00a..aba57365 100644
--- a/arrow/ipc/file_reader.go
+++ b/arrow/ipc/file_reader.go
@@ -242,6 +242,7 @@ func NewMappedFileReader(data []byte, opts ...Option)
(*FileReader, error) {
)
if err := f.init(cfg); err != nil {
+ _ = f.Close()
return nil, err
}
return &f, nil
@@ -263,6 +264,7 @@ func NewFileReader(r ReadAtSeeker, opts ...Option)
(*FileReader, error) {
)
if err := f.init(cfg); err != nil {
+ _ = f.Close()
return nil, err
}
return &f, nil
@@ -333,8 +335,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
}
@@ -411,6 +415,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..89c98999 100644
--- a/arrow/ipc/file_writer.go
+++ b/arrow/ipc/file_writer.go
@@ -246,6 +246,8 @@ type FileWriter struct {
headerStarted bool
footerWritten bool
+ closed bool
+ closeErr error
pw PayloadWriter
@@ -285,9 +287,16 @@ func NewFileWriter(w io.Writer, opts ...Option)
(*FileWriter, error) {
}
func (f *FileWriter) Close() error {
+ if f.closed {
+ return f.closeErr
+ }
+ f.closed = true
+ defer f.releaseDictionaries()
+
err := f.checkStarted()
if err != nil {
- return fmt.Errorf("arrow/ipc: could not write empty file: %w",
err)
+ f.closeErr = fmt.Errorf("arrow/ipc: could not write empty file:
%w", err)
+ return f.closeErr
}
if f.footerWritten {
@@ -296,14 +305,26 @@ func (f *FileWriter) Close() error {
err = f.pw.Close()
if err != nil {
- return fmt.Errorf("arrow/ipc: could not close payload writer:
%w", err)
+ f.closeErr = fmt.Errorf("arrow/ipc: could not close payload
writer: %w", err)
+ return f.closeErr
}
f.footerWritten = true
return nil
}
+func (f *FileWriter) releaseDictionaries() {
+ for _, d := range f.lastWrittenDicts {
+ d.Release()
+ }
+ f.lastWrittenDicts = nil
+}
+
func (f *FileWriter) Write(rec arrow.RecordBatch) error {
+ if f.closed {
+ return errFileWriterClosed
+ }
+
schema := rec.Schema()
if schema == nil || !schema.Equal(f.schema) {
return errInconsistentSchema
diff --git a/arrow/ipc/ipc.go b/arrow/ipc/ipc.go
index 744a59e7..da9cf333 100644
--- a/arrow/ipc/ipc.go
+++ b/arrow/ipc/ipc.go
@@ -29,6 +29,7 @@ const (
errNotArrowFile = errString("arrow/ipc: not an Arrow file")
errInconsistentFileMetadata = errString("arrow/ipc: file is smaller
than indicated metadata size")
errInconsistentSchema = errString("arrow/ipc: tried to write
record batch with different schema")
+ errFileWriterClosed = errString("arrow/ipc: file writer is
already closed")
errMaxRecursion = errString("arrow/ipc: max recursion depth
reached")
errBigArray = errString("arrow/ipc: array larger than
2^31-1 in length")
diff --git a/arrow/ipc/ipc_test.go b/arrow/ipc/ipc_test.go
index 895e70d8..f1f310bd 100644
--- a/arrow/ipc/ipc_test.go
+++ b/arrow/ipc/ipc_test.go
@@ -430,6 +430,29 @@ func TestDictionary(t *testing.T) {
ipcReader.Release()
}
+func TestFileWriterRejectsWriteAfterClose(t *testing.T) {
+ pool := memory.NewCheckedAllocator(memory.NewGoAllocator())
+ defer pool.AssertSize(t, 0)
+
+ schema := arrow.NewSchema([]arrow.Field{{Name: "field", Type:
&arrow.DictionaryType{
+ IndexType: arrow.PrimitiveTypes.Int8,
+ ValueType: arrow.BinaryTypes.String,
+ }}}, nil)
+ bldr := array.NewBuilder(pool, schema.Field(0).Type)
+ defer bldr.Release()
+ require.NoError(t, bldr.UnmarshalJSON([]byte(`["value"]`)))
+ arr := bldr.NewArray()
+ defer arr.Release()
+ record := array.NewRecordBatch(schema, []arrow.Array{arr}, 1)
+ defer record.Release()
+
+ writer, err := ipc.NewFileWriter(&bytes.Buffer{},
ipc.WithSchema(schema), ipc.WithAllocator(pool))
+ require.NoError(t, err)
+ require.NoError(t, writer.Write(record))
+ require.NoError(t, writer.Close())
+ require.ErrorContains(t, writer.Write(record), "already closed")
+}
+
// ARROW-18326
func TestDictionaryDeltas(t *testing.T) {
pool := memory.NewCheckedAllocator(memory.NewGoAllocator())
diff --git a/arrow/ipc/writer_test.go b/arrow/ipc/writer_test.go
index dc13032b..957d812c 100644
--- a/arrow/ipc/writer_test.go
+++ b/arrow/ipc/writer_test.go
@@ -59,10 +59,18 @@ func (failingCompressor) Type() flatbuf.CompressionType {
return flatbuf.CompressionTypeZSTD
}
+type failingWriter struct {
+ err error
+}
+
func (shortWriteWriter) Write(p []byte) (int, error) {
return len(p) - 1, io.ErrShortWrite
}
+func (w failingWriter) Write([]byte) (int, error) {
+ return 0, w.err
+}
+
func TestPayloadWriteRejectsShortWrites(t *testing.T) {
mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
defer mem.AssertSize(t, 0)
@@ -104,6 +112,16 @@ func TestWriterCloseFailureIsTerminal(t *testing.T) {
require.Equal(t, 1, payloadWriter.closeCall)
}
+func TestFileWriterCloseFailureIsTerminal(t *testing.T) {
+ schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
+ want := errors.New("write failed")
+ writer, err := NewFileWriter(failingWriter{err: want},
WithSchema(schema))
+ require.NoError(t, err)
+
+ require.ErrorIs(t, writer.Close(), want)
+ require.ErrorIs(t, writer.Close(), want)
+}
+
func TestWriterSchemaFailureIsTerminal(t *testing.T) {
schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
builder := array.NewRecordBuilder(memory.DefaultAllocator, schema)
diff --git a/arrow/scalar/parse.go b/arrow/scalar/parse.go
index 41de7717..f7616202 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
}
@@ -82,7 +100,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
}
}
@@ -90,7 +108,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
@@ -116,7 +154,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))))
@@ -128,11 +166,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))
@@ -186,6 +230,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 ade9ed7c..1d8dd13a 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:
@@ -757,6 +758,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
@@ -867,6 +875,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 666e45e5..46abb7fc 100644
--- a/arrow/scalar/scalar_test.go
+++ b/arrow/scalar/scalar_test.go
@@ -1376,6 +1376,43 @@ func TestMakeArrayFromScalarSupportsZeroLength(t
*testing.T) {
require.NoError(t, array.ValidateFull(nullArr))
}
+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)
@@ -1414,6 +1451,116 @@ 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()
+}
+
type typedNilFromScalar struct{}
func (s *typedNilFromScalar) FromStructScalar(*scalar.Struct) error {
@@ -1442,6 +1589,55 @@ type PartialScalarTest struct {
Bad []complex64
}
+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)