This is an automated email from the ASF dual-hosted git repository.
zeroshade pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/arrow.git
The following commit(s) were added to refs/heads/master by this push:
new 40ec956469 ARROW-17678: [Go] Filter kernels for Record Batches and
Tables (#14156)
40ec956469 is described below
commit 40ec95646962cccdcd62032c80e8506d4c275bc6
Author: Matt Topol <[email protected]>
AuthorDate: Tue Sep 20 10:29:32 2022 -0400
ARROW-17678: [Go] Filter kernels for Record Batches and Tables (#14156)
Authored-by: Matt Topol <[email protected]>
Signed-off-by: Matt Topol <[email protected]>
---
go/arrow/array/compare.go | 2 +-
go/arrow/array/util.go | 34 ++++
go/arrow/compute/executor.go | 8 +-
go/arrow/compute/internal/exec/utils.go | 55 ++++++
go/arrow/compute/internal/exec/utils_test.go | 109 +++++++++++
go/arrow/compute/selection.go | 182 +++++++++++++++++-
go/arrow/compute/vector_selection_test.go | 273 +++++++++++++++++++++++++++
go/arrow/table.go | 10 +-
8 files changed, 665 insertions(+), 8 deletions(-)
diff --git a/go/arrow/array/compare.go b/go/arrow/array/compare.go
index 78075cd0f4..828ed477c7 100644
--- a/go/arrow/array/compare.go
+++ b/go/arrow/array/compare.go
@@ -117,7 +117,7 @@ func ChunkedEqual(left, right *arrow.Chunked) bool {
return false
}
- var isequal bool
+ var isequal bool = true
chunkedBinaryApply(left, right, func(left arrow.Array, lbeg, lend
int64, right arrow.Array, rbeg, rend int64) bool {
isequal = SliceEqual(left, lbeg, lend, right, rbeg, rend)
return isequal
diff --git a/go/arrow/array/util.go b/go/arrow/array/util.go
index 53634e4ec5..a40b3e0ba5 100644
--- a/go/arrow/array/util.go
+++ b/go/arrow/array/util.go
@@ -237,6 +237,19 @@ func RecordToJSON(rec arrow.Record, w io.Writer) error {
return nil
}
+func TableFromJSON(mem memory.Allocator, sc *arrow.Schema, recJSON []string,
opt ...FromJSONOption) (arrow.Table, error) {
+ batches := make([]arrow.Record, len(recJSON))
+ for i, batchJSON := range recJSON {
+ batch, _, err := RecordFromJSON(mem, sc,
strings.NewReader(batchJSON), opt...)
+ if err != nil {
+ return nil, err
+ }
+ defer batch.Release()
+ batches[i] = batch
+ }
+ return NewTableFromRecords(sc, batches), nil
+}
+
func getDictArrayData(mem memory.Allocator, valueType arrow.DataType,
memoTable hashing.MemoTable, startOffset int) (*Data, error) {
dictLen := memoTable.Size() - startOffset
buffers := []*memory.Buffer{nil, nil}
@@ -300,6 +313,27 @@ func DictArrayFromJSON(mem memory.Allocator, dt
*arrow.DictionaryType, indicesJS
return NewDictionaryArray(dt, indices, dict), nil
}
+func ChunkedFromJSON(mem memory.Allocator, dt arrow.DataType, chunkStrs
[]string, opts ...FromJSONOption) (*arrow.Chunked, error) {
+ chunks := make([]arrow.Array, len(chunkStrs))
+ defer func() {
+ for _, c := range chunks {
+ if c != nil {
+ c.Release()
+ }
+ }
+ }()
+
+ var err error
+ for i, c := range chunkStrs {
+ chunks[i], _, err = FromJSON(mem, dt, strings.NewReader(c),
opts...)
+ if err != nil {
+ return nil, err
+ }
+ }
+
+ return arrow.NewChunked(dt, chunks), nil
+}
+
func getMaxBufferLen(dt arrow.DataType, length int) int {
bufferLen := int(bitutil.BytesForBits(int64(length)))
diff --git a/go/arrow/compute/executor.go b/go/arrow/compute/executor.go
index c20bd9e468..11340b295a 100644
--- a/go/arrow/compute/executor.go
+++ b/go/arrow/compute/executor.go
@@ -970,7 +970,13 @@ func (v *vectorExecutor) WrapResults(ctx context.Context,
out <-chan Datum, hasC
)
toChunked := func() {
- acc = output.(ArrayLikeDatum).Chunks()
+ out := output.(ArrayLikeDatum).Chunks()
+ acc = make([]arrow.Array, 0, len(out))
+ for _, o := range out {
+ if o.Len() > 0 {
+ acc = append(acc, o)
+ }
+ }
if output.Kind() != KindChunked {
output.Release()
}
diff --git a/go/arrow/compute/internal/exec/utils.go
b/go/arrow/compute/internal/exec/utils.go
index 9192b85d4a..0d6a21ef02 100644
--- a/go/arrow/compute/internal/exec/utils.go
+++ b/go/arrow/compute/internal/exec/utils.go
@@ -18,6 +18,7 @@ package exec
import (
"fmt"
+ "math"
"reflect"
"unsafe"
@@ -174,3 +175,57 @@ func ArrayFromSlice[T NumericTypes](mem memory.Allocator,
data []T) arrow.Array
bldr.AppendValues(data, nil)
return bldr.NewArray()
}
+
+func RechunkArraysConsistently(groups [][]arrow.Array) [][]arrow.Array {
+ if len(groups) <= 1 {
+ return groups
+ }
+
+ var totalLen int
+ for _, a := range groups[0] {
+ totalLen += a.Len()
+ }
+
+ if totalLen == 0 {
+ return groups
+ }
+
+ rechunked := make([][]arrow.Array, len(groups))
+ offsets := make([]int, len(groups))
+ // scan all array vectors at once, rechunking along the way
+ var start int64
+ for start < int64(totalLen) {
+ // first compute max possible length for next chunk
+ chunkLength := math.MaxInt64
+ for i, g := range groups {
+ offset := offsets[i]
+ // skip any done arrays including 0-length
+ for offset == g[0].Len() {
+ g = g[1:]
+ offset = 0
+ }
+ arr := g[0]
+ chunkLength = Min(chunkLength, arr.Len()-offset)
+
+ offsets[i] = offset
+ groups[i] = g
+ }
+
+ // now slice all the arrays along this chunk size
+ for i, g := range groups {
+ offset := offsets[i]
+ arr := g[0]
+ if offset == 0 && arr.Len() == chunkLength {
+ // slice spans entire array
+ arr.Retain()
+ rechunked[i] = append(rechunked[i], arr)
+ } else {
+ rechunked[i] = append(rechunked[i],
array.NewSlice(arr, int64(offset), int64(offset+chunkLength)))
+ }
+ offsets[i] += chunkLength
+ }
+
+ start += int64(chunkLength)
+ }
+ return rechunked
+}
diff --git a/go/arrow/compute/internal/exec/utils_test.go
b/go/arrow/compute/internal/exec/utils_test.go
new file mode 100644
index 0000000000..1917429f7d
--- /dev/null
+++ b/go/arrow/compute/internal/exec/utils_test.go
@@ -0,0 +1,109 @@
+// 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.
+
+package exec_test
+
+import (
+ "testing"
+
+ "github.com/apache/arrow/go/v10/arrow"
+ "github.com/apache/arrow/go/v10/arrow/array"
+ "github.com/apache/arrow/go/v10/arrow/compute/internal/exec"
+ "github.com/apache/arrow/go/v10/arrow/memory"
+ "github.com/stretchr/testify/assert"
+)
+
+func TestRechunkConsistentArraysTrivial(t *testing.T) {
+ var groups [][]arrow.Array
+ rechunked := exec.RechunkArraysConsistently(groups)
+ assert.Zero(t, rechunked)
+
+ mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+ defer mem.AssertSize(t, 0)
+
+ a1 := exec.ArrayFromSlice(mem, []int16{})
+ defer a1.Release()
+ a2 := exec.ArrayFromSlice(mem, []int16{})
+ defer a2.Release()
+ b1 := exec.ArrayFromSlice(mem, []int32{})
+ defer b1.Release()
+ groups = [][]arrow.Array{{a1, a2}, {}, {b1}}
+ rechunked = exec.RechunkArraysConsistently(groups)
+ assert.Len(t, rechunked, 3)
+
+ for _, arrvec := range rechunked {
+ for _, arr := range arrvec {
+ assert.Zero(t, arr.Len())
+ }
+ }
+}
+
+func assertEqual[T exec.NumericTypes](t *testing.T, mem memory.Allocator, arr
arrow.Array, data []T) {
+ exp := exec.ArrayFromSlice(mem, data)
+ defer exp.Release()
+ assert.Truef(t, array.Equal(exp, arr), "expected: %s\ngot: %s", exp,
arr)
+}
+
+func TestRechunkArraysConsistentlyPlain(t *testing.T) {
+ mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+ defer mem.AssertSize(t, 0)
+
+ a1 := exec.ArrayFromSlice(mem, []int16{1, 2, 3})
+ defer a1.Release()
+ a2 := exec.ArrayFromSlice(mem, []int16{4, 5})
+ defer a2.Release()
+ a3 := exec.ArrayFromSlice(mem, []int16{6, 7, 8, 9})
+ defer a3.Release()
+
+ b1 := exec.ArrayFromSlice(mem, []int32{41, 42})
+ defer b1.Release()
+ b2 := exec.ArrayFromSlice(mem, []int32{43, 44, 45})
+ defer b2.Release()
+ b3 := exec.ArrayFromSlice(mem, []int32{46, 47})
+ defer b3.Release()
+ b4 := exec.ArrayFromSlice(mem, []int32{48, 49})
+ defer b4.Release()
+
+ groups := [][]arrow.Array{{a1, a2, a3}, {b1, b2, b3, b4}}
+ rechunked := exec.RechunkArraysConsistently(groups)
+ assert.Len(t, rechunked, 2)
+ ra := rechunked[0]
+ rb := rechunked[1]
+
+ assert.Len(t, ra, 5)
+ assertEqual(t, mem, ra[0], []int16{1, 2})
+ ra[0].Release()
+ assertEqual(t, mem, ra[1], []int16{3})
+ ra[1].Release()
+ assertEqual(t, mem, ra[2], []int16{4, 5})
+ ra[2].Release()
+ assertEqual(t, mem, ra[3], []int16{6, 7})
+ ra[3].Release()
+ assertEqual(t, mem, ra[4], []int16{8, 9})
+ ra[4].Release()
+
+ assert.Len(t, rb, 5)
+ assertEqual(t, mem, rb[0], []int32{41, 42})
+ rb[0].Release()
+ assertEqual(t, mem, rb[1], []int32{43})
+ rb[1].Release()
+ assertEqual(t, mem, rb[2], []int32{44, 45})
+ rb[2].Release()
+ assertEqual(t, mem, rb[3], []int32{46, 47})
+ rb[3].Release()
+ assertEqual(t, mem, rb[4], []int32{48, 49})
+ rb[4].Release()
+}
diff --git a/go/arrow/compute/selection.go b/go/arrow/compute/selection.go
index ac5b4e4c65..fd7d941184 100644
--- a/go/arrow/compute/selection.go
+++ b/go/arrow/compute/selection.go
@@ -32,15 +32,39 @@ var (
filterMetaFunc = NewMetaFunction("filter", Binary(), filterDoc,
func(ctx context.Context, opts FunctionOptions, args ...Datum)
(Datum, error) {
if args[1].(ArrayLikeDatum).Type().ID() != arrow.BOOL {
- return nil, fmt.Errorf("%w: fitler argument
must be boolean type",
+ return nil, fmt.Errorf("%w: filter argument
must be boolean type",
arrow.ErrNotImplemented)
}
switch args[0].Kind() {
case KindRecord:
- return nil, fmt.Errorf("%w: record batch
filtering", arrow.ErrNotImplemented)
+ filtOpts, ok := opts.(*FilterOptions)
+ if !ok {
+ return nil, fmt.Errorf("%w: invalid
options type", arrow.ErrInvalid)
+ }
+
+ if filter, ok := args[1].(*ArrayDatum); ok {
+ filterArr := filter.MakeArray()
+ defer filterArr.Release()
+ rec, err := FilterRecordBatch(ctx,
args[0].(*RecordDatum).Value, filterArr, filtOpts)
+ if err != nil {
+ return nil, err
+ }
+ return &RecordDatum{Value: rec}, nil
+ }
+ return nil, fmt.Errorf("%w: record batch
filtering only implemented for Array filter", arrow.ErrNotImplemented)
case KindTable:
- return nil, fmt.Errorf("%w: table filtering",
arrow.ErrNotImplemented)
+ filtOpts, ok := opts.(*FilterOptions)
+ if !ok {
+ return nil, fmt.Errorf("%w: invalid
options type", arrow.ErrInvalid)
+ }
+
+ tbl, err := FilterTable(ctx,
args[0].(*TableDatum).Value, args[1], filtOpts)
+ if err != nil {
+ return nil, err
+ }
+ return &TableDatum{Value: tbl}, nil
+
default:
return CallFunction(ctx, "array_filter", opts,
args...)
}
@@ -343,3 +367,155 @@ func FilterArray(ctx context.Context, values, filter
arrow.Array, options Filter
defer outDatum.Release()
return outDatum.(*ArrayDatum).MakeArray(), nil
}
+
+func FilterRecordBatch(ctx context.Context, batch arrow.Record, filter
arrow.Array, opts *FilterOptions) (arrow.Record, error) {
+ if batch.NumRows() != int64(filter.Len()) {
+ return nil, fmt.Errorf("%w: filter inputs must all be the same
length", arrow.ErrInvalid)
+ }
+
+ var filterSpan exec.ArraySpan
+ filterSpan.SetMembers(filter.Data())
+
+ indices, err := kernels.GetTakeIndices(exec.GetAllocator(ctx),
&filterSpan, opts.NullSelection)
+ if err != nil {
+ return nil, err
+ }
+ defer indices.Release()
+
+ indicesArr := array.MakeFromData(indices)
+ defer indicesArr.Release()
+
+ cols := make([]arrow.Array, batch.NumCols())
+ defer func() {
+ for _, c := range cols {
+ if c != nil {
+ c.Release()
+ }
+ }
+ }()
+ eg, cctx := errgroup.WithContext(ctx)
+ eg.SetLimit(GetExecCtx(ctx).NumParallel)
+ for i, col := range batch.Columns() {
+ i, col := i, col
+ eg.Go(func() error {
+ out, err := TakeArrayOpts(cctx, col, indicesArr,
kernels.TakeOptions{BoundsCheck: false})
+ if err != nil {
+ return err
+ }
+ cols[i] = out
+ return nil
+ })
+ }
+
+ if err := eg.Wait(); err != nil {
+ return nil, err
+ }
+
+ return array.NewRecord(batch.Schema(), cols, int64(indicesArr.Len())),
nil
+}
+
+func FilterTable(ctx context.Context, tbl arrow.Table, filter Datum, opts
*FilterOptions) (arrow.Table, error) {
+ if tbl.NumRows() != filter.Len() {
+ return nil, fmt.Errorf("%w: filter inputs must all be the same
length", arrow.ErrInvalid)
+ }
+
+ if tbl.NumRows() == 0 {
+ cols := make([]arrow.Column, tbl.NumCols())
+ for i := 0; i < int(tbl.NumCols()); i++ {
+ cols[i] = *tbl.Column(i)
+ }
+ return array.NewTable(tbl.Schema(), cols, 0), nil
+ }
+
+ // last input element will be the filter array
+ nCols := tbl.NumCols()
+ inputs := make([][]arrow.Array, nCols+1)
+ for i := int64(0); i < nCols; i++ {
+ inputs[i] = tbl.Column(int(i)).Data().Chunks()
+ }
+
+ switch ft := filter.(type) {
+ case *ArrayDatum:
+ inputs[nCols] = ft.Chunks()
+ defer inputs[nCols][0].Release()
+ case *ChunkedDatum:
+ inputs[nCols] = ft.Chunks()
+ default:
+ return nil, fmt.Errorf("%w: filter should be array-like",
arrow.ErrNotImplemented)
+ }
+
+ // rechunk inputs to allow consistent iteration over the respective
chunks
+ inputs = exec.RechunkArraysConsistently(inputs)
+
+ // instead of filtering each column with the boolean filter
+ // (which would be slow if the table has a large number of columns)
+ // convert each filter chunk to indices and take() the column
+ mem := GetAllocator(ctx)
+ outCols := make([][]arrow.Array, nCols)
+ // pre-size the output
+ nChunks := len(inputs[nCols])
+ for i := range outCols {
+ outCols[i] = make([]arrow.Array, nChunks)
+ }
+ var outNumRows int64
+ var cancel context.CancelFunc
+ ctx, cancel = context.WithCancel(ctx)
+ defer cancel()
+
+ eg, cctx := errgroup.WithContext(ctx)
+ eg.SetLimit(GetExecCtx(cctx).NumParallel)
+
+ var filterSpan exec.ArraySpan
+ for i, filterChunk := range inputs[nCols] {
+ filterSpan.SetMembers(filterChunk.Data())
+ indices, err := kernels.GetTakeIndices(mem, &filterSpan,
opts.NullSelection)
+ if err != nil {
+ return nil, err
+ }
+ defer indices.Release()
+ filterChunk.Release()
+ if indices.Len() == 0 {
+ for col := int64(0); col < nCols; col++ {
+ inputs[col][i].Release()
+ }
+ continue
+ }
+
+ // take from all input columns
+ outNumRows += int64(indices.Len())
+ indicesDatum := NewDatum(indices)
+ defer indicesDatum.Release()
+
+ for col := int64(0); col < nCols; col++ {
+ columnChunk := inputs[col][i]
+ defer columnChunk.Release()
+ i := i
+ col := col
+ eg.Go(func() error {
+ columnDatum := NewDatum(columnChunk)
+ defer columnDatum.Release()
+ out, err := Take(cctx,
kernels.TakeOptions{BoundsCheck: false}, columnDatum, indicesDatum)
+ if err != nil {
+ return err
+ }
+ defer out.Release()
+ outCols[col][i] = out.(*ArrayDatum).MakeArray()
+ return nil
+ })
+ }
+ }
+
+ if err := eg.Wait(); err != nil {
+ return nil, err
+ }
+
+ outChunks := make([]arrow.Column, nCols)
+ for i, chunks := range outCols {
+ chk := arrow.NewChunked(tbl.Column(i).DataType(), chunks)
+ outChunks[i] = *arrow.NewColumn(tbl.Schema().Field(i), chk)
+ defer outChunks[i].Release()
+ chk.Release()
+ }
+
+ return array.NewTable(tbl.Schema(), outChunks, outNumRows), nil
+}
diff --git a/go/arrow/compute/vector_selection_test.go
b/go/arrow/compute/vector_selection_test.go
index 6b01a11b74..31ca6b6e64 100644
--- a/go/arrow/compute/vector_selection_test.go
+++ b/go/arrow/compute/vector_selection_test.go
@@ -814,6 +814,276 @@ func (f *FilterKernelWithStruct) TestStruct() {
f.assertFilterJSON(dt, structJSON, `[true, false, true, false]`,
`[null, {"a": 2, "b": "hello"}]`)
}
+type FilterKernelWithRecordBatch struct {
+ FilterKernelTestSuite
+}
+
+func (f *FilterKernelWithRecordBatch) doFilter(sc *arrow.Schema, batchJSON,
selection string, opts compute.FilterOptions) (arrow.Record, error) {
+ rec, _, err := array.RecordFromJSON(f.mem, sc,
strings.NewReader(batchJSON), array.WithUseNumber())
+ if err != nil {
+ return nil, err
+ }
+ defer rec.Release()
+
+ batch := compute.NewDatum(rec)
+ defer batch.Release()
+
+ filter, _, _ := array.FromJSON(f.mem, arrow.FixedWidthTypes.Boolean,
strings.NewReader(selection))
+ defer filter.Release()
+ filterDatum := compute.NewDatum(filter)
+ defer filterDatum.Release()
+
+ outDatum, err := compute.Filter(context.TODO(), batch, filterDatum,
opts)
+ if err != nil {
+ return nil, err
+ }
+
+ return outDatum.(*compute.RecordDatum).Value, nil
+}
+
+func (f *FilterKernelWithRecordBatch) assertFilter(sc *arrow.Schema,
batchJSON, selection string, opts compute.FilterOptions, expectedBatch string) {
+ actual, err := f.doFilter(sc, batchJSON, selection, opts)
+ f.Require().NoError(err)
+ defer actual.Release()
+
+ expected, _, err := array.RecordFromJSON(f.mem, sc,
strings.NewReader(expectedBatch), array.WithUseNumber())
+ f.Require().NoError(err)
+ defer expected.Release()
+
+ f.Truef(array.RecordEqual(expected, actual), "expected: %s\ngot: %s",
expected, actual)
+}
+
+func (f *FilterKernelWithRecordBatch) TestFilterRecord() {
+ fields := []arrow.Field{
+ {Name: "a", Type: arrow.PrimitiveTypes.Int32, Nullable: true},
+ {Name: "b", Type: arrow.BinaryTypes.String, Nullable: true},
+ }
+ sc := arrow.NewSchema(fields, nil)
+
+ batchJSON := `[
+ {"a": null, "b": "yo"},
+ {"a": 1, "b": ""},
+ {"a": 2, "b": "hello"},
+ {"a": 4, "b": "eh"}
+ ]`
+
+ for _, opts := range []compute.FilterOptions{f.emitNulls, f.dropOpts} {
+ f.assertFilter(sc, batchJSON, `[false, false, false, false]`,
opts, `[]`)
+ f.assertFilter(sc, batchJSON, `[true, true, true, true]`, opts,
batchJSON)
+ f.assertFilter(sc, batchJSON, `[true, false, true, false]`,
opts, `[
+ {"a": null, "b": "yo"},
+ {"a": 2, "b": "hello"}
+ ]`)
+ }
+
+ f.assertFilter(sc, batchJSON, `[false, true, true, null]`, f.dropOpts,
`[
+ {"a": 1, "b": ""},
+ {"a": 2, "b": "hello"}
+ ]`)
+
+ f.assertFilter(sc, batchJSON, `[false, true, true, null]`, f.emitNulls,
`[
+ {"a": 1, "b": ""},
+ {"a": 2, "b": "hello"},
+ {"a": null, "b": null}
+ ]`)
+}
+
+type FilterKernelWithChunked struct {
+ FilterKernelTestSuite
+}
+
+func (f *FilterKernelWithChunked) filterWithArray(dt arrow.DataType, values
[]string, filterStr string) (*arrow.Chunked, error) {
+ chk, err := array.ChunkedFromJSON(f.mem, dt, values)
+ f.Require().NoError(err)
+ defer chk.Release()
+
+ input := compute.NewDatum(chk)
+ defer input.Release()
+
+ filter, _, _ := array.FromJSON(f.mem, arrow.FixedWidthTypes.Boolean,
strings.NewReader(filterStr))
+ defer filter.Release()
+
+ filterDatum := compute.NewDatum(filter)
+ defer filterDatum.Release()
+
+ out, err := compute.Filter(context.TODO(), input, filterDatum,
*compute.DefaultFilterOptions())
+ if err != nil {
+ return nil, err
+ }
+ return out.(*compute.ChunkedDatum).Value, nil
+}
+
+func (f *FilterKernelWithChunked) filterWithChunked(dt arrow.DataType, values,
filter []string) (*arrow.Chunked, error) {
+ chk, err := array.ChunkedFromJSON(f.mem, dt, values)
+ f.Require().NoError(err)
+ defer chk.Release()
+
+ input := compute.NewDatum(chk)
+ defer input.Release()
+
+ filtChk, err := array.ChunkedFromJSON(f.mem,
arrow.FixedWidthTypes.Boolean, filter)
+ f.Require().NoError(err)
+ defer filtChk.Release()
+
+ filtDatum := compute.NewDatum(filtChk)
+ defer filtDatum.Release()
+
+ out, err := compute.Filter(context.TODO(), input, filtDatum,
*compute.DefaultFilterOptions())
+ if err != nil {
+ return nil, err
+ }
+ return out.(*compute.ChunkedDatum).Value, nil
+}
+
+func (f *FilterKernelWithChunked) assertFilter(dt arrow.DataType, values
[]string, filter string, expected []string) {
+ actual, err := f.filterWithArray(dt, values, filter)
+ f.Require().NoError(err)
+ defer actual.Release()
+
+ expectedResult, _ := array.ChunkedFromJSON(f.mem, dt, expected)
+ defer expectedResult.Release()
+ if !f.True(array.ChunkedEqual(expectedResult, actual)) {
+ var s strings.Builder
+ s.WriteString("expected: \n")
+ for _, c := range expectedResult.Chunks() {
+ fmt.Fprintf(&s, "%s\n", c)
+ }
+ s.WriteString("actual: \n")
+ for _, c := range actual.Chunks() {
+ fmt.Fprintf(&s, "%s\n", c)
+ }
+ f.T().Log(s.String())
+ }
+}
+
+func (f *FilterKernelWithChunked) assertChunkedFilter(dt arrow.DataType,
values, filter, expected []string) {
+ actual, err := f.filterWithChunked(dt, values, filter)
+ f.Require().NoError(err)
+ defer actual.Release()
+
+ expectedResult, _ := array.ChunkedFromJSON(f.mem, dt, expected)
+ defer expectedResult.Release()
+ if !f.True(array.ChunkedEqual(expectedResult, actual)) {
+ var s strings.Builder
+ s.WriteString("expected: \n")
+ for _, c := range expectedResult.Chunks() {
+ fmt.Fprintf(&s, "%s\n", c)
+ }
+ s.WriteString("actual: \n")
+ for _, c := range actual.Chunks() {
+ fmt.Fprintf(&s, "%s\n", c)
+ }
+ f.T().Log(s.String())
+ }
+}
+
+func (f *FilterKernelWithChunked) TestFilterChunked() {
+ f.assertFilter(arrow.PrimitiveTypes.Int8, []string{`[]`}, `[]`,
[]string{})
+ f.assertChunkedFilter(arrow.PrimitiveTypes.Int8, []string{`[]`},
[]string{`[]`}, []string{})
+
+ f.assertFilter(arrow.PrimitiveTypes.Int8, []string{`[7]`, `[8, 9]`},
`[false, true, false]`, []string{`[8]`})
+ f.assertChunkedFilter(arrow.PrimitiveTypes.Int8, []string{`[7]`, `[8,
9]`}, []string{`[false]`, `[true, false]`}, []string{`[8]`})
+ f.assertChunkedFilter(arrow.PrimitiveTypes.Int8, []string{`[7]`, `[8,
9]`}, []string{`[false, true]`, `[false]`}, []string{`[8]`})
+
+ _, err := f.filterWithArray(arrow.PrimitiveTypes.Int8, []string{`[7]`,
`[8, 9]`}, `[false, true, false, true, true]`)
+ f.ErrorIs(err, arrow.ErrInvalid)
+ _, err = f.filterWithChunked(arrow.PrimitiveTypes.Int8, []string{`[7]`,
`[8, 9]`}, []string{`[ false, true, false]`, `[true, true]`})
+ f.ErrorIs(err, arrow.ErrInvalid)
+}
+
+type FilterKernelWithTable struct {
+ FilterKernelTestSuite
+}
+
+func (f *FilterKernelWithTable) filterWithArray(sc *arrow.Schema, values
[]string, filter string, opts compute.FilterOptions) (arrow.Table, error) {
+ tbl, err := array.TableFromJSON(f.mem, sc, values)
+ if err != nil {
+ return nil, err
+ }
+ defer tbl.Release()
+
+ filterArr, _, _ := array.FromJSON(f.mem, arrow.FixedWidthTypes.Boolean,
strings.NewReader(filter))
+ defer filterArr.Release()
+
+ out, err := compute.Filter(context.TODO(), &compute.TableDatum{Value:
tbl}, &compute.ArrayDatum{Value: filterArr.Data()}, opts)
+ if err != nil {
+ return nil, err
+ }
+ return out.(*compute.TableDatum).Value, nil
+}
+
+func (f *FilterKernelWithTable) filterWithChunked(sc *arrow.Schema, values,
filter []string, opts compute.FilterOptions) (arrow.Table, error) {
+ tbl, err := array.TableFromJSON(f.mem, sc, values)
+ if err != nil {
+ return nil, err
+ }
+ defer tbl.Release()
+
+ filtChk, err := array.ChunkedFromJSON(f.mem,
arrow.FixedWidthTypes.Boolean, filter)
+ f.Require().NoError(err)
+ defer filtChk.Release()
+
+ out, err := compute.Filter(context.TODO(), &compute.TableDatum{Value:
tbl}, &compute.ChunkedDatum{Value: filtChk}, opts)
+ if err != nil {
+ return nil, err
+ }
+ return out.(*compute.TableDatum).Value, nil
+}
+
+func (f *FilterKernelWithTable) assertChunkedFilter(sc *arrow.Schema,
tableJSON, filter []string, opts compute.FilterOptions, expTable []string) {
+ actual, err := f.filterWithChunked(sc, tableJSON, filter, opts)
+ f.Require().NoError(err)
+ defer actual.Release()
+
+ expected, err := array.TableFromJSON(f.mem, sc, expTable)
+ f.Require().NoError(err)
+ defer expected.Release()
+
+ f.Truef(array.TableEqual(expected, actual), "expected: %s\ngot: %s",
expected, actual)
+}
+
+func (f *FilterKernelWithTable) assertFilter(sc *arrow.Schema, tableJSON
[]string, filter string, opts compute.FilterOptions, expectedTable []string) {
+ actual, err := f.filterWithArray(sc, tableJSON, filter, opts)
+ f.Require().NoError(err)
+ defer actual.Release()
+
+ expected, err := array.TableFromJSON(f.mem, sc, expectedTable)
+ f.Require().NoError(err)
+ defer expected.Release()
+
+ f.Truef(array.TableEqual(expected, actual), "expected: %s\ngot: %s",
expected, actual)
+}
+
+func (f *FilterKernelWithTable) TestFilterTable() {
+ fields := []arrow.Field{
+ {Name: "a", Type: arrow.PrimitiveTypes.Int32, Nullable: true},
+ {Name: "b", Type: arrow.BinaryTypes.String, Nullable: true},
+ }
+ sc := arrow.NewSchema(fields, nil)
+ tableJSON := []string{`[
+ {"a": null, "b": "yo"},
+ {"a": 1, "b": ""}
+ ]`, `[
+ {"a": 2, "b": "hello"},
+ {"a": 4, "b": "eh"}
+ ]`}
+
+ for _, opt := range []compute.FilterOptions{f.emitNulls, f.dropOpts} {
+ f.assertFilter(sc, tableJSON, `[false, false, false, false]`,
opt, []string{})
+ f.assertChunkedFilter(sc, tableJSON, []string{`[false]`,
`[false, false, false]`}, opt, []string{})
+ f.assertFilter(sc, tableJSON, `[true, true, true, true]`, opt,
tableJSON)
+ f.assertChunkedFilter(sc, tableJSON, []string{`[true]`, `[true,
true, true]`}, opt, tableJSON)
+ }
+
+ expectedEmitNull := []string{`[{"a": 1, "b": ""}]`, `[{"a": 2, "b":
"hello"},{"a": null, "b": null}]`}
+ f.assertFilter(sc, tableJSON, `[false, true, true, null]`, f.emitNulls,
expectedEmitNull)
+ f.assertChunkedFilter(sc, tableJSON, []string{`[false, true, true]`,
`[null]`}, f.emitNulls, expectedEmitNull)
+
+ expectedDrop := []string{`[{"a": 1, "b": ""}]`, `[{"a": 2, "b":
"hello"}]`}
+ f.assertFilter(sc, tableJSON, `[false, true, true, null]`, f.dropOpts,
expectedDrop)
+ f.assertChunkedFilter(sc, tableJSON, []string{`[false, true, true]`,
`[null]`}, f.dropOpts, expectedDrop)
+}
+
type TakeKernelTestTyped struct {
TakeKernelTestSuite
@@ -1101,4 +1371,7 @@ func TestFilterKernels(t *testing.T) {
suite.Run(t, new(FilterKernelWithUnion))
suite.Run(t, new(FilterKernelExtension))
suite.Run(t, new(FilterKernelWithStruct))
+ suite.Run(t, new(FilterKernelWithRecordBatch))
+ suite.Run(t, new(FilterKernelWithChunked))
+ suite.Run(t, new(FilterKernelWithTable))
}
diff --git a/go/arrow/table.go b/go/arrow/table.go
index c4a6351cce..e0c2caf515 100644
--- a/go/arrow/table.go
+++ b/go/arrow/table.go
@@ -140,16 +140,20 @@ type Chunked struct {
// NewChunked panics if the chunks do not have the same data type.
func NewChunked(dtype DataType, chunks []Array) *Chunked {
arr := &Chunked{
- chunks: make([]Array, len(chunks)),
+ chunks: make([]Array, 0, len(chunks)),
refCount: 1,
dtype: dtype,
}
- for i, chunk := range chunks {
+ for _, chunk := range chunks {
+ if chunk == nil {
+ continue
+ }
+
if !TypeEqual(chunk.DataType(), dtype) {
panic("arrow/array: mismatch data type")
}
chunk.Retain()
- arr.chunks[i] = chunk
+ arr.chunks = append(arr.chunks, chunk)
arr.length += chunk.Len()
arr.nulls += chunk.NullN()
}