From 36599a8ff2fd8da3c29ff022d9ba7c1a90f8294f Mon Sep 17 00:00:00 2001 From: serramatutu Date: Thu, 30 Apr 2026 12:43:09 +0200 Subject: [PATCH 01/18] Separate public API from impl of comparison functions This commit separates the actual implementation from the public `*Equal` functions. Now, all the public API does is convert the `opts ...EqualOption` into an `opt equalOption` struct and pass it into the implementation. This is useful to avoid having to reconstruct back and forth between the two in nested call stacks. All the implementations care about is `equalOption`, and `EqualOption` remains as a convenient thing only for the public API. --- arrow/array/compare.go | 70 +++++++++++++++++++++++++++++------------- 1 file changed, 49 insertions(+), 21 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index d8a9552e9..19252364a 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -26,7 +26,11 @@ import ( ) // RecordEqual reports whether the two provided records are equal. -func RecordEqual(left, right arrow.RecordBatch) bool { +func RecordEqual(left, right arrow.RecordBatch, opts ...EqualOption) bool { + return recordEqual(left, right, newEqualOption(opts...)) +} + +func recordEqual(left, right arrow.RecordBatch, opt equalOption) bool { switch { case left.NumCols() != right.NumCols(): return false @@ -37,7 +41,7 @@ func RecordEqual(left, right arrow.RecordBatch) bool { for i := range left.Columns() { lc := left.Column(i) rc := right.Column(i) - if !Equal(lc, rc) { + if !equal(lc, rc, opt) { return false } } @@ -47,6 +51,10 @@ func RecordEqual(left, right arrow.RecordBatch) bool { // RecordApproxEqual reports whether the two provided records are approximately equal. // For non-floating point columns, it is equivalent to RecordEqual. func RecordApproxEqual(left, right arrow.RecordBatch, opts ...EqualOption) bool { + return recordApproxEqual(left, right, newEqualOption(opts...)) +} + +func recordApproxEqual(left, right arrow.RecordBatch, opt equalOption) bool { switch { case left.NumCols() != right.NumCols(): return false @@ -54,8 +62,6 @@ func RecordApproxEqual(left, right arrow.RecordBatch, opts ...EqualOption) bool return false } - opt := newEqualOption(opts...) - for i := range left.Columns() { lc := left.Column(i) rc := right.Column(i) @@ -106,7 +112,11 @@ func chunkedBinaryApply(left, right *arrow.Chunked, fn func(left arrow.Array, lb } // ChunkedEqual reports whether two chunked arrays are equal regardless of their chunkings -func ChunkedEqual(left, right *arrow.Chunked) bool { +func ChunkedEqual(left, right *arrow.Chunked, opts ...EqualOption) bool { + return chunkedEqual(left, right, newEqualOption(opts...)) +} + +func chunkedEqual(left, right *arrow.Chunked, opt equalOption) bool { switch { case left == right: return true @@ -130,6 +140,10 @@ func ChunkedEqual(left, right *arrow.Chunked) bool { // ChunkedApproxEqual reports whether two chunked arrays are approximately equal regardless of their chunkings // for non-floating point arrays, this is equivalent to ChunkedEqual func ChunkedApproxEqual(left, right *arrow.Chunked, opts ...EqualOption) bool { + return chunkedApproxEqual(left, right, newEqualOption(opts...)) +} + +func chunkedApproxEqual(left, right *arrow.Chunked, opt equalOption) bool { switch { case left == right: return true @@ -143,7 +157,7 @@ func ChunkedApproxEqual(left, right *arrow.Chunked, opts ...EqualOption) bool { var isequal = true chunkedBinaryApply(left, right, func(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64) bool { - isequal = SliceApproxEqual(left, lbeg, lend, right, rbeg, rend, opts...) + isequal = sliceApproxEqual(left, lbeg, lend, right, rbeg, rend, opt) return isequal }) @@ -151,7 +165,11 @@ func ChunkedApproxEqual(left, right *arrow.Chunked, opts ...EqualOption) bool { } // TableEqual returns if the two tables have the same data in the same schema -func TableEqual(left, right arrow.Table) bool { +func TableEqual(left, right arrow.Table, opts ...EqualOption) bool { + return tableEqual(left, right, newEqualOption(opts...)) +} + +func tableEqual(left, right arrow.Table, opt equalOption) bool { switch { case left.NumCols() != right.NumCols(): return false @@ -166,15 +184,19 @@ func TableEqual(left, right arrow.Table) bool { return false } - if !ChunkedEqual(lc.Data(), rc.Data()) { + if !chunkedEqual(lc.Data(), rc.Data(), opt) { return false } } return true } -// TableEqual returns if the two tables have the approximately equal data in the same schema +// TableApproxEqual returns if the two tables have the approximately equal data in the same schema func TableApproxEqual(left, right arrow.Table, opts ...EqualOption) bool { + return tableApproxEqual(left, right, newEqualOption(opts...)) +} + +func tableApproxEqual(left, right arrow.Table, opt equalOption) bool { switch { case left.NumCols() != right.NumCols(): return false @@ -189,7 +211,7 @@ func TableApproxEqual(left, right arrow.Table, opts ...EqualOption) bool { return false } - if !ChunkedApproxEqual(lc.Data(), rc.Data(), opts...) { + if !chunkedApproxEqual(lc.Data(), rc.Data(), opt) { return false } } @@ -197,9 +219,13 @@ func TableApproxEqual(left, right arrow.Table, opts ...EqualOption) bool { } // Equal reports whether the two provided arrays are equal. -func Equal(left, right arrow.Array) bool { +func Equal(left, right arrow.Array, opts ...EqualOption) bool { + return equal(left, right, newEqualOption(opts...)) +} + +func equal(left, right arrow.Array, opt equalOption) bool { switch { - case !baseArrayEqual(left, right): + case !baseArrayEqual(left, right, opt): return false case left.Len() == 0: return true @@ -345,26 +371,29 @@ func Equal(left, right arrow.Array) bool { return arrayDenseUnionEqual(l, r) case *RunEndEncoded: r := right.(*RunEndEncoded) - return arrayRunEndEncodedEqual(l, r) + return arrayRunEndEncodedEqual(l, r, opt) default: panic(fmt.Errorf("arrow/array: unknown array type %T", l)) } } // SliceEqual reports whether slices left[lbeg:lend] and right[rbeg:rend] are equal. -func SliceEqual(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64) bool { +func SliceEqual(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64, opts ...EqualOption) bool { + return sliceEqual(left, lbeg, lend, right, rbeg, rend, newEqualOption(opts...)) +} + +func sliceEqual(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64, opt equalOption) bool { l := NewSlice(left, lbeg, lend) defer l.Release() r := NewSlice(right, rbeg, rend) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) } // SliceApproxEqual reports whether slices left[lbeg:lend] and right[rbeg:rend] are approximately equal. func SliceApproxEqual(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64, opts ...EqualOption) bool { - opt := newEqualOption(opts...) - return sliceApproxEqual(left, lbeg, lend, right, rbeg, rend, opt) + return sliceApproxEqual(left, lbeg, lend, right, rbeg, rend, newEqualOption(opts...)) } func sliceApproxEqual(left arrow.Array, lbeg, lend int64, right arrow.Array, rbeg, rend int64, opt equalOption) bool { @@ -455,13 +484,12 @@ func WithUnorderedMapKeys(v bool) EqualOption { // ApproxEqual reports whether the two provided arrays are approximately equal. // For non-floating point arrays, it is equivalent to Equal. func ApproxEqual(left, right arrow.Array, opts ...EqualOption) bool { - opt := newEqualOption(opts...) - return arrayApproxEqual(left, right, opt) + return arrayApproxEqual(left, right, newEqualOption(opts...)) } func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { switch { - case !baseArrayEqual(left, right): + case !baseArrayEqual(left, right, opt): return false case left.Len() == 0: return true @@ -616,7 +644,7 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { } } -func baseArrayEqual(left, right arrow.Array) bool { +func baseArrayEqual(left, right arrow.Array, opt equalOption) bool { switch { case left.Len() != right.Len(): return false From 29b75f050a628ddefd51c259d473f243268e293d Mon Sep 17 00:00:00 2001 From: serramatutu Date: Thu, 30 Apr 2026 12:48:06 +0200 Subject: [PATCH 02/18] Add `nullable` to `equalOption` This allows for 2 things: 1. Users can now explicitly pass into `EqualOption` if they want the comparison functions to compare the nullable buffer or not. 2. Struct and record comparison can change the value of the `nullable` option depending on `innerField.Nullable`, making the comparison semantically accurate. --- arrow/array/compare.go | 64 ++++++++++++++++++++++++++++--------- arrow/array/compare_test.go | 43 +++++++++++++++++++++++++ 2 files changed, 92 insertions(+), 15 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index 19252364a..c8e42aff8 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -39,6 +39,14 @@ func recordEqual(left, right arrow.RecordBatch, opt equalOption) bool { } for i := range left.Columns() { + lf := left.Schema().Field(i) + rf := left.Schema().Field(i) + if !lf.Equal(rf) { + return false + } + + opt.nullable = lf.Nullable + lc := left.Column(i) rc := right.Column(i) if !equal(lc, rc, opt) { @@ -63,6 +71,14 @@ func recordApproxEqual(left, right arrow.RecordBatch, opt equalOption) bool { } for i := range left.Columns() { + lf := left.Schema().Field(i) + rf := left.Schema().Field(i) + if !lf.Equal(rf) { + return false + } + + opt.nullable = lf.Nullable + lc := left.Column(i) rc := right.Column(i) if !arrayApproxEqual(lc, rc, opt) { @@ -122,7 +138,7 @@ func chunkedEqual(left, right *arrow.Chunked, opt equalOption) bool { return true case left.Len() != right.Len(): return false - case left.NullN() != right.NullN(): + case opt.nullable && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): return false @@ -149,7 +165,7 @@ func chunkedApproxEqual(left, right *arrow.Chunked, opt equalOption) bool { return true case left.Len() != right.Len(): return false - case left.NullN() != right.NullN(): + case opt.nullable && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): return false @@ -178,12 +194,16 @@ func tableEqual(left, right arrow.Table, opt equalOption) bool { } for i := 0; int64(i) < left.NumCols(); i++ { - lc := left.Column(i) - rc := right.Column(i) - if !lc.Field().Equal(rc.Field()) { + lf := left.Schema().Field(i) + rf := left.Schema().Field(i) + if !lf.Equal(rf) { return false } + opt.nullable = lf.Nullable + + lc := left.Column(i) + rc := right.Column(i) if !chunkedEqual(lc.Data(), rc.Data(), opt) { return false } @@ -205,12 +225,16 @@ func tableApproxEqual(left, right arrow.Table, opt equalOption) bool { } for i := 0; int64(i) < left.NumCols(); i++ { - lc := left.Column(i) - rc := right.Column(i) - if !lc.Field().Equal(rc.Field()) { + lf := left.Schema().Field(i) + rf := left.Schema().Field(i) + if !lf.Equal(rf) { return false } + opt.nullable = lf.Nullable + + lc := left.Column(i) + rc := right.Column(i) if !chunkedApproxEqual(lc.Data(), rc.Data(), opt) { return false } @@ -229,7 +253,7 @@ func equal(left, right arrow.Array, opt equalOption) bool { return false case left.Len() == 0: return true - case left.NullN() == left.Len(): + case opt.nullable && left.NullN() == left.Len(): return true } @@ -411,6 +435,7 @@ type equalOption struct { atol float64 // absolute tolerance nansEq bool // whether NaNs are considered equal. unorderedMapKeys bool // whether maps are allowed to have different entries order + nullable bool // whether the fields being compared are considered nullable } func (eq equalOption) f16(f1, f2 float16.Num) bool { @@ -446,8 +471,9 @@ func (eq equalOption) f64(v1, v2 float64) bool { func newEqualOption(opts ...EqualOption) equalOption { eq := equalOption{ - atol: defaultAbsoluteTolerance, - nansEq: false, + atol: defaultAbsoluteTolerance, + nansEq: false, + nullable: true, } for _, opt := range opts { opt(&eq) @@ -481,6 +507,14 @@ func WithUnorderedMapKeys(v bool) EqualOption { } } +// WithNullable sets whether the comparison function will consider both fields as nullable. If they're non-nullable, their +// valids buffer will be ignored for comparison and the underlying values will be used instead +func WithNullable(v bool) EqualOption { + return func(o *equalOption) { + o.nullable = v + } +} + // ApproxEqual reports whether the two provided arrays are approximately equal. // For non-floating point arrays, it is equivalent to Equal. func ApproxEqual(left, right arrow.Array, opts ...EqualOption) bool { @@ -493,7 +527,7 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { return false case left.Len() == 0: return true - case left.NullN() == left.Len(): + case opt.nullable && left.NullN() == left.Len(): return true } @@ -648,11 +682,11 @@ func baseArrayEqual(left, right arrow.Array, opt equalOption) bool { switch { case left.Len() != right.Len(): return false - case left.NullN() != right.NullN(): + case opt.nullable && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): // We do not check for metadata as in the C++ implementation. return false - case !validityBitmapEqual(left, right): + case opt.nullable && !validityBitmapEqual(left, right): return false } return true @@ -882,7 +916,7 @@ func arrayApproxEqualSingleMapEntry(left, right *Struct, opt equalOption) bool { switch { case left.Len() != right.Len(): return false - case left.NullN() != right.NullN(): + case opt.nullable && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): // We do not check for metadata as in the C++ implementation. return false diff --git a/arrow/array/compare_test.go b/arrow/array/compare_test.go index 0671def62..9de3445f6 100644 --- a/arrow/array/compare_test.go +++ b/arrow/array/compare_test.go @@ -391,6 +391,49 @@ func TestArrayApproxEqualFloats(t *testing.T) { } } +func TestArrayEqualNonNullable(t *testing.T) { + for name, recs := range arrdata.Records { + t.Run(name, func(t *testing.T) { + rec := recs[0] + + // Clone the schema and make everything non-nullable + fields := rec.Schema().Fields() + meta := rec.Schema().Metadata() + for i := range fields { + fields[i].Nullable = false + } + schema := arrow.NewSchema(fields, &meta) + + for i, rawCol := range rec.Columns() { + // make a clone of the column with NullN=0 + col := array.MakeFromData(array.NewData( + rawCol.DataType(), + rawCol.Len(), + rawCol.Data().Buffers(), + rawCol.Data().Children(), + 0, + 0, + )) + t.Run(schema.Field(i).Name, func(t *testing.T) { + arr := col + if !array.Equal(arr, arr, array.WithNullable(false)) { + t.Fatalf("identical arrays should compare equal:\narray=%v", arr) + } + sub1 := array.NewSlice(arr, 1, int64(arr.Len())) + defer sub1.Release() + + sub2 := array.NewSlice(arr, 0, int64(arr.Len()-1)) + defer sub2.Release() + + if array.Equal(sub1, sub2) && name != "nulls" { + t.Fatalf("non-identical arrays should not compare equal:\nsub1=%v\nsub2=%v\narrf=%v\n", sub1, sub2, arr) + } + }) + } + }) + } +} + func testStringMap(mem memory.Allocator, m map[string]string, keys []string) *array.Map { dt := arrow.MapOf(arrow.BinaryTypes.String, arrow.BinaryTypes.String) builder := array.NewMapBuilderWithType(mem, dt) From fbfe0ca30a4b0c06e31a991b5718f788aad8fad0 Mon Sep 17 00:00:00 2001 From: serramatutu Date: Thu, 30 Apr 2026 13:16:33 +0200 Subject: [PATCH 03/18] Check for nullable opt in most array compare implementations --- arrow/array/binary.go | 12 +- arrow/array/boolean.go | 4 +- arrow/array/compare.go | 188 ++++++++++++++++---------------- arrow/array/decimal.go | 4 +- arrow/array/dictionary.go | 4 +- arrow/array/encoded.go | 4 +- arrow/array/extension.go | 4 +- arrow/array/fixed_size_list.go | 6 +- arrow/array/fixedsize_binary.go | 4 +- arrow/array/interval.go | 12 +- arrow/array/list.go | 24 ++-- arrow/array/map.go | 4 +- arrow/array/numeric_generic.go | 4 +- arrow/array/string.go | 12 +- arrow/array/struct.go | 5 +- arrow/array/timestamp.go | 4 +- 16 files changed, 150 insertions(+), 145 deletions(-) diff --git a/arrow/array/binary.go b/arrow/array/binary.go index a8e77ae9c..4ddfdbf6c 100644 --- a/arrow/array/binary.go +++ b/arrow/array/binary.go @@ -223,9 +223,9 @@ func (a *Binary) ValidateFull() error { return nil } -func arrayEqualBinary(left, right *Binary) bool { +func arrayEqualBinary(left, right *Binary, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !bytes.Equal(left.Value(i), right.Value(i)) { @@ -417,9 +417,9 @@ func (a *LargeBinary) ValidateFull() error { return nil } -func arrayEqualLargeBinary(left, right *LargeBinary) bool { +func arrayEqualLargeBinary(left, right *LargeBinary, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !bytes.Equal(left.Value(i), right.Value(i)) { @@ -539,10 +539,10 @@ func (a *BinaryView) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualBinaryView(left, right *BinaryView) bool { +func arrayEqualBinaryView(left, right *BinaryView, opt equalOption) bool { leftBufs, rightBufs := left.dataBuffers, right.dataBuffers for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !left.ValueHeader(i).Equals(leftBufs, right.ValueHeader(i), rightBufs) { diff --git a/arrow/array/boolean.go b/arrow/array/boolean.go index d579fa0c8..e57dd3e6e 100644 --- a/arrow/array/boolean.go +++ b/arrow/array/boolean.go @@ -109,9 +109,9 @@ func (a *Boolean) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualBoolean(left, right *Boolean) bool { +func arrayEqualBoolean(left, right *Boolean, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { diff --git a/arrow/array/compare.go b/arrow/array/compare.go index c8e42aff8..ec56cb6d7 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -266,127 +266,127 @@ func equal(left, right arrow.Array, opt equalOption) bool { return true case *Boolean: r := right.(*Boolean) - return arrayEqualBoolean(l, r) + return arrayEqualBoolean(l, r, opt) case *FixedSizeBinary: r := right.(*FixedSizeBinary) - return arrayEqualFixedSizeBinary(l, r) + return arrayEqualFixedSizeBinary(l, r, opt) case *Binary: r := right.(*Binary) - return arrayEqualBinary(l, r) + return arrayEqualBinary(l, r, opt) case *String: r := right.(*String) - return arrayEqualString(l, r) + return arrayEqualString(l, r, opt) case *LargeBinary: r := right.(*LargeBinary) - return arrayEqualLargeBinary(l, r) + return arrayEqualLargeBinary(l, r, opt) case *LargeString: r := right.(*LargeString) - return arrayEqualLargeString(l, r) + return arrayEqualLargeString(l, r, opt) case *BinaryView: r := right.(*BinaryView) - return arrayEqualBinaryView(l, r) + return arrayEqualBinaryView(l, r, opt) case *StringView: r := right.(*StringView) - return arrayEqualStringView(l, r) + return arrayEqualStringView(l, r, opt) case *Int8: r := right.(*Int8) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int16: r := right.(*Int16) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int32: r := right.(*Int32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int64: r := right.(*Int64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint8: r := right.(*Uint8) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint16: r := right.(*Uint16) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint32: r := right.(*Uint32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint64: r := right.(*Uint64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Float16: r := right.(*Float16) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Float32: r := right.(*Float32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Float64: r := right.(*Float64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Decimal32: r := right.(*Decimal32) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal64: r := right.(*Decimal64) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal128: r := right.(*Decimal128) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal256: r := right.(*Decimal256) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Date32: r := right.(*Date32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Date64: r := right.(*Date64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Time32: r := right.(*Time32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Time64: r := right.(*Time64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Timestamp: r := right.(*Timestamp) - return arrayEqualTimestamp(l, r) + return arrayEqualTimestamp(l, r, opt) case *List: r := right.(*List) - return arrayEqualList(l, r) + return arrayEqualList(l, r, opt) case *LargeList: r := right.(*LargeList) - return arrayEqualLargeList(l, r) + return arrayEqualLargeList(l, r, opt) case *ListView: r := right.(*ListView) - return arrayEqualListView(l, r) + return arrayEqualListView(l, r, opt) case *LargeListView: r := right.(*LargeListView) - return arrayEqualLargeListView(l, r) + return arrayEqualLargeListView(l, r, opt) case *FixedSizeList: r := right.(*FixedSizeList) - return arrayEqualFixedSizeList(l, r) + return arrayEqualFixedSizeList(l, r, opt) case *Struct: r := right.(*Struct) - return arrayEqualStruct(l, r) + return arrayEqualStruct(l, r, opt) case *MonthInterval: r := right.(*MonthInterval) - return arrayEqualMonthInterval(l, r) + return arrayEqualMonthInterval(l, r, opt) case *DayTimeInterval: r := right.(*DayTimeInterval) - return arrayEqualDayTimeInterval(l, r) + return arrayEqualDayTimeInterval(l, r, opt) case *MonthDayNanoInterval: r := right.(*MonthDayNanoInterval) - return arrayEqualMonthDayNanoInterval(l, r) + return arrayEqualMonthDayNanoInterval(l, r, opt) case *Duration: r := right.(*Duration) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Map: r := right.(*Map) - return arrayEqualMap(l, r) + return arrayEqualMap(l, r, opt) case ExtensionArray: r := right.(ExtensionArray) - return arrayEqualExtension(l, r) + return arrayEqualExtension(l, r, opt) case *Dictionary: r := right.(*Dictionary) - return arrayEqualDict(l, r) + return arrayEqualDict(l, r, opt) case *SparseUnion: r := right.(*SparseUnion) return arraySparseUnionEqual(l, r) @@ -540,52 +540,52 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { return true case *Boolean: r := right.(*Boolean) - return arrayEqualBoolean(l, r) + return arrayEqualBoolean(l, r, opt) case *FixedSizeBinary: r := right.(*FixedSizeBinary) - return arrayEqualFixedSizeBinary(l, r) + return arrayEqualFixedSizeBinary(l, r, opt) case *Binary: r := right.(*Binary) - return arrayEqualBinary(l, r) + return arrayEqualBinary(l, r, opt) case *String: r := right.(*String) - return arrayApproxEqualString(l, r) + return arrayApproxEqualString(l, r, opt) case *LargeBinary: r := right.(*LargeBinary) - return arrayEqualLargeBinary(l, r) + return arrayEqualLargeBinary(l, r, opt) case *LargeString: r := right.(*LargeString) - return arrayApproxEqualLargeString(l, r) + return arrayApproxEqualLargeString(l, r, opt) case *BinaryView: r := right.(*BinaryView) - return arrayEqualBinaryView(l, r) + return arrayEqualBinaryView(l, r, opt) case *StringView: r := right.(*StringView) - return arrayApproxEqualStringView(l, r) + return arrayApproxEqualStringView(l, r, opt) case *Int8: r := right.(*Int8) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int16: r := right.(*Int16) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int32: r := right.(*Int32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Int64: r := right.(*Int64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint8: r := right.(*Uint8) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint16: r := right.(*Uint16) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint32: r := right.(*Uint32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Uint64: r := right.(*Uint64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Float16: r := right.(*Float16) return arrayApproxEqualFloat16(l, r, opt) @@ -597,31 +597,31 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { return arrayApproxEqualFloat64(l, r, opt) case *Decimal32: r := right.(*Decimal32) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal64: r := right.(*Decimal64) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal128: r := right.(*Decimal128) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Decimal256: r := right.(*Decimal256) - return arrayEqualDecimal(l, r) + return arrayEqualDecimal(l, r, opt) case *Date32: r := right.(*Date32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Date64: r := right.(*Date64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Time32: r := right.(*Time32) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Time64: r := right.(*Time64) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Timestamp: r := right.(*Timestamp) - return arrayEqualTimestamp(l, r) + return arrayEqualTimestamp(l, r, opt) case *List: r := right.(*List) return arrayApproxEqualList(l, r, opt) @@ -642,16 +642,16 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { return arrayApproxEqualStruct(l, r, opt) case *MonthInterval: r := right.(*MonthInterval) - return arrayEqualMonthInterval(l, r) + return arrayEqualMonthInterval(l, r, opt) case *DayTimeInterval: r := right.(*DayTimeInterval) - return arrayEqualDayTimeInterval(l, r) + return arrayEqualDayTimeInterval(l, r, opt) case *MonthDayNanoInterval: r := right.(*MonthDayNanoInterval) - return arrayEqualMonthDayNanoInterval(l, r) + return arrayEqualMonthDayNanoInterval(l, r, opt) case *Duration: r := right.(*Duration) - return arrayEqualFixedWidth(l, r) + return arrayEqualFixedWidth(l, r, opt) case *Map: r := right.(*Map) if opt.unorderedMapKeys { @@ -706,9 +706,9 @@ func validityBitmapEqual(left, right arrow.Array) bool { return true } -func arrayApproxEqualString(left, right *String) bool { +func arrayApproxEqualString(left, right *String, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if stripNulls(left.Value(i)) != stripNulls(right.Value(i)) { @@ -718,9 +718,9 @@ func arrayApproxEqualString(left, right *String) bool { return true } -func arrayApproxEqualLargeString(left, right *LargeString) bool { +func arrayApproxEqualLargeString(left, right *LargeString, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if stripNulls(left.Value(i)) != stripNulls(right.Value(i)) { @@ -730,9 +730,9 @@ func arrayApproxEqualLargeString(left, right *LargeString) bool { return true } -func arrayApproxEqualStringView(left, right *StringView) bool { +func arrayApproxEqualStringView(left, right *StringView, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if stripNulls(left.Value(i)) != stripNulls(right.Value(i)) { @@ -744,7 +744,7 @@ func arrayApproxEqualStringView(left, right *StringView) bool { func arrayApproxEqualFloat16(left, right *Float16, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !opt.f16(left.Value(i), right.Value(i)) { @@ -756,7 +756,7 @@ func arrayApproxEqualFloat16(left, right *Float16, opt equalOption) bool { func arrayApproxEqualFloat32(left, right *Float32, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !opt.f32(left.Value(i), right.Value(i)) { @@ -768,7 +768,7 @@ func arrayApproxEqualFloat32(left, right *Float32, opt equalOption) bool { func arrayApproxEqualFloat64(left, right *Float64, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !opt.f64(left.Value(i), right.Value(i)) { @@ -780,7 +780,7 @@ func arrayApproxEqualFloat64(left, right *Float64, opt equalOption) bool { func arrayApproxEqualList(left, right *List, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -799,7 +799,7 @@ func arrayApproxEqualList(left, right *List, opt equalOption) bool { func arrayApproxEqualLargeList(left, right *LargeList, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -818,7 +818,7 @@ func arrayApproxEqualLargeList(left, right *LargeList, opt equalOption) bool { func arrayApproxEqualListView(left, right *ListView, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -837,7 +837,7 @@ func arrayApproxEqualListView(left, right *ListView, opt equalOption) bool { func arrayApproxEqualLargeListView(left, right *LargeListView, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -856,7 +856,7 @@ func arrayApproxEqualLargeListView(left, right *LargeListView, opt equalOption) func arrayApproxEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -874,11 +874,15 @@ func arrayApproxEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) } func arrayApproxEqualStruct(left, right *Struct, opt equalOption) bool { - return bitutils.VisitSetBitRuns( - left.NullBitmapBytes(), - int64(left.Offset()), int64(left.Len()), - approxEqualStructRun(left, right, opt), - ) == nil + visitFn := approxEqualStructRun(left, right, opt) + if opt.nullable { + return bitutils.VisitSetBitRuns( + left.NullBitmapBytes(), + int64(left.Offset()), int64(left.Len()), + visitFn, + ) == nil + } + return visitFn(0, int64(left.Len())) == nil } func approxEqualStructRun(left, right *Struct, opt equalOption) bitutils.VisitFn { @@ -895,7 +899,7 @@ func approxEqualStructRun(left, right *Struct, opt equalOption) bitutils.VisitFn // arrayApproxEqualMap doesn't care about the order of keys (in Go map traversal order is undefined) func arrayApproxEqualMap(left, right *Map, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !arrayApproxEqualSingleMapEntry(left.newListValue(i).(*Struct), right.newListValue(i).(*Struct), opt) { @@ -926,7 +930,7 @@ func arrayApproxEqualSingleMapEntry(left, right *Struct, opt equalOption) bool { used := make(map[int]bool, right.Len()) for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } @@ -936,7 +940,7 @@ func arrayApproxEqualSingleMapEntry(left, right *Struct, opt equalOption) bool { if used[j] { continue } - if right.IsNull(j) { + if opt.nullable && right.IsNull(j) { used[j] = true continue } diff --git a/arrow/array/decimal.go b/arrow/array/decimal.go index 704b1d932..b4e24f69e 100644 --- a/arrow/array/decimal.go +++ b/arrow/array/decimal.go @@ -110,9 +110,9 @@ func (a *baseDecimal[T]) MarshalJSON() ([]byte, error) { func arrayEqualDecimal[T interface { decimal.DecimalTypes decimal.Num[T] -}](left, right *baseDecimal[T]) bool { +}](left, right *baseDecimal[T], opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } diff --git a/arrow/array/dictionary.go b/arrow/array/dictionary.go index 38a43f77c..9645b788e 100644 --- a/arrow/array/dictionary.go +++ b/arrow/array/dictionary.go @@ -303,8 +303,8 @@ func (d *Dictionary) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualDict(l, r *Dictionary) bool { - return Equal(l.Dictionary(), r.Dictionary()) && Equal(l.indices, r.indices) +func arrayEqualDict(l, r *Dictionary, opt equalOption) bool { + return equal(l.Dictionary(), r.Dictionary(), opt) && equal(l.indices, r.indices, opt) } func arrayApproxEqualDict(l, r *Dictionary, opt equalOption) bool { diff --git a/arrow/array/encoded.go b/arrow/array/encoded.go index 8d628ffc2..4cfc346fe 100644 --- a/arrow/array/encoded.go +++ b/arrow/array/encoded.go @@ -260,14 +260,14 @@ func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayRunEndEncodedEqual(l, r *RunEndEncoded) bool { +func arrayRunEndEncodedEqual(l, r *RunEndEncoded, opt equalOption) bool { // types were already checked before getting here, so we know // the encoded types are equal mr := encoded.NewMergedRuns([2]arrow.Array{l, r}) for mr.Next() { lIndex := mr.IndexIntoArray(0) rIndex := mr.IndexIntoArray(1) - if !SliceEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1) { + if !sliceEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1, opt) { return false } } diff --git a/arrow/array/extension.go b/arrow/array/extension.go index e509b5e0f..21ed3b249 100644 --- a/arrow/array/extension.go +++ b/arrow/array/extension.go @@ -46,12 +46,12 @@ type ExtensionArray interface { // two extension arrays are equal if their data types are equal and // their underlying storage arrays are equal. -func arrayEqualExtension(l, r ExtensionArray) bool { +func arrayEqualExtension(l, r ExtensionArray, opt equalOption) bool { if !arrow.TypeEqual(l.DataType(), r.DataType()) { return false } - return Equal(l.Storage(), r.Storage()) + return equal(l.Storage(), r.Storage(), opt) } // two extension arrays are approximately equal if their data types are diff --git a/arrow/array/fixed_size_list.go b/arrow/array/fixed_size_list.go index d382ebe93..4b3ff286f 100644 --- a/arrow/array/fixed_size_list.go +++ b/arrow/array/fixed_size_list.go @@ -84,9 +84,9 @@ func (a *FixedSizeList) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func arrayEqualFixedSizeList(left, right *FixedSizeList) bool { +func arrayEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -94,7 +94,7 @@ func arrayEqualFixedSizeList(left, right *FixedSizeList) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) }() if !o { return false diff --git a/arrow/array/fixedsize_binary.go b/arrow/array/fixedsize_binary.go index 31d507c5b..44ca44883 100644 --- a/arrow/array/fixedsize_binary.go +++ b/arrow/array/fixedsize_binary.go @@ -106,9 +106,9 @@ func (a *FixedSizeBinary) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualFixedSizeBinary(left, right *FixedSizeBinary) bool { +func arrayEqualFixedSizeBinary(left, right *FixedSizeBinary, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !bytes.Equal(left.Value(i), right.Value(i)) { diff --git a/arrow/array/interval.go b/arrow/array/interval.go index 2c029c252..128cdfadf 100644 --- a/arrow/array/interval.go +++ b/arrow/array/interval.go @@ -120,9 +120,9 @@ func (a *MonthInterval) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualMonthInterval(left, right *MonthInterval) bool { +func arrayEqualMonthInterval(left, right *MonthInterval, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { @@ -423,9 +423,9 @@ func (a *DayTimeInterval) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualDayTimeInterval(left, right *DayTimeInterval) bool { +func arrayEqualDayTimeInterval(left, right *DayTimeInterval, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { @@ -727,9 +727,9 @@ func (a *MonthDayNanoInterval) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualMonthDayNanoInterval(left, right *MonthDayNanoInterval) bool { +func arrayEqualMonthDayNanoInterval(left, right *MonthDayNanoInterval, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { diff --git a/arrow/array/list.go b/arrow/array/list.go index d72887bb8..df45223f0 100644 --- a/arrow/array/list.go +++ b/arrow/array/list.go @@ -129,9 +129,9 @@ func (a *List) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayEqualList(left, right *List) bool { +func arrayEqualList(left, right *List, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -139,7 +139,7 @@ func arrayEqualList(left, right *List) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) }() if !o { return false @@ -261,9 +261,9 @@ func (a *LargeList) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayEqualLargeList(left, right *LargeList) bool { +func arrayEqualLargeList(left, right *LargeList, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -271,7 +271,7 @@ func arrayEqualLargeList(left, right *LargeList) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) }() if !o { return false @@ -736,9 +736,9 @@ func (a *ListView) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayEqualListView(left, right *ListView) bool { +func arrayEqualListView(left, right *ListView, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -746,7 +746,7 @@ func arrayEqualListView(left, right *ListView) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) }() if !o { return false @@ -883,9 +883,9 @@ func (a *LargeListView) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayEqualLargeListView(left, right *LargeListView) bool { +func arrayEqualLargeListView(left, right *LargeListView, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } o := func() bool { @@ -893,7 +893,7 @@ func arrayEqualLargeListView(left, right *LargeListView) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return Equal(l, r) + return equal(l, r, opt) }() if !o { return false diff --git a/arrow/array/map.go b/arrow/array/map.go index 71d4d1382..af6c28183 100644 --- a/arrow/array/map.go +++ b/arrow/array/map.go @@ -106,9 +106,9 @@ func (a *Map) Release() { a.items.Release() } -func arrayEqualMap(left, right *Map) bool { +func arrayEqualMap(left, right *Map, opt equalOption) bool { // since Map is implemented using a list, we can just use arrayEqualList - return arrayEqualList(left.List, right.List) + return arrayEqualList(left.List, right.List, opt) } type MapBuilder struct { diff --git a/arrow/array/numeric_generic.go b/arrow/array/numeric_generic.go index 1b671fc76..49a369cdd 100644 --- a/arrow/array/numeric_generic.go +++ b/arrow/array/numeric_generic.go @@ -436,9 +436,9 @@ func NewDate64Data(data arrow.ArrayData) *Date64 { func (a *Date64) Date64Values() []arrow.Date64 { return a.Values() } -func arrayEqualFixedWidth[T arrow.FixedWidthType](left, right arrow.TypedArray[T]) bool { +func arrayEqualFixedWidth[T arrow.FixedWidthType](left, right arrow.TypedArray[T], opt equalOption) bool { for i := range left.Len() { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { diff --git a/arrow/array/string.go b/arrow/array/string.go index 7c2ab0744..f0d9be401 100644 --- a/arrow/array/string.go +++ b/arrow/array/string.go @@ -231,9 +231,9 @@ func (a *String) ValidateFull() error { return nil } -func arrayEqualString(left, right *String) bool { +func arrayEqualString(left, right *String, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { @@ -435,9 +435,9 @@ func (a *LargeString) ValidateFull() error { return nil } -func arrayEqualLargeString(left, right *LargeString) bool { +func arrayEqualLargeString(left, right *LargeString, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { @@ -541,10 +541,10 @@ func (a *StringView) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualStringView(left, right *StringView) bool { +func arrayEqualStringView(left, right *StringView, opt equalOption) bool { leftBufs, rightBufs := left.dataBuffers, right.dataBuffers for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if !left.ValueHeader(i).Equals(leftBufs, right.ValueHeader(i), rightBufs) { diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 07505de1c..203dc1f70 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -239,10 +239,11 @@ func (a *Struct) MarshalJSON() ([]byte, error) { return buf.Bytes(), nil } -func arrayEqualStruct(left, right *Struct) bool { +func arrayEqualStruct(left, right *Struct, opt equalOption) bool { for i, lf := range left.fields { rf := right.fields[i] - if !Equal(lf, rf) { + opt.nullable = left.data.dtype.(*arrow.StructType).Field(i).Nullable + if !equal(lf, rf, opt) { return false } } diff --git a/arrow/array/timestamp.go b/arrow/array/timestamp.go index 5ac0ee6cc..55e9b5358 100644 --- a/arrow/array/timestamp.go +++ b/arrow/array/timestamp.go @@ -125,9 +125,9 @@ func (a *Timestamp) MarshalJSON() ([]byte, error) { return json.Marshal(vals) } -func arrayEqualTimestamp(left, right *Timestamp) bool { +func arrayEqualTimestamp(left, right *Timestamp, opt equalOption) bool { for i := 0; i < left.Len(); i++ { - if left.IsNull(i) { + if opt.nullable && left.IsNull(i) { continue } if left.Value(i) != right.Value(i) { From 1da53ae8c922b7253478e2326642434258724417 Mon Sep 17 00:00:00 2001 From: serramatutu Date: Tue, 28 Apr 2026 15:35:46 +0200 Subject: [PATCH 04/18] Add stricter tests to null JSON in record and struct --- arrow/array/record_test.go | 47 ++++++++++++++++++++++------ arrow/array/struct_test.go | 6 ++++ arrow/array/util_test.go | 63 +++++++++++++++++++++++++------------- 3 files changed, 86 insertions(+), 30 deletions(-) diff --git a/arrow/array/record_test.go b/arrow/array/record_test.go index a3924382a..eabb84f3e 100644 --- a/arrow/array/record_test.go +++ b/arrow/array/record_test.go @@ -17,6 +17,7 @@ package array_test import ( + "bytes" "fmt" "reflect" "strings" @@ -25,6 +26,7 @@ import ( "github.com/apache/arrow-go/v18/arrow" "github.com/apache/arrow-go/v18/arrow/array" "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/internal/json" "github.com/stretchr/testify/assert" ) @@ -485,9 +487,9 @@ func TestRecordBuilder(t *testing.T) { mapDt.SetItemNullable(false) schema := arrow.NewSchema( []arrow.Field{ - {Name: "f1-i32", Type: arrow.PrimitiveTypes.Int32}, - {Name: "f2-f64", Type: arrow.PrimitiveTypes.Float64}, - {Name: "map", Type: mapDt}, + {Name: "f1-i32", Type: arrow.PrimitiveTypes.Int32, Nullable: true}, + {Name: "f2-f64-notnull", Type: arrow.PrimitiveTypes.Float64, Nullable: false}, + {Name: "map", Type: mapDt, Nullable: true}, }, nil, ) @@ -498,11 +500,14 @@ func TestRecordBuilder(t *testing.T) { b.Retain() b.Release() - b.Field(0).(*array.Int32Builder).AppendValues([]int32{1, 2, 3}, nil) + b.Field(0).(*array.Int32Builder).AppendNull() + b.Field(0).(*array.Int32Builder).AppendValues([]int32{2, 3}, nil) b.Field(0).(*array.Int32Builder).AppendValues([]int32{4, 5}, nil) - b.Field(1).(*array.Float64Builder).AppendValues([]float64{1, 2, 3, 4, 5}, nil) + + b.Field(1).(*array.Float64Builder).AppendValues([]float64{1.1, 2.2, 3.3, 4.4, 5.5}, nil) + mb := b.Field(2).(*array.MapBuilder) - for i := 0; i < 5; i++ { + for i := range 5 { mb.Append(true) if i%3 == 0 { @@ -511,6 +516,12 @@ func TestRecordBuilder(t *testing.T) { } } + err := b.UnmarshalJSON([]byte(`{"f1-i32": 6, "f2-f64-notnull": 6.6, "map": [{"key": "4": "value": "d"}]}`)) + assert.NoError(t, err) + + err = b.UnmarshalJSON([]byte(`{"f1-i32": null, "f2-f64-notnull": null, "map": null}`)) + assert.NoError(t, err) + rec := b.NewRecordBatch() defer rec.Release() @@ -518,7 +529,7 @@ func TestRecordBuilder(t *testing.T) { t.Fatalf("invalid schema: got=%#v, want=%#v", got, want) } - if got, want := rec.NumRows(), int64(5); got != want { + if got, want := rec.NumRows(), int64(7); got != want { t.Fatalf("invalid number of rows: got=%d, want=%d", got, want) } if got, want := rec.NumCols(), int64(3); got != want { @@ -527,9 +538,27 @@ func TestRecordBuilder(t *testing.T) { if got, want := rec.ColumnName(0), schema.Field(0).Name; got != want { t.Fatalf("invalid column name: got=%q, want=%q", got, want) } - if got, want := rec.Column(2).String(), `[{["0" "2" "3"] ["a" "b" "c"]} {[] []} {[] []} {["3" "2" "3"] ["a" "b" "c"]} {[] []}]`; got != want { - t.Fatalf("invalid column name: got=%q, want=%q", got, want) + + if got, want := rec.Column(0).String(), `[(null) 2 3 4 5 6 (null)]`; got != want { + t.Fatalf("invalid column values: got=%q, want=%q", got, want) + } + if got, want := rec.Column(1).String(), `[1.1 2.2 3.3 4.4 5.5 6.6 0]`; got != want { + t.Fatalf("invalid column values: got=%q, want=%q", got, want) } + if got, want := rec.Column(2).String(), `[{["0" "2" "3"] ["a" "b" "c"]} {[] []} {[] []} {["3" "2" "3"] ["a" "b" "c"]} {[] []} {["4"] ["d"]} (null)]`; got != want { + t.Fatalf("invalid column values: got=%q, want=%q", got, want) + } + + // roundtripping from JSON with array.FromJSON should work + arr := array.RecordToStructArray(rec) + defer arr.Release() + jsonStr, err := json.Marshal(arr) + assert.NoError(t, err) + + roundtripped, _, err := array.FromJSON(mem, arr.DataType(), bytes.NewReader(jsonStr)) + defer roundtripped.Release() + assert.NoError(t, err) + assert.Truef(t, array.Equal(arr, roundtripped), "JSON round trip returns different array: got=%q, want=%d", arr, roundtripped) } func TestRecordBuilderResize(t *testing.T) { diff --git a/arrow/array/struct_test.go b/arrow/array/struct_test.go index 216b353ab..6027d0288 100644 --- a/arrow/array/struct_test.go +++ b/arrow/array/struct_test.go @@ -513,6 +513,12 @@ func TestStructArrayUnmarshalJSONMissingFields(t *testing.T) { panic: false, want: `{[(null)] [3] {[(null)] [(null)] ["test"]}}`, }, + { + name: "explicit null in required field", + jsonInput: `[{"f2": 3, "f3": {"f3_3": null}}]`, + panic: false, + want: `{[(null)] [3] {[(null)] [(null)] [""]}}`, + }, } for _, tc := range tests { diff --git a/arrow/array/util_test.go b/arrow/array/util_test.go index eb3de6a8b..84103d454 100644 --- a/arrow/array/util_test.go +++ b/arrow/array/util_test.go @@ -452,29 +452,50 @@ func TestArrRecordsJSONRoundTrip(t *testing.T) { continue } t.Run(k, func(t *testing.T) { - var buf bytes.Buffer - assert.NotPanics(t, func() { - enc := json.NewEncoder(&buf) - for _, r := range v { - if err := enc.Encode(r); err != nil { - panic(err) - } + for _, nullable := range []bool{true, false} { + var name string + if nullable { + name = "nullable" + } else { + name = "non-nullable" } - }) - - rdr := bytes.NewReader(buf.Bytes()) - var cur int64 - - mem := memory.NewCheckedAllocator(memory.NewGoAllocator()) - defer mem.AssertSize(t, 0) - - for _, r := range v { - rec, off, err := array.RecordFromJSON(mem, r.Schema(), rdr, array.WithStartOffset(cur)) - assert.NoError(t, err) - defer rec.Release() - assert.Truef(t, array.RecordApproxEqual(r, rec), "expected: %s\ngot: %s\n", r, rec) - cur += off + t.Run(name, func(t *testing.T) { + fields := v[0].Schema().Fields() + for i := range fields { + fields[i].Nullable = nullable + } + meta := v[0].Schema().Metadata() + schema := arrow.NewSchema(fields, &meta) + + var buf bytes.Buffer + assert.NotPanics(t, func() { + enc := json.NewEncoder(&buf) + for _, rawBatch := range v { + batch := array.NewRecordBatch(schema, rawBatch.Columns(), rawBatch.NumRows()) + if err := enc.Encode(batch); err != nil { + panic(err) + } + } + }) + + rdr := bytes.NewReader(buf.Bytes()) + var cur int64 + + mem := memory.NewCheckedAllocator(memory.NewGoAllocator()) + defer mem.AssertSize(t, 0) + + for _, rawBatch := range v { + batch := array.NewRecordBatch(schema, rawBatch.Columns(), rawBatch.NumRows()) + + rec, off, err := array.RecordFromJSON(mem, schema, rdr, array.WithStartOffset(cur)) + assert.NoError(t, err) + defer rec.Release() + + assert.Truef(t, array.RecordApproxEqual(batch, rec), "expected: %s\ngot: %s\n", batch, rec) + cur += off + } + }) } }) } From 7dc69eefb1926f0cfb671a96f80e82a8183df60a Mon Sep 17 00:00:00 2001 From: serramatutu Date: Tue, 28 Apr 2026 15:36:34 +0200 Subject: [PATCH 05/18] `AppendEmptyValue()` if field is nullable Modified `StructBuilder` and `RecordBuilder` to append the empty default value (usually zero) to the inner field if it's marked as nullable and the consumed value is null. --- arrow/array/record.go | 13 ++++++++++++- arrow/array/struct.go | 12 +++++++++++- internal/json/json.go | 5 +++++ 3 files changed, 28 insertions(+), 2 deletions(-) diff --git a/arrow/array/record.go b/arrow/array/record.go index b7f84180a..2f8948d4c 100644 --- a/arrow/array/record.go +++ b/arrow/array/record.go @@ -455,10 +455,21 @@ func (b *RecordBuilder) UnmarshalOne(dec *json.Decoder) error { } continue } + idx := indices[0] - if err := b.fields[indices[0]].UnmarshalOne(dec); err != nil { + var next json.RawMessage + if err := dec.Decode(&next); err != nil { return err } + + if json.IsNullMessage(next) && !b.schema.Field(idx).Nullable { + b.fields[idx].AppendEmptyValue() + } else { + sub := json.NewDecoder(bytes.NewReader(next)) + if err := b.fields[idx].UnmarshalOne(sub); err != nil { + return err + } + } } // consume the closing '}' diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 203dc1f70..76611dca0 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -488,9 +488,19 @@ func (b *StructBuilder) UnmarshalOne(dec *json.Decoder) error { continue } - if err := b.fields[idx].UnmarshalOne(dec); err != nil { + var next json.RawMessage + if err := dec.Decode(&next); err != nil { return err } + + if json.IsNullMessage(next) && !b.dtype.(*arrow.StructType).Field(idx).Nullable { + b.fields[idx].AppendEmptyValue() + } else { + sub := json.NewDecoder(bytes.NewReader(next)) + if err := b.fields[idx].UnmarshalOne(sub); err != nil { + return err + } + } } // Append null values to all optional fields that were not presented in the json input diff --git a/internal/json/json.go b/internal/json/json.go index b4c4c9f6e..7fdd0d863 100644 --- a/internal/json/json.go +++ b/internal/json/json.go @@ -20,6 +20,7 @@ package json import ( + "bytes" "io" "github.com/goccy/go-json" @@ -49,3 +50,7 @@ func NewDecoder(r io.Reader) *Decoder { func NewEncoder(w io.Writer) *Encoder { return json.NewEncoder(w) } + +func IsNullMessage(m RawMessage) bool { + return bytes.Equal(m, []byte("null")) +} From 554eddc6eba807023a9c2a501df998b3c4fe1cd8 Mon Sep 17 00:00:00 2001 From: serramatutu Date: Wed, 29 Apr 2026 10:50:43 +0200 Subject: [PATCH 06/18] Add `nullable` argument to `GetOneForMarshal()` The array implementations need a way of knowing whether to ignore the valids buffer or not. By default, it shouldn't ignore it if the array is being serialized by itself, like with `MarshalJSON()`. However, if the array is a part of a `Field` in a `RecordBatch` or `Struct`, then the value of the valids buffer might need to be ignored. This will be implemented in the next commit. --- arrow/array.go | 2 +- arrow/array/binary.go | 18 +++++++++--------- arrow/array/boolean.go | 4 ++-- arrow/array/decimal.go | 8 ++++---- arrow/array/decimal128_test.go | 2 +- arrow/array/decimal256_test.go | 2 +- arrow/array/dictionary.go | 8 ++++---- arrow/array/encoded.go | 10 +++++----- arrow/array/extension.go | 4 ++-- arrow/array/fixed_size_list.go | 6 +++--- arrow/array/fixedsize_binary.go | 4 ++-- arrow/array/float16.go | 4 ++-- arrow/array/interval.go | 16 ++++++++-------- arrow/array/list.go | 32 ++++++++++++++++---------------- arrow/array/null.go | 2 +- arrow/array/numeric_generic.go | 32 ++++++++++++++++---------------- arrow/array/string.go | 16 ++++++++-------- arrow/array/struct.go | 10 +++++----- arrow/array/timestamp.go | 6 +++--- arrow/array/union.go | 24 ++++++++++++------------ arrow/array/util.go | 2 +- arrow/extensions/bool8.go | 4 ++-- arrow/extensions/json.go | 14 +++++++++----- arrow/extensions/uuid.go | 6 +++--- arrow/extensions/variant.go | 4 ++-- 25 files changed, 122 insertions(+), 118 deletions(-) diff --git a/arrow/array.go b/arrow/array.go index d42ca6d05..891697e9b 100644 --- a/arrow/array.go +++ b/arrow/array.go @@ -111,7 +111,7 @@ type Array interface { ValueStr(i int) string // Get single value to be marshalled with `json.Marshal` - GetOneForMarshal(i int) interface{} + GetOneForMarshal(i int, nullable bool) interface{} Data() ArrayData diff --git a/arrow/array/binary.go b/arrow/array/binary.go index 4ddfdbf6c..ba9ce287b 100644 --- a/arrow/array/binary.go +++ b/arrow/array/binary.go @@ -152,8 +152,8 @@ func (a *Binary) setData(data *Data) { } } -func (a *Binary) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *Binary) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } return a.Value(i) @@ -162,7 +162,7 @@ func (a *Binary) GetOneForMarshal(i int) interface{} { func (a *Binary) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string @@ -346,8 +346,8 @@ func (a *LargeBinary) setData(data *Data) { } } -func (a *LargeBinary) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *LargeBinary) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } return a.Value(i) @@ -356,7 +356,7 @@ func (a *LargeBinary) GetOneForMarshal(i int) interface{} { func (a *LargeBinary) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string @@ -522,8 +522,8 @@ func (a *BinaryView) ValueStr(i int) string { return base64.StdEncoding.EncodeToString(a.Value(i)) } -func (a *BinaryView) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *BinaryView) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } return a.Value(i) @@ -532,7 +532,7 @@ func (a *BinaryView) GetOneForMarshal(i int) interface{} { func (a *BinaryView) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string diff --git a/arrow/array/boolean.go b/arrow/array/boolean.go index e57dd3e6e..c555536f5 100644 --- a/arrow/array/boolean.go +++ b/arrow/array/boolean.go @@ -90,8 +90,8 @@ func (a *Boolean) setData(data *Data) { } } -func (a *Boolean) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *Boolean) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.Value(i) } return nil diff --git a/arrow/array/decimal.go b/arrow/array/decimal.go index b4e24f69e..3227f66e3 100644 --- a/arrow/array/decimal.go +++ b/arrow/array/decimal.go @@ -55,7 +55,7 @@ func (a *baseDecimal[T]) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - return a.GetOneForMarshal(i).(string) + return a.GetOneForMarshal(i, true).(string) } func (a *baseDecimal[T]) Values() []T { return a.values } @@ -89,8 +89,8 @@ func (a *baseDecimal[T]) setData(data *Data) { } } -func (a *baseDecimal[T]) GetOneForMarshal(i int) any { - if a.IsNull(i) { +func (a *baseDecimal[T]) GetOneForMarshal(i int, nullable bool) any { + if nullable && a.IsNull(i) { return nil } @@ -102,7 +102,7 @@ func (a *baseDecimal[T]) GetOneForMarshal(i int) any { func (a *baseDecimal[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } diff --git a/arrow/array/decimal128_test.go b/arrow/array/decimal128_test.go index e642d0374..dcb234a63 100644 --- a/arrow/array/decimal128_test.go +++ b/arrow/array/decimal128_test.go @@ -279,7 +279,7 @@ func TestDecimal128GetOneForMarshal(t *testing.T) { } for i := range cases { - assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i), "unexpected value at index %d", i) + assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i, true), "unexpected value at index %d", i) } } diff --git a/arrow/array/decimal256_test.go b/arrow/array/decimal256_test.go index b5674253e..0d6e9f140 100644 --- a/arrow/array/decimal256_test.go +++ b/arrow/array/decimal256_test.go @@ -288,6 +288,6 @@ func TestDecimal256GetOneForMarshal(t *testing.T) { } for i := range cases { - assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i), "unexpected value at index %d", i) + assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i, true), "unexpected value at index %d", i) } } diff --git a/arrow/array/dictionary.go b/arrow/array/dictionary.go index 9645b788e..966131047 100644 --- a/arrow/array/dictionary.go +++ b/arrow/array/dictionary.go @@ -287,18 +287,18 @@ func (d *Dictionary) GetValueIndex(i int) int { return -1 } -func (d *Dictionary) GetOneForMarshal(i int) interface{} { - if d.IsNull(i) { +func (d *Dictionary) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && d.IsNull(i) { return nil } vidx := d.GetValueIndex(i) - return d.Dictionary().GetOneForMarshal(vidx) + return d.Dictionary().GetOneForMarshal(vidx, nullable) } func (d *Dictionary) MarshalJSON() ([]byte, error) { vals := make([]any, d.Len()) for i := range d.Len() { - vals[i] = d.GetOneForMarshal(i) + vals[i] = d.GetOneForMarshal(i, true) } return json.Marshal(vals) } diff --git a/arrow/array/encoded.go b/arrow/array/encoded.go index 4cfc346fe..bcfeded76 100644 --- a/arrow/array/encoded.go +++ b/arrow/array/encoded.go @@ -219,13 +219,13 @@ func (r *RunEndEncoded) String() string { buf.WriteByte(',') } - value := r.values.GetOneForMarshal(i) + value := r.values.GetOneForMarshal(i, true) if byts, ok := value.(json.RawMessage); ok { value = string(byts) } var runEnd int - switch e := r.ends.GetOneForMarshal(i).(type) { + switch e := r.ends.GetOneForMarshal(i, true).(type) { case int16: runEnd = int(e) - r.data.offset case int32: @@ -240,8 +240,8 @@ func (r *RunEndEncoded) String() string { return buf.String() } -func (r *RunEndEncoded) GetOneForMarshal(i int) interface{} { - return r.values.GetOneForMarshal(r.GetPhysicalIndex(i)) +func (r *RunEndEncoded) GetOneForMarshal(i int, nullable bool) interface{} { + return r.values.GetOneForMarshal(r.GetPhysicalIndex(i), nullable) } func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { @@ -252,7 +252,7 @@ func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(r.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(r.GetOneForMarshal(i, true)); err != nil { return nil, err } } diff --git a/arrow/array/extension.go b/arrow/array/extension.go index 21ed3b249..48c4f03ad 100644 --- a/arrow/array/extension.go +++ b/arrow/array/extension.go @@ -116,8 +116,8 @@ func (e *ExtensionArrayBase) String() string { return fmt.Sprintf("(%s)%s", e.data.dtype, e.storage) } -func (e *ExtensionArrayBase) GetOneForMarshal(i int) interface{} { - return e.storage.GetOneForMarshal(i) +func (e *ExtensionArrayBase) GetOneForMarshal(i int, nullable bool) interface{} { + return e.storage.GetOneForMarshal(i, nullable) } func (e *ExtensionArrayBase) MarshalJSON() ([]byte, error) { diff --git a/arrow/array/fixed_size_list.go b/arrow/array/fixed_size_list.go index 4b3ff286f..10dce4c74 100644 --- a/arrow/array/fixed_size_list.go +++ b/arrow/array/fixed_size_list.go @@ -51,7 +51,7 @@ func (a *FixedSizeList) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - return string(a.GetOneForMarshal(i).(json.RawMessage)) + return string(a.GetOneForMarshal(i, true).(json.RawMessage)) } func (a *FixedSizeList) String() string { @@ -123,8 +123,8 @@ func (a *FixedSizeList) Release() { a.values.Release() } -func (a *FixedSizeList) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *FixedSizeList) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } slice := a.newListValue(i) diff --git a/arrow/array/fixedsize_binary.go b/arrow/array/fixedsize_binary.go index 44ca44883..b9cf2bfdc 100644 --- a/arrow/array/fixedsize_binary.go +++ b/arrow/array/fixedsize_binary.go @@ -86,8 +86,8 @@ func (a *FixedSizeBinary) setData(data *Data) { } } -func (a *FixedSizeBinary) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *FixedSizeBinary) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } diff --git a/arrow/array/float16.go b/arrow/array/float16.go index 41276803b..8536df402 100644 --- a/arrow/array/float16.go +++ b/arrow/array/float16.go @@ -77,8 +77,8 @@ func (a *Float16) setData(data *Data) { } } -func (a *Float16) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *Float16) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.values[i].Float32() } return nil diff --git a/arrow/array/interval.go b/arrow/array/interval.go index 128cdfadf..b2aad56f1 100644 --- a/arrow/array/interval.go +++ b/arrow/array/interval.go @@ -94,8 +94,8 @@ func (a *MonthInterval) setData(data *Data) { } } -func (a *MonthInterval) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *MonthInterval) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.values[i] } return nil @@ -361,7 +361,7 @@ func (a *DayTimeInterval) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i)) + data, err := json.Marshal(a.GetOneForMarshal(i, true)) if err != nil { panic(err) } @@ -399,8 +399,8 @@ func (a *DayTimeInterval) setData(data *Data) { } } -func (a *DayTimeInterval) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *DayTimeInterval) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.values[i] } return nil @@ -663,7 +663,7 @@ func (a *MonthDayNanoInterval) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i)) + data, err := json.Marshal(a.GetOneForMarshal(i, true)) if err != nil { panic(err) } @@ -703,8 +703,8 @@ func (a *MonthDayNanoInterval) setData(data *Data) { } } -func (a *MonthDayNanoInterval) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *MonthDayNanoInterval) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.values[i] } return nil diff --git a/arrow/array/list.go b/arrow/array/list.go index df45223f0..16eac56f9 100644 --- a/arrow/array/list.go +++ b/arrow/array/list.go @@ -61,7 +61,7 @@ func (a *List) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i).(json.RawMessage)) + return string(a.GetOneForMarshal(i, true).(json.RawMessage)) } func (a *List) String() string { @@ -98,8 +98,8 @@ func (a *List) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *List) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *List) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } @@ -121,7 +121,7 @@ func (a *List) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } @@ -193,7 +193,7 @@ func (a *LargeList) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i).(json.RawMessage)) + return string(a.GetOneForMarshal(i, true).(json.RawMessage)) } func (a *LargeList) String() string { @@ -230,8 +230,8 @@ func (a *LargeList) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *LargeList) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *LargeList) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } @@ -253,7 +253,7 @@ func (a *LargeList) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } @@ -664,7 +664,7 @@ func (a *ListView) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i).(json.RawMessage)) + return string(a.GetOneForMarshal(i, true).(json.RawMessage)) } func (a *ListView) String() string { @@ -705,8 +705,8 @@ func (a *ListView) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *ListView) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *ListView) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } @@ -728,7 +728,7 @@ func (a *ListView) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } @@ -811,7 +811,7 @@ func (a *LargeListView) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i).(json.RawMessage)) + return string(a.GetOneForMarshal(i, true).(json.RawMessage)) } func (a *LargeListView) String() string { @@ -852,8 +852,8 @@ func (a *LargeListView) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *LargeListView) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *LargeListView) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } @@ -875,7 +875,7 @@ func (a *LargeListView) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } diff --git a/arrow/array/null.go b/arrow/array/null.go index 8f8f58056..01001776c 100644 --- a/arrow/array/null.go +++ b/arrow/array/null.go @@ -80,7 +80,7 @@ func (a *Null) setData(data *Data) { a.data.nulls = a.data.length } -func (a *Null) GetOneForMarshal(i int) interface{} { +func (a *Null) GetOneForMarshal(i int, nullable bool) interface{} { return nil } diff --git a/arrow/array/numeric_generic.go b/arrow/array/numeric_generic.go index 49a369cdd..c7616cf20 100644 --- a/arrow/array/numeric_generic.go +++ b/arrow/array/numeric_generic.go @@ -83,8 +83,8 @@ func (a *numericArray[T]) ValueStr(i int) string { return fmt.Sprintf("%v", a.values[i]) } -func (a *numericArray[T]) GetOneForMarshal(i int) any { - if a.IsNull(i) { +func (a *numericArray[T]) GetOneForMarshal(i int, nullable bool) any { + if nullable && a.IsNull(i) { return nil } @@ -107,8 +107,8 @@ type oneByteArrs[T int8 | uint8] struct { numericArray[T] } -func (a *oneByteArrs[T]) GetOneForMarshal(i int) any { - if a.IsNull(i) { +func (a *oneByteArrs[T]) GetOneForMarshal(i int, nullable bool) any { + if nullable && a.IsNull(i) { return nil } @@ -141,8 +141,8 @@ func (a *floatArray[T]) ValueStr(i int) string { return strconv.FormatFloat(float64(a.Value(i)), 'g', -1, bitWidth) } -func (a *floatArray[T]) GetOneForMarshal(i int) any { - if a.IsNull(i) { +func (a *floatArray[T]) GetOneForMarshal(i int, nullable bool) any { + if nullable && a.IsNull(i) { return nil } @@ -160,7 +160,7 @@ func (a *floatArray[T]) GetOneForMarshal(i int) any { func (a *floatArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } @@ -176,7 +176,7 @@ type dateArray[T interface { func (d *dateArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, d.Len()) for i := range d.values { - vals[i] = d.GetOneForMarshal(i) + vals[i] = d.GetOneForMarshal(i, true) } return json.Marshal(vals) } @@ -189,8 +189,8 @@ func (d *dateArray[T]) ValueStr(i int) string { return d.values[i].FormattedString() } -func (d *dateArray[T]) GetOneForMarshal(i int) interface{} { - if d.IsNull(i) { +func (d *dateArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && d.IsNull(i) { return nil } @@ -212,7 +212,7 @@ type timeArray[T interface { func (a *timeArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } @@ -225,8 +225,8 @@ func (a *timeArray[T]) ValueStr(i int) string { return a.values[i].FormattedString(a.DataType().(timeType).TimeUnit()) } -func (a *timeArray[T]) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *timeArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } @@ -248,7 +248,7 @@ func (a *Duration) DurationValues() []arrow.Duration { return a.Values() } func (a *Duration) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } @@ -261,8 +261,8 @@ func (a *Duration) ValueStr(i int) string { return fmt.Sprintf("%d%s", a.values[i], a.DataType().(timeType).TimeUnit()) } -func (a *Duration) GetOneForMarshal(i int) any { - if a.IsNull(i) { +func (a *Duration) GetOneForMarshal(i int, nullable bool) any { + if nullable && a.IsNull(i) { return nil } return fmt.Sprintf("%d%s", a.values[i], a.DataType().(timeType).TimeUnit()) diff --git a/arrow/array/string.go b/arrow/array/string.go index f0d9be401..e3ed07236 100644 --- a/arrow/array/string.go +++ b/arrow/array/string.go @@ -154,8 +154,8 @@ func (a *String) setData(data *Data) { } } -func (a *String) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *String) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.Value(i) } return nil @@ -362,8 +362,8 @@ func (a *LargeString) setData(data *Data) { } } -func (a *LargeString) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *LargeString) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.Value(i) } return nil @@ -372,7 +372,7 @@ func (a *LargeString) GetOneForMarshal(i int) interface{} { func (a *LargeString) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } @@ -526,8 +526,8 @@ func (a *StringView) ValueStr(i int) string { return a.Value(i) } -func (a *StringView) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *StringView) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } return a.Value(i) @@ -536,7 +536,7 @@ func (a *StringView) GetOneForMarshal(i int) interface{} { func (a *StringView) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 76611dca0..5feb5faec 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -130,7 +130,7 @@ func (a *Struct) ValueStr(i int) string { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i)) + data, err := json.Marshal(a.GetOneForMarshal(i, true)) if err != nil { panic(err) } @@ -209,15 +209,15 @@ func (a *Struct) setData(data *Data) { } } -func (a *Struct) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *Struct) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } tmp := make(map[string]interface{}) fieldList := a.data.dtype.(*arrow.StructType).Fields() for j, d := range a.fields { - tmp[fieldList[j].Name] = d.GetOneForMarshal(i) + tmp[fieldList[j].Name] = d.GetOneForMarshal(i, true) } return tmp } @@ -231,7 +231,7 @@ func (a *Struct) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } diff --git a/arrow/array/timestamp.go b/arrow/array/timestamp.go index 55e9b5358..613941f7f 100644 --- a/arrow/array/timestamp.go +++ b/arrow/array/timestamp.go @@ -109,8 +109,8 @@ func (a *Timestamp) ValueStr(i int) string { return toTime(a.values[i]).Format(layout) } -func (a *Timestamp) GetOneForMarshal(i int) interface{} { - if val := a.ValueStr(i); val != NullValueStr { +func (a *Timestamp) GetOneForMarshal(i int, nullable bool) interface{} { + if val := a.ValueStr(i); !nullable || val != NullValueStr { return val } return nil @@ -119,7 +119,7 @@ func (a *Timestamp) GetOneForMarshal(i int) interface{} { func (a *Timestamp) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) diff --git a/arrow/array/union.go b/arrow/array/union.go index 58524a768..1d40ad66b 100644 --- a/arrow/array/union.go +++ b/arrow/array/union.go @@ -320,17 +320,17 @@ func (a *SparseUnion) setData(data *Data) { debug.Assert(a.data.buffers[0] == nil, "arrow/array: validity bitmap for sparse unions should be nil") } -func (a *SparseUnion) GetOneForMarshal(i int) interface{} { +func (a *SparseUnion) GetOneForMarshal(i int, nullable bool) interface{} { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) data := a.Field(childID) - if data.IsNull(i) { + if nullable && data.IsNull(i) { return nil } - return []interface{}{typeID, data.GetOneForMarshal(i)} + return []interface{}{typeID, data.GetOneForMarshal(i, nullable)} } func (a *SparseUnion) MarshalJSON() ([]byte, error) { @@ -342,7 +342,7 @@ func (a *SparseUnion) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } @@ -355,7 +355,7 @@ func (a *SparseUnion) ValueStr(i int) string { return NullValueStr } - val := a.GetOneForMarshal(i) + val := a.GetOneForMarshal(i, true) if val == nil { // child is nil return NullValueStr @@ -380,7 +380,7 @@ func (a *SparseUnion) String() string { field := fieldList[a.ChildID(i)] f := a.Field(a.ChildID(i)) - fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(i)) + fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(i, true)) } b.WriteByte(']') return b.String() @@ -613,18 +613,18 @@ func (a *DenseUnion) setData(data *Data) { } } -func (a *DenseUnion) GetOneForMarshal(i int) interface{} { +func (a *DenseUnion) GetOneForMarshal(i int, nullable bool) interface{} { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) data := a.Field(childID) offset := int(a.RawValueOffsets()[i]) - if data.IsNull(offset) { + if nullable && data.IsNull(offset) { return nil } - return []interface{}{typeID, data.GetOneForMarshal(offset)} + return []interface{}{typeID, data.GetOneForMarshal(offset, nullable)} } func (a *DenseUnion) MarshalJSON() ([]byte, error) { @@ -636,7 +636,7 @@ func (a *DenseUnion) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { return nil, err } } @@ -649,7 +649,7 @@ func (a *DenseUnion) ValueStr(i int) string { return NullValueStr } - val := a.GetOneForMarshal(i) + val := a.GetOneForMarshal(i, true) if val == nil { // child in nil return NullValueStr @@ -676,7 +676,7 @@ func (a *DenseUnion) String() string { field := fieldList[a.ChildID(i)] f := a.Field(a.ChildID(i)) - fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(int(offsets[i]))) + fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(int(offsets[i]), true)) } b.WriteByte(']') return b.String() diff --git a/arrow/array/util.go b/arrow/array/util.go index 136ed3537..8fc61a6c2 100644 --- a/arrow/array/util.go +++ b/arrow/array/util.go @@ -284,7 +284,7 @@ func RecordToJSON(rec arrow.RecordBatch, w io.Writer) error { cols := make(map[string]interface{}) for i := 0; int64(i) < rec.NumRows(); i++ { for j, c := range rec.Columns() { - cols[fields[j].Name] = c.GetOneForMarshal(i) + cols[fields[j].Name] = c.GetOneForMarshal(i, true) } if err := enc.Encode(cols); err != nil { return err diff --git a/arrow/extensions/bool8.go b/arrow/extensions/bool8.go index 97038a1bf..aaf51f0c5 100644 --- a/arrow/extensions/bool8.go +++ b/arrow/extensions/bool8.go @@ -114,8 +114,8 @@ func (a *Bool8Array) MarshalJSON() ([]byte, error) { return json.Marshal(values) } -func (a *Bool8Array) GetOneForMarshal(i int) interface{} { - if a.IsNull(i) { +func (a *Bool8Array) GetOneForMarshal(i int, nullable bool) interface{} { + if nullable && a.IsNull(i) { return nil } return a.Value(i) diff --git a/arrow/extensions/json.go b/arrow/extensions/json.go index 3f46b50e3..b9cc50a73 100644 --- a/arrow/extensions/json.go +++ b/arrow/extensions/json.go @@ -116,9 +116,7 @@ func (a *JSONArray) ValueBytes(i int) []byte { return b } -// ValueJSON wraps the underlying string value as a json.RawMessage, -// or returns nil if the array value is null. -func (a *JSONArray) ValueJSON(i int) json.RawMessage { +func (a *JSONArray) valueJSON(i int, nullable bool) json.RawMessage { var val json.RawMessage if a.IsValid(i) { val = json.RawMessage(a.Storage().(array.StringLike).Value(i)) @@ -126,6 +124,12 @@ func (a *JSONArray) ValueJSON(i int) json.RawMessage { return val } +// ValueJSON wraps the underlying string value as a json.RawMessage, +// or returns nil if the array value is null. +func (a *JSONArray) ValueJSON(i int) json.RawMessage { + return a.valueJSON(i, true) +} + // MarshalJSON implements json.Marshaler. // Marshaling json.RawMessage is a no-op, except that nil values will // be marshaled as a JSON null. @@ -138,8 +142,8 @@ func (a *JSONArray) MarshalJSON() ([]byte, error) { } // GetOneForMarshal implements arrow.Array. -func (a *JSONArray) GetOneForMarshal(i int) interface{} { - return a.ValueJSON(i) +func (a *JSONArray) GetOneForMarshal(i int, nullable bool) interface{} { + return a.valueJSON(i, nullable) } var ( diff --git a/arrow/extensions/uuid.go b/arrow/extensions/uuid.go index 9aac02253..cbe56ecd6 100644 --- a/arrow/extensions/uuid.go +++ b/arrow/extensions/uuid.go @@ -190,13 +190,13 @@ func (a *UUIDArray) ValueStr(i int) string { func (a *UUIDArray) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range vals { - vals[i] = a.GetOneForMarshal(i) + vals[i] = a.GetOneForMarshal(i, true) } return json.Marshal(vals) } -func (a *UUIDArray) GetOneForMarshal(i int) interface{} { - if a.IsValid(i) { +func (a *UUIDArray) GetOneForMarshal(i int, nullable bool) interface{} { + if !nullable || a.IsValid(i) { return a.Value(i) } return nil diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index 379822c4b..1f41ef076 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -599,8 +599,8 @@ func (v *VariantArray) MarshalJSON() ([]byte, error) { return json.Marshal(values) } -func (v *VariantArray) GetOneForMarshal(i int) any { - if v.IsNull(i) { +func (v *VariantArray) GetOneForMarshal(i int, nullable bool) any { + if nullable && v.IsNull(i) { return nil } From 6fb44389dd2f495c56501f1b19ec56175efd5917 Mon Sep 17 00:00:00 2001 From: serramatutu Date: Wed, 29 Apr 2026 18:24:30 +0200 Subject: [PATCH 07/18] Make struct and record use `field.Nullable` when serializing --- arrow/array/struct.go | 5 +++-- arrow/array/util.go | 2 +- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 5feb5faec..63ccbf599 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -215,9 +215,10 @@ func (a *Struct) GetOneForMarshal(i int, nullable bool) interface{} { } tmp := make(map[string]interface{}) - fieldList := a.data.dtype.(*arrow.StructType).Fields() + dtype := a.data.dtype.(*arrow.StructType) + fieldList := dtype.Fields() for j, d := range a.fields { - tmp[fieldList[j].Name] = d.GetOneForMarshal(i, true) + tmp[fieldList[j].Name] = d.GetOneForMarshal(i, dtype.Field(j).Nullable) } return tmp } diff --git a/arrow/array/util.go b/arrow/array/util.go index 8fc61a6c2..b3375152a 100644 --- a/arrow/array/util.go +++ b/arrow/array/util.go @@ -284,7 +284,7 @@ func RecordToJSON(rec arrow.RecordBatch, w io.Writer) error { cols := make(map[string]interface{}) for i := 0; int64(i) < rec.NumRows(); i++ { for j, c := range rec.Columns() { - cols[fields[j].Name] = c.GetOneForMarshal(i, true) + cols[fields[j].Name] = c.GetOneForMarshal(i, rec.Schema().Field(j).Nullable) } if err := enc.Encode(cols); err != nil { return err From ad78a63378a50b4857600442109077d8f9f7ef13 Mon Sep 17 00:00:00 2001 From: serramatutu Date: Wed, 29 Apr 2026 18:25:58 +0200 Subject: [PATCH 08/18] Fix tests that were implicitly depending on wrong nullable semantics This commit fixes some tests that were implicitly setting `valid=false` on non-nullable fields, which now causes legitimate test failures when roundtripping to JSON. --- arrow/compute/vector_sort_test.go | 46 +++++++++++++------------- arrow/extensions/uuid_test.go | 2 +- arrow/internal/arrdata/arrdata.go | 46 +++++++++++++------------- arrow/internal/arrjson/arrjson_test.go | 4 +-- arrow/ipc/cmd/arrow-ls/main_test.go | 6 ++-- 5 files changed, 52 insertions(+), 52 deletions(-) diff --git a/arrow/compute/vector_sort_test.go b/arrow/compute/vector_sort_test.go index 39bf5e95f..5a15428e7 100644 --- a/arrow/compute/vector_sort_test.go +++ b/arrow/compute/vector_sort_test.go @@ -1349,8 +1349,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("NoNull", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Uint8}, - {Name: "b", Type: arrow.PrimitiveTypes.Uint32}, + {Name: "a", Type: arrow.PrimitiveTypes.Uint8, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Uint32, Nullable: true}, }, nil) jsonRows := `[ {"a": 3, "b": 5}, @@ -1373,8 +1373,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("Null", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Uint8}, - {Name: "b", Type: arrow.PrimitiveTypes.Uint32}, + {Name: "a", Type: arrow.PrimitiveTypes.Uint8, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Uint32, Nullable: true}, }, nil) jsonRows := `[ {"a": null, "b": 5}, @@ -1396,8 +1396,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("NaN", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Float32}, - {Name: "b", Type: arrow.PrimitiveTypes.Float64}, + {Name: "a", Type: arrow.PrimitiveTypes.Float32, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Float64, Nullable: true}, }, nil) ba := array.NewFloat32Builder(mem) defer ba.Release() @@ -1426,8 +1426,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("NaNAndNull", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Float32}, - {Name: "b", Type: arrow.PrimitiveTypes.Float64}, + {Name: "a", Type: arrow.PrimitiveTypes.Float32, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Float64, Nullable: true}, }, nil) ba := array.NewFloat32Builder(mem) defer ba.Release() @@ -1460,8 +1460,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("Boolean", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.FixedWidthTypes.Boolean}, - {Name: "b", Type: arrow.FixedWidthTypes.Boolean}, + {Name: "a", Type: arrow.FixedWidthTypes.Boolean, Nullable: true}, + {Name: "b", Type: arrow.FixedWidthTypes.Boolean, Nullable: true}, }, nil) jsonRows := `[ {"a": true, "b": null}, @@ -1486,9 +1486,9 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { ts := &arrow.TimestampType{Unit: arrow.Microsecond} fsb3 := &arrow.FixedSizeBinaryType{ByteWidth: 3} schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: ts}, - {Name: "b", Type: arrow.BinaryTypes.LargeString}, - {Name: "c", Type: fsb3}, + {Name: "a", Type: ts, Nullable: true}, + {Name: "b", Type: arrow.BinaryTypes.LargeString, Nullable: true}, + {Name: "c", Type: fsb3, Nullable: true}, }, nil) ba := array.NewTimestampBuilder(mem, ts) defer ba.Release() @@ -1535,8 +1535,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { d128 := &arrow.Decimal128Type{Precision: 3, Scale: 1} d256 := &arrow.Decimal256Type{Precision: 4, Scale: 2} schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: d128}, - {Name: "b", Type: d256}, + {Name: "a", Type: d128, Nullable: true}, + {Name: "b", Type: d256, Nullable: true}, }, nil) jsonRows := `[ {"a": "12.3", "b": "12.34"}, @@ -1561,8 +1561,8 @@ func TestVectorSortIndicesCppRecordBatchParity(t *testing.T) { t.Run("DuplicateSortKeys", func(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Float32}, - {Name: "b", Type: arrow.PrimitiveTypes.Float64}, + {Name: "a", Type: arrow.PrimitiveTypes.Float32, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Float64, Nullable: true}, }, nil) ba := array.NewFloat32Builder(mem) defer ba.Release() @@ -1610,8 +1610,8 @@ func TestVectorSortIndicesCppTableParity(t *testing.T) { ctx := context.Background() schemaAB := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Uint8}, - {Name: "b", Type: arrow.PrimitiveTypes.Uint32}, + {Name: "a", Type: arrow.PrimitiveTypes.Uint8, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Uint32, Nullable: true}, }, nil) t.Run("EmptyTable", func(t *testing.T) { @@ -1667,8 +1667,8 @@ func TestVectorSortIndicesCppTableParity(t *testing.T) { t.Run("BinaryLikeTwoChunks", func(t *testing.T) { fsb3 := &arrow.FixedSizeBinaryType{ByteWidth: 3} s := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.BinaryTypes.LargeString}, - {Name: "b", Type: fsb3}, + {Name: "a", Type: arrow.BinaryTypes.LargeString, Nullable: true}, + {Name: "b", Type: fsb3, Nullable: true}, }, nil) buildBatch := func(a []string, b [][]byte, bNulls []bool) arrow.RecordBatch { ab := array.NewLargeStringBuilder(mem) @@ -1719,8 +1719,8 @@ func TestVectorSortIndicesCppTableParity(t *testing.T) { t.Run("HeterogenousChunking", func(t *testing.T) { s := arrow.NewSchema([]arrow.Field{ - {Name: "a", Type: arrow.PrimitiveTypes.Float32}, - {Name: "b", Type: arrow.PrimitiveTypes.Float64}, + {Name: "a", Type: arrow.PrimitiveTypes.Float32, Nullable: true}, + {Name: "b", Type: arrow.PrimitiveTypes.Float64, Nullable: true}, }, nil) a0, _, err := array.FromJSON(mem, arrow.PrimitiveTypes.Float32, strings.NewReader("[null, 1]")) require.NoError(t, err) diff --git a/arrow/extensions/uuid_test.go b/arrow/extensions/uuid_test.go index a76b77a91..36bb812b9 100644 --- a/arrow/extensions/uuid_test.go +++ b/arrow/extensions/uuid_test.go @@ -62,7 +62,7 @@ func TestUUIDExtensionBuilder(t *testing.T) { func TestUUIDExtensionRecordBuilder(t *testing.T) { schema := arrow.NewSchema([]arrow.Field{ - {Name: "uuid", Type: extensions.NewUUIDType()}, + {Name: "uuid", Type: extensions.NewUUIDType(), Nullable: true}, }, nil) builder := array.NewRecordBuilder(memory.DefaultAllocator, schema) builder.Field(0).(*extensions.UUIDBuilder).Append(testUUID) diff --git a/arrow/internal/arrdata/arrdata.go b/arrow/internal/arrdata/arrdata.go index 095571a8f..d95ee7cc7 100644 --- a/arrow/internal/arrdata/arrdata.go +++ b/arrow/internal/arrdata/arrdata.go @@ -192,59 +192,59 @@ func makeStructsRecords() []arrow.RecordBatch { mem := memory.NewGoAllocator() fields := []arrow.Field{ - {Name: "f1", Type: arrow.PrimitiveTypes.Int32}, - {Name: "f2", Type: arrow.BinaryTypes.String}, + {Name: "f1", Type: arrow.PrimitiveTypes.Int32, Nullable: true}, + {Name: "f2", Type: arrow.BinaryTypes.String, Nullable: true}, } dtype := arrow.StructOf(fields...) schema := arrow.NewSchema([]arrow.Field{{Name: "struct_nullable", Type: dtype, Nullable: true}}, nil) - mask := []bool{true, false, false, true, true, true, false, true} + innerValids := []bool{true, false, false, true, true} chunks := [][]arrow.Array{ { structOf(mem, dtype, [][]arrow.Array{ { - arrayOf(mem, []int32{-1, -2, -3, -4, -5}, mask[:5]), - arrayOf(mem, []string{"111", "222", "333", "444", "555"}, mask[:5]), + arrayOf(mem, []int32{-1, -2, -3, -4, -5}, innerValids), + arrayOf(mem, []string{"111", "222", "333", "444", "555"}, innerValids), }, { - arrayOf(mem, []int32{-11, -12, -13, -14, -15}, mask[:5]), - arrayOf(mem, []string{"1111", "1222", "1333", "1444", "1555"}, mask[:5]), + arrayOf(mem, []int32{-11, -12, -13, -14, -15}, innerValids), + arrayOf(mem, []string{"1111", "1222", "1333", "1444", "1555"}, innerValids), }, { - arrayOf(mem, []int32{-21, -22, -23, -24, -25}, mask[:5]), - arrayOf(mem, []string{"2111", "2222", "2333", "2444", "2555"}, mask[:5]), + arrayOf(mem, []int32{-21, -22, -23, -24, -25}, innerValids), + arrayOf(mem, []string{"2111", "2222", "2333", "2444", "2555"}, innerValids), }, { - arrayOf(mem, []int32{-31, -32, -33, -34, -35}, mask[:5]), - arrayOf(mem, []string{"3111", "3222", "3333", "3444", "3555"}, mask[:5]), + arrayOf(mem, []int32{-31, -32, -33, -34, -35}, innerValids), + arrayOf(mem, []string{"3111", "3222", "3333", "3444", "3555"}, innerValids), }, { - arrayOf(mem, []int32{-41, -42, -43, -44, -45}, mask[:5]), - arrayOf(mem, []string{"4111", "4222", "4333", "4444", "4555"}, mask[:5]), + arrayOf(mem, []int32{-41, -42, -43, -44, -45}, innerValids), + arrayOf(mem, []string{"4111", "4222", "4333", "4444", "4555"}, innerValids), }, }, []bool{true, false, true, true, true}), }, { structOf(mem, dtype, [][]arrow.Array{ { - arrayOf(mem, []int32{1, 2, 3, 4, 5}, mask[:5]), - arrayOf(mem, []string{"-111", "-222", "-333", "-444", "-555"}, mask[:5]), + arrayOf(mem, []int32{1, 2, 3, 4, 5}, innerValids), + arrayOf(mem, []string{"-111", "-222", "-333", "-444", "-555"}, innerValids), }, { - arrayOf(mem, []int32{11, 12, 13, 14, 15}, mask[:5]), - arrayOf(mem, []string{"-1111", "-1222", "-1333", "-1444", "-1555"}, mask[:5]), + arrayOf(mem, []int32{11, 12, 13, 14, 15}, innerValids), + arrayOf(mem, []string{"-1111", "-1222", "-1333", "-1444", "-1555"}, innerValids), }, { - arrayOf(mem, []int32{21, 22, 23, 24, 25}, mask[:5]), - arrayOf(mem, []string{"-2111", "-2222", "-2333", "-2444", "-2555"}, mask[:5]), + arrayOf(mem, []int32{21, 22, 23, 24, 25}, innerValids), + arrayOf(mem, []string{"-2111", "-2222", "-2333", "-2444", "-2555"}, innerValids), }, { - arrayOf(mem, []int32{31, 32, 33, 34, 35}, mask[:5]), - arrayOf(mem, []string{"-3111", "-3222", "-3333", "-3444", "-3555"}, mask[:5]), + arrayOf(mem, []int32{31, 32, 33, 34, 35}, innerValids), + arrayOf(mem, []string{"-3111", "-3222", "-3333", "-3444", "-3555"}, innerValids), }, { - arrayOf(mem, []int32{41, 42, 43, 44, 45}, mask[:5]), - arrayOf(mem, []string{"-4111", "-4222", "-4333", "-4444", "-4555"}, mask[:5]), + arrayOf(mem, []int32{41, 42, 43, 44, 45}, innerValids), + arrayOf(mem, []string{"-4111", "-4222", "-4333", "-4444", "-4555"}, innerValids), }, }, []bool{true, false, false, true, true}), }, diff --git a/arrow/internal/arrjson/arrjson_test.go b/arrow/internal/arrjson/arrjson_test.go index 7e2f386fc..faeecfc01 100644 --- a/arrow/internal/arrjson/arrjson_test.go +++ b/arrow/internal/arrjson/arrjson_test.go @@ -948,7 +948,7 @@ func makeStructsWantJSONs() string { "isSigned": true, "bitWidth": 32 }, - "nullable": false, + "nullable": true, "children": [] }, { @@ -956,7 +956,7 @@ func makeStructsWantJSONs() string { "type": { "name": "utf8" }, - "nullable": false, + "nullable": true, "children": [] } ] diff --git a/arrow/ipc/cmd/arrow-ls/main_test.go b/arrow/ipc/cmd/arrow-ls/main_test.go index f90e4a800..0f3b5377c 100644 --- a/arrow/ipc/cmd/arrow-ls/main_test.go +++ b/arrow/ipc/cmd/arrow-ls/main_test.go @@ -59,7 +59,7 @@ records: 3 name: "structs", want: `schema: fields: 1 - - struct_nullable: type=struct, nullable + - struct_nullable: type=struct, nullable records: 2 `, }, @@ -221,7 +221,7 @@ records: 3 name: "structs", want: `schema: fields: 1 - - struct_nullable: type=struct, nullable + - struct_nullable: type=struct, nullable records: 2 `, }, @@ -230,7 +230,7 @@ records: 2 want: `version: V5 schema: fields: 1 - - struct_nullable: type=struct, nullable + - struct_nullable: type=struct, nullable records: 2 `, }, From 0cbfe2acdba93451ab92157fcf6ff0018574df8c Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 13:58:11 -0400 Subject: [PATCH 09/18] Post-rebase fixes for CI: JSON int precision, GetOneForMarshal, TinyGo record.go/struct.go: call UseNumber() on the per-field sub-decoder in the nullable decode path so large int64 keeps full precision (regression against #816, surfaced by the rebase). json_reader_test.go: pass the nullable flag to the two-arg GetOneForMarshal. internal/json/json_stdlib.go: add IsNullMessage to the tinygo/stdlib json variant (it was only added to the goccy build), fixing the TinyGo 'undefined: json.IsNullMessage' example build. --- arrow/array/json_reader_test.go | 2 +- arrow/array/record.go | 1 + arrow/array/struct.go | 1 + internal/json/json_stdlib.go | 5 +++++ 4 files changed, 8 insertions(+), 1 deletion(-) diff --git a/arrow/array/json_reader_test.go b/arrow/array/json_reader_test.go index 3d0def659..fc7cedea9 100644 --- a/arrow/array/json_reader_test.go +++ b/arrow/array/json_reader_test.go @@ -253,7 +253,7 @@ func recordBatchToNDJSON(t *testing.T, rec arrow.RecordBatch) string { defer arr.Release() for pos := range arr.Len() { - s, err := json.Marshal(arr.GetOneForMarshal(pos)) + s, err := json.Marshal(arr.GetOneForMarshal(pos, true)) assert.NoError(t, err) sb.Write(s) sb.WriteByte('\n') diff --git a/arrow/array/record.go b/arrow/array/record.go index 2f8948d4c..4b7b17bd4 100644 --- a/arrow/array/record.go +++ b/arrow/array/record.go @@ -466,6 +466,7 @@ func (b *RecordBuilder) UnmarshalOne(dec *json.Decoder) error { b.fields[idx].AppendEmptyValue() } else { sub := json.NewDecoder(bytes.NewReader(next)) + sub.UseNumber() if err := b.fields[idx].UnmarshalOne(sub); err != nil { return err } diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 63ccbf599..17f6152a3 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -498,6 +498,7 @@ func (b *StructBuilder) UnmarshalOne(dec *json.Decoder) error { b.fields[idx].AppendEmptyValue() } else { sub := json.NewDecoder(bytes.NewReader(next)) + sub.UseNumber() if err := b.fields[idx].UnmarshalOne(sub); err != nil { return err } diff --git a/internal/json/json_stdlib.go b/internal/json/json_stdlib.go index 3031029d8..06071fada 100644 --- a/internal/json/json_stdlib.go +++ b/internal/json/json_stdlib.go @@ -20,6 +20,7 @@ package json import ( + "bytes" "io" "encoding/json" @@ -49,3 +50,7 @@ func NewDecoder(r io.Reader) *Decoder { func NewEncoder(w io.Writer) *Encoder { return json.NewEncoder(w) } + +func IsNullMessage(m RawMessage) bool { + return bytes.Equal(m, []byte("null")) +} From 719f973951c7abea09742b07da50fbe2e849ff24 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 15:34:50 -0400 Subject: [PATCH 10/18] Make nested nullability field-local in comparison and union marshal A non-nullable list/union/struct can legitimately hold nullable children, so equality and JSON serialization must honor each child/element field's own nullability rather than propagating the parent's. Previously the parent nullability was pushed down, so null child slots were compared and serialized by their arbitrary underlying bytes -- which round-trips inconsistently and breaks cross-language integration (union, nested_large_offsets) even when values are logically equal. compare.go: unions and run-end-encoded arrays have no top-level validity bitmap, so skip their top-level null-count/validity checks (and the all-null shortcut); recurse into list/struct/map/union children with the child field's nullability. union.go: SparseUnion/DenseUnion GetOneForMarshal use the selected child field's nullability and emit [typeID, null] for a null child so it round-trips as null. --- arrow/array/compare.go | 55 +++++++++++++++++++++++++++++++++--------- arrow/array/union.go | 24 ++++++++++-------- 2 files changed, 57 insertions(+), 22 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index ec56cb6d7..07f475bf7 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -253,7 +253,7 @@ func equal(left, right arrow.Array, opt equalOption) bool { return false case left.Len() == 0: return true - case opt.nullable && left.NullN() == left.Len(): + case opt.nullable && hasTopLevelValidityBitmap(left.DataType().ID()) && left.NullN() == left.Len(): return true } @@ -527,7 +527,7 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { return false case left.Len() == 0: return true - case opt.nullable && left.NullN() == left.Len(): + case opt.nullable && hasTopLevelValidityBitmap(left.DataType().ID()) && left.NullN() == left.Len(): return true } @@ -678,14 +678,34 @@ func arrayApproxEqual(left, right arrow.Array, opt equalOption) bool { } } +func withNullable(opt equalOption, nullable bool) equalOption { + opt.nullable = nullable + return opt +} + +// hasTopLevelValidityBitmap reports whether arrays of the given type id carry +// their own top-level validity bitmap. Union and run-end-encoded arrays do not: +// their nullness is encoded entirely in their children, so top-level null-count +// and validity comparisons are not meaningful for them and can legitimately +// differ between logically-equal arrays (e.g. built from JSON vs. read from IPC). +func hasTopLevelValidityBitmap(id arrow.Type) bool { + switch id { + case arrow.SPARSE_UNION, arrow.DENSE_UNION, arrow.RUN_END_ENCODED: + return false + } + return true +} + func baseArrayEqual(left, right arrow.Array, opt equalOption) bool { switch { case left.Len() != right.Len(): return false - case opt.nullable && left.NullN() != right.NullN(): - return false case !arrow.TypeEqual(left.DataType(), right.DataType()): // We do not check for metadata as in the C++ implementation. return false + case !hasTopLevelValidityBitmap(left.DataType().ID()): + return true + case opt.nullable && left.NullN() != right.NullN(): + return false case opt.nullable && !validityBitmapEqual(left, right): return false } @@ -779,6 +799,7 @@ func arrayApproxEqualFloat64(left, right *Float64, opt equalOption) bool { } func arrayApproxEqualList(left, right *List, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -788,7 +809,7 @@ func arrayApproxEqualList(left, right *List, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return arrayApproxEqual(l, r, opt) + return arrayApproxEqual(l, r, childOpt) }() if !o { return false @@ -798,6 +819,7 @@ func arrayApproxEqualList(left, right *List, opt equalOption) bool { } func arrayApproxEqualLargeList(left, right *LargeList, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -807,7 +829,7 @@ func arrayApproxEqualLargeList(left, right *LargeList, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return arrayApproxEqual(l, r, opt) + return arrayApproxEqual(l, r, childOpt) }() if !o { return false @@ -817,6 +839,7 @@ func arrayApproxEqualLargeList(left, right *LargeList, opt equalOption) bool { } func arrayApproxEqualListView(left, right *ListView, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -826,7 +849,7 @@ func arrayApproxEqualListView(left, right *ListView, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return arrayApproxEqual(l, r, opt) + return arrayApproxEqual(l, r, childOpt) }() if !o { return false @@ -836,6 +859,7 @@ func arrayApproxEqualListView(left, right *ListView, opt equalOption) bool { } func arrayApproxEqualLargeListView(left, right *LargeListView, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -845,7 +869,7 @@ func arrayApproxEqualLargeListView(left, right *LargeListView, opt equalOption) defer l.Release() r := right.newListValue(i) defer r.Release() - return arrayApproxEqual(l, r, opt) + return arrayApproxEqual(l, r, childOpt) }() if !o { return false @@ -855,6 +879,7 @@ func arrayApproxEqualLargeListView(left, right *LargeListView, opt equalOption) } func arrayApproxEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -864,7 +889,7 @@ func arrayApproxEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) defer l.Release() r := right.newListValue(i) defer r.Release() - return arrayApproxEqual(l, r, opt) + return arrayApproxEqual(l, r, childOpt) }() if !o { return false @@ -886,9 +911,11 @@ func arrayApproxEqualStruct(left, right *Struct, opt equalOption) bool { } func approxEqualStructRun(left, right *Struct, opt equalOption) bitutils.VisitFn { + st := left.DataType().(*arrow.StructType) return func(pos int64, length int64) error { for i := range left.fields { - if !sliceApproxEqual(left.fields[i], pos, pos+length, right.fields[i], pos, pos+length, opt) { + childOpt := withNullable(opt, st.Field(i).Nullable) + if !sliceApproxEqual(left.fields[i], pos, pos+length, right.fields[i], pos, pos+length, childOpt) { return arrow.ErrInvalid } } @@ -928,6 +955,10 @@ func arrayApproxEqualSingleMapEntry(left, right *Struct, opt equalOption) bool { return true } + st := left.DataType().(*arrow.StructType) + keyOpt := withNullable(opt, st.Field(0).Nullable) + valOpt := withNullable(opt, st.Field(1).Nullable) + used := make(map[int]bool, right.Len()) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { @@ -948,12 +979,12 @@ func arrayApproxEqualSingleMapEntry(left, right *Struct, opt equalOption) bool { rBeg, rEnd := int64(j), int64(j+1) // check keys (field 0) - if !sliceApproxEqual(left.Field(0), lBeg, lEnd, right.Field(0), rBeg, rEnd, opt) { + if !sliceApproxEqual(left.Field(0), lBeg, lEnd, right.Field(0), rBeg, rEnd, keyOpt) { continue } // only now check the values - if sliceApproxEqual(left.Field(1), lBeg, lEnd, right.Field(1), rBeg, rEnd, opt) { + if sliceApproxEqual(left.Field(1), lBeg, lEnd, right.Field(1), rBeg, rEnd, valOpt) { found = true used[j] = true break diff --git a/arrow/array/union.go b/arrow/array/union.go index 1d40ad66b..c10e9964d 100644 --- a/arrow/array/union.go +++ b/arrow/array/union.go @@ -320,17 +320,18 @@ func (a *SparseUnion) setData(data *Data) { debug.Assert(a.data.buffers[0] == nil, "arrow/array: validity bitmap for sparse unions should be nil") } -func (a *SparseUnion) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *SparseUnion) GetOneForMarshal(i int, nullable bool) any { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) data := a.Field(childID) - if nullable && data.IsNull(i) { - return nil + childNullable := a.unionType.Fields()[childID].Nullable + if childNullable && data.IsNull(i) { + return []any{typeID, nil} } - return []interface{}{typeID, data.GetOneForMarshal(i, nullable)} + return []any{typeID, data.GetOneForMarshal(i, childNullable)} } func (a *SparseUnion) MarshalJSON() ([]byte, error) { @@ -471,8 +472,9 @@ func arraySparseUnionApproxEqual(l, r *SparseUnion, opt equalOption) bool { } childNum := childIDs[typeID] + childOpt := withNullable(opt, l.unionType.Fields()[childNum].Nullable) eq := sliceApproxEqual(l.children[childNum], int64(i+l.data.offset), int64(i+l.data.offset+1), - r.children[childNum], int64(i+r.data.offset), int64(i+r.data.offset+1), opt) + r.children[childNum], int64(i+r.data.offset), int64(i+r.data.offset+1), childOpt) if !eq { return false } @@ -613,18 +615,19 @@ func (a *DenseUnion) setData(data *Data) { } } -func (a *DenseUnion) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *DenseUnion) GetOneForMarshal(i int, nullable bool) any { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) data := a.Field(childID) offset := int(a.RawValueOffsets()[i]) - if nullable && data.IsNull(offset) { - return nil + childNullable := a.unionType.Fields()[childID].Nullable + if childNullable && data.IsNull(offset) { + return []any{typeID, nil} } - return []interface{}{typeID, data.GetOneForMarshal(offset, nullable)} + return []any{typeID, data.GetOneForMarshal(offset, childNullable)} } func (a *DenseUnion) MarshalJSON() ([]byte, error) { @@ -715,8 +718,9 @@ func arrayDenseUnionApproxEqual(l, r *DenseUnion, opt equalOption) bool { } childNum := childIDs[typeID] + childOpt := withNullable(opt, l.unionType.Fields()[childNum].Nullable) eq := sliceApproxEqual(l.children[childNum], int64(leftOffsets[i]), int64(leftOffsets[i]+1), - r.children[childNum], int64(rightOffsets[i]), int64(rightOffsets[i]+1), opt) + r.children[childNum], int64(rightOffsets[i]), int64(rightOffsets[i]+1), childOpt) if !eq { return false } From 98c1bbe563caf387215e08867f6a4fc74935d172 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 15:53:33 -0400 Subject: [PATCH 11/18] Make exact Equal and chunked comparison field-local too Follow-up to the field-local nullability change (addresses review of 79cf1808): the exact Equal path and the chunked/table helpers still used parent/top-level nullability. arrayEqualList/LargeList/ListView/LargeListView and arrayEqualFixedSizeList now recurse into elements with the element field's nullability (arrayEqualStruct was already field-local; arrayEqualMap delegates to arrayEqualList). chunkedEqual and chunkedApproxEqual now skip the top-level NullN check for unions and run-end-encoded arrays via hasTopLevelValidityBitmap, so Table comparisons don't reject logically-equal union/REE columns. --- arrow/array/compare.go | 4 ++-- arrow/array/fixed_size_list.go | 3 ++- arrow/array/list.go | 12 ++++++++---- 3 files changed, 12 insertions(+), 7 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index 07f475bf7..3c9cdc8e2 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -138,7 +138,7 @@ func chunkedEqual(left, right *arrow.Chunked, opt equalOption) bool { return true case left.Len() != right.Len(): return false - case opt.nullable && left.NullN() != right.NullN(): + case opt.nullable && hasTopLevelValidityBitmap(left.DataType().ID()) && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): return false @@ -165,7 +165,7 @@ func chunkedApproxEqual(left, right *arrow.Chunked, opt equalOption) bool { return true case left.Len() != right.Len(): return false - case opt.nullable && left.NullN() != right.NullN(): + case opt.nullable && hasTopLevelValidityBitmap(left.DataType().ID()) && left.NullN() != right.NullN(): return false case !arrow.TypeEqual(left.DataType(), right.DataType()): return false diff --git a/arrow/array/fixed_size_list.go b/arrow/array/fixed_size_list.go index 10dce4c74..d4ddfe8f3 100644 --- a/arrow/array/fixed_size_list.go +++ b/arrow/array/fixed_size_list.go @@ -85,6 +85,7 @@ func (a *FixedSizeList) setData(data *Data) { } func arrayEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -94,7 +95,7 @@ func arrayEqualFixedSizeList(left, right *FixedSizeList, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return equal(l, r, opt) + return equal(l, r, childOpt) }() if !o { return false diff --git a/arrow/array/list.go b/arrow/array/list.go index 16eac56f9..a777eab6b 100644 --- a/arrow/array/list.go +++ b/arrow/array/list.go @@ -130,6 +130,7 @@ func (a *List) MarshalJSON() ([]byte, error) { } func arrayEqualList(left, right *List, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -139,7 +140,7 @@ func arrayEqualList(left, right *List, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return equal(l, r, opt) + return equal(l, r, childOpt) }() if !o { return false @@ -262,6 +263,7 @@ func (a *LargeList) MarshalJSON() ([]byte, error) { } func arrayEqualLargeList(left, right *LargeList, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -271,7 +273,7 @@ func arrayEqualLargeList(left, right *LargeList, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return equal(l, r, opt) + return equal(l, r, childOpt) }() if !o { return false @@ -737,6 +739,7 @@ func (a *ListView) MarshalJSON() ([]byte, error) { } func arrayEqualListView(left, right *ListView, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -746,7 +749,7 @@ func arrayEqualListView(left, right *ListView, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return equal(l, r, opt) + return equal(l, r, childOpt) }() if !o { return false @@ -884,6 +887,7 @@ func (a *LargeListView) MarshalJSON() ([]byte, error) { } func arrayEqualLargeListView(left, right *LargeListView, opt equalOption) bool { + childOpt := withNullable(opt, left.DataType().(arrow.ListLikeType).ElemField().Nullable) for i := 0; i < left.Len(); i++ { if opt.nullable && left.IsNull(i) { continue @@ -893,7 +897,7 @@ func arrayEqualLargeListView(left, right *LargeListView, opt equalOption) bool { defer l.Release() r := right.newListValue(i) defer r.Release() - return equal(l, r, opt) + return equal(l, r, childOpt) }() if !o { return false From c084d402fcfb2114cd559a9d01151da80fa0ef68 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 16:02:00 -0400 Subject: [PATCH 12/18] Preserve nullability option in exact chunked comparison Addresses review of 6f92ce2c: chunkedEqual called exported SliceEqual, which rebuilds default options and drops the caller's nullable setting, so TableEqual/ChunkedEqual with WithNullable(false) re-compared each chunk slice as nullable. Call the private sliceEqual(..., opt) instead, matching chunkedApproxEqual. --- arrow/array/compare.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index 3c9cdc8e2..f30f1a62e 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -146,7 +146,7 @@ func chunkedEqual(left, right *arrow.Chunked, opt equalOption) bool { var isequal = 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) + isequal = sliceEqual(left, lbeg, lend, right, rbeg, rend, opt) return isequal }) From 39d9f566ebac4c03790341f7f24917011d8c1d24 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 18:57:05 -0400 Subject: [PATCH 13/18] Address review feedback on nullability-aware comparison - array/compare.go: read the right-hand field from right.Schema() (not left.Schema()) in recordEqual, recordApproxEqual, tableEqual, and tableApproxEqual, so field/schema differences (name, nullability, metadata) between the two sides are actually detected instead of a field being compared against itself. - array/record_test.go: fix a malformed JSON map-entry literal (missing comma) in TestRecordBuilder. - ipc/metadata_test.go: TestUnrecognizedExtensionType now compares against the real preserved extension metadata (name arrow.uuid, empty serialized value). The previous placeholder only passed because record comparison ignored field metadata before the compare.go fix. --- arrow/array/compare.go | 8 ++++---- arrow/array/record_test.go | 2 +- arrow/ipc/metadata_test.go | 2 +- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index f30f1a62e..b34c025a2 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -40,7 +40,7 @@ func recordEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) - rf := left.Schema().Field(i) + rf := right.Schema().Field(i) if !lf.Equal(rf) { return false } @@ -72,7 +72,7 @@ func recordApproxEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) - rf := left.Schema().Field(i) + rf := right.Schema().Field(i) if !lf.Equal(rf) { return false } @@ -195,7 +195,7 @@ func tableEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) - rf := left.Schema().Field(i) + rf := right.Schema().Field(i) if !lf.Equal(rf) { return false } @@ -226,7 +226,7 @@ func tableApproxEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) - rf := left.Schema().Field(i) + rf := right.Schema().Field(i) if !lf.Equal(rf) { return false } diff --git a/arrow/array/record_test.go b/arrow/array/record_test.go index eabb84f3e..d07325c73 100644 --- a/arrow/array/record_test.go +++ b/arrow/array/record_test.go @@ -516,7 +516,7 @@ func TestRecordBuilder(t *testing.T) { } } - err := b.UnmarshalJSON([]byte(`{"f1-i32": 6, "f2-f64-notnull": 6.6, "map": [{"key": "4": "value": "d"}]}`)) + err := b.UnmarshalJSON([]byte(`{"f1-i32": 6, "f2-f64-notnull": 6.6, "map": [{"key": "4", "value": "d"}]}`)) assert.NoError(t, err) err = b.UnmarshalJSON([]byte(`{"f1-i32": null, "f2-f64-notnull": null, "map": null}`)) diff --git a/arrow/ipc/metadata_test.go b/arrow/ipc/metadata_test.go index ac8820dcc..eb72f972c 100644 --- a/arrow/ipc/metadata_test.go +++ b/arrow/ipc/metadata_test.go @@ -257,7 +257,7 @@ func TestUnrecognizedExtensionType(t *testing.T) { // create a record batch with the same data, but the field should contain the // extension metadata and be of the storage type instead of being the extension type. - extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"uuid", "uuid-serialized"}) + extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"arrow.uuid", ""}) batchNoExt := array.NewRecordBatch( arrow.NewSchema([]arrow.Field{ {Name: "f0", Type: storageArr.DataType(), Nullable: true, Metadata: extMetadata}, From 13e84dd8a4cc10e09c3742c6d42cf16b319376ff Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 21:51:44 -0400 Subject: [PATCH 14/18] Compare record/table fields ignoring metadata The right-schema comparison fix routed RecordEqual/RecordApproxEqual/ TableEqual/TableApproxEqual through Field.Equal, which also compares field and type metadata. That rejected logically-equal records differing only in metadata -- for example a PARQUET:field_id attached during a parquet round-trip -- breaking parquet/pqarrow tests such as TestForceLargeTypes. Introduce fieldEqualIgnoringMetadata (name + nullability + metadata- insensitive TypeEqual) and use it in the four comparison functions. It still detects the nullability differences the schema check was added for, while restoring the metadata-insensitive behavior these APIs had before. This makes the earlier metadata_test.go expected-metadata tweak unnecessary, so revert it. --- arrow/array/compare.go | 17 +++++++++++++---- arrow/ipc/metadata_test.go | 2 +- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index b34c025a2..77dd76551 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -25,6 +25,15 @@ import ( "github.com/apache/arrow-go/v18/internal/bitutils" ) +// fieldEqualIgnoringMetadata reports whether two fields describe the same column +// for record/table equality: same name, nullability, and type. Field and type +// metadata are intentionally ignored so records that differ only in metadata +// (for example a PARQUET:field_id attached during a round-trip) still compare as +// equal, preserving the pre-existing behavior of these comparison APIs. +func fieldEqualIgnoringMetadata(l, r arrow.Field) bool { + return l.Name == r.Name && l.Nullable == r.Nullable && arrow.TypeEqual(l.Type, r.Type) +} + // RecordEqual reports whether the two provided records are equal. func RecordEqual(left, right arrow.RecordBatch, opts ...EqualOption) bool { return recordEqual(left, right, newEqualOption(opts...)) @@ -41,7 +50,7 @@ func recordEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !lf.Equal(rf) { + if !fieldEqualIgnoringMetadata(lf, rf) { return false } @@ -73,7 +82,7 @@ func recordApproxEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !lf.Equal(rf) { + if !fieldEqualIgnoringMetadata(lf, rf) { return false } @@ -196,7 +205,7 @@ func tableEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !lf.Equal(rf) { + if !fieldEqualIgnoringMetadata(lf, rf) { return false } @@ -227,7 +236,7 @@ func tableApproxEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !lf.Equal(rf) { + if !fieldEqualIgnoringMetadata(lf, rf) { return false } diff --git a/arrow/ipc/metadata_test.go b/arrow/ipc/metadata_test.go index eb72f972c..ac8820dcc 100644 --- a/arrow/ipc/metadata_test.go +++ b/arrow/ipc/metadata_test.go @@ -257,7 +257,7 @@ func TestUnrecognizedExtensionType(t *testing.T) { // create a record batch with the same data, but the field should contain the // extension metadata and be of the storage type instead of being the extension type. - extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"arrow.uuid", ""}) + extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"uuid", "uuid-serialized"}) batchNoExt := array.NewRecordBatch( arrow.NewSchema([]arrow.Field{ {Name: "f0", Type: storageArr.DataType(), Nullable: true, Metadata: extMetadata}, From 45904705fe38a4bc91ab69ce1db8497b83ee178a Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 2 Jul 2026 22:01:10 -0400 Subject: [PATCH 15/18] Address roborev findings on record/table field comparison - compare.go: drop the arrow.TypeEqual call from the schema-field check. TypeEqual is not metadata-insensitive for large-list / list-view element types (they fall through to reflect.DeepEqual or Field.Equal), so relying on it could still reject records that differ only in child-field metadata. Field type and values are already validated per column by baseArrayEqual, so the schema-field check now compares only name and nullability (renamed schemaFieldEqual). - ipc/metadata_test.go: use the real preserved extension metadata (arrow.uuid, empty serialized value) and assert on the read-back field metadata explicitly, since record comparison no longer inspects field metadata. --- arrow/array/compare.go | 22 ++++++++++++---------- arrow/ipc/metadata_test.go | 5 ++++- 2 files changed, 16 insertions(+), 11 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index 77dd76551..348df5009 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -25,13 +25,15 @@ import ( "github.com/apache/arrow-go/v18/internal/bitutils" ) -// fieldEqualIgnoringMetadata reports whether two fields describe the same column -// for record/table equality: same name, nullability, and type. Field and type -// metadata are intentionally ignored so records that differ only in metadata +// schemaFieldEqual reports whether two fields agree on the schema-level +// attributes that record/table equality must match between the two sides: name +// and nullability. The field type and values are validated per column by the +// array comparison (baseArrayEqual), so they are not re-checked here, and field +// metadata is intentionally not compared so records that differ only in metadata // (for example a PARQUET:field_id attached during a round-trip) still compare as -// equal, preserving the pre-existing behavior of these comparison APIs. -func fieldEqualIgnoringMetadata(l, r arrow.Field) bool { - return l.Name == r.Name && l.Nullable == r.Nullable && arrow.TypeEqual(l.Type, r.Type) +// equal, matching the pre-existing behavior of these comparison APIs. +func schemaFieldEqual(l, r arrow.Field) bool { + return l.Name == r.Name && l.Nullable == r.Nullable } // RecordEqual reports whether the two provided records are equal. @@ -50,7 +52,7 @@ func recordEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !fieldEqualIgnoringMetadata(lf, rf) { + if !schemaFieldEqual(lf, rf) { return false } @@ -82,7 +84,7 @@ func recordApproxEqual(left, right arrow.RecordBatch, opt equalOption) bool { for i := range left.Columns() { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !fieldEqualIgnoringMetadata(lf, rf) { + if !schemaFieldEqual(lf, rf) { return false } @@ -205,7 +207,7 @@ func tableEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !fieldEqualIgnoringMetadata(lf, rf) { + if !schemaFieldEqual(lf, rf) { return false } @@ -236,7 +238,7 @@ func tableApproxEqual(left, right arrow.Table, opt equalOption) bool { for i := 0; int64(i) < left.NumCols(); i++ { lf := left.Schema().Field(i) rf := right.Schema().Field(i) - if !fieldEqualIgnoringMetadata(lf, rf) { + if !schemaFieldEqual(lf, rf) { return false } diff --git a/arrow/ipc/metadata_test.go b/arrow/ipc/metadata_test.go index ac8820dcc..63e700b8b 100644 --- a/arrow/ipc/metadata_test.go +++ b/arrow/ipc/metadata_test.go @@ -257,12 +257,15 @@ func TestUnrecognizedExtensionType(t *testing.T) { // create a record batch with the same data, but the field should contain the // extension metadata and be of the storage type instead of being the extension type. - extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"uuid", "uuid-serialized"}) + extMetadata := arrow.NewMetadata([]string{ExtensionTypeKeyName, ExtensionMetadataKeyName}, []string{"arrow.uuid", ""}) batchNoExt := array.NewRecordBatch( arrow.NewSchema([]arrow.Field{ {Name: "f0", Type: storageArr.DataType(), Nullable: true, Metadata: extMetadata}, }, nil), []arrow.Array{storageArr}, 4) defer batchNoExt.Release() + // RecordEqual ignores field metadata, so explicitly verify the unrecognized + // extension metadata is preserved on the read-back field. + assert.Truef(t, rec.Schema().Field(0).Metadata.Equal(extMetadata), "expected metadata %v, got %v", extMetadata, rec.Schema().Field(0).Metadata) assert.Truef(t, array.RecordEqual(rec, batchNoExt), "expected: %s\ngot: %s\n", batchNoExt, rec) } From 94644b37904c8f3867b1a8ccee9c65099436d754 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Sat, 4 Jul 2026 13:44:04 -0400 Subject: [PATCH 16/18] Avoid breaking arrow.Array with optional NullableMarshaler GetOneForMarshal previously grew a nullable bool parameter on the public arrow.Array interface, breaking external callers and custom Array implementations. Restore the single-argument GetOneForMarshal(i int) and move the field-local nullability behavior to a new optional exported interface, arrow.NullableMarshaler, with GetOneForMarshalNullable(i int, nullable bool). Containers (struct, sparse/dense union, dictionary, run-end-encoded, extension) and RecordToJSON propagate each field's Nullable flag through an unexported getOneForMarshalNullable helper that type-asserts to NullableMarshaler and falls back to GetOneForMarshal for arrays that do not implement it, preserving the previous validity-bitmap behavior. ExtensionArrayBase deliberately does not implement NullableMarshaler so that the optional interface is never promoted onto embedding extension arrays and cannot bypass a concrete type's GetOneForMarshal override. The helper instead handles the one field-local case that matters for plain extension arrays: a null slot in a non-nullable field serializes the storage value rather than JSON null, while still honoring any custom GetOneForMarshal that returns a non-null logical value. --- arrow/array.go | 17 ++++++++++++- arrow/array/binary.go | 24 +++++++++++++----- arrow/array/boolean.go | 6 ++++- arrow/array/decimal.go | 10 +++++--- arrow/array/decimal128_test.go | 2 +- arrow/array/decimal256_test.go | 2 +- arrow/array/dictionary.go | 11 ++++++--- arrow/array/encoded.go | 17 +++++++++---- arrow/array/extension.go | 9 +++++-- arrow/array/fixed_size_list.go | 8 ++++-- arrow/array/fixedsize_binary.go | 6 ++++- arrow/array/float16.go | 6 ++++- arrow/array/interval.go | 22 +++++++++++++---- arrow/array/json_reader_test.go | 2 +- arrow/array/list.go | 40 +++++++++++++++++++++--------- arrow/array/null.go | 6 ++++- arrow/array/numeric_generic.go | 44 +++++++++++++++++++++++++-------- arrow/array/string.go | 22 +++++++++++++---- arrow/array/struct.go | 12 ++++++--- arrow/array/timestamp.go | 8 ++++-- arrow/array/union.go | 28 +++++++++++++-------- arrow/array/util.go | 30 +++++++++++++++++++++- arrow/extensions/bool8.go | 6 ++++- arrow/extensions/json.go | 8 ++++-- arrow/extensions/uuid.go | 8 ++++-- arrow/extensions/variant.go | 6 ++++- 26 files changed, 275 insertions(+), 85 deletions(-) diff --git a/arrow/array.go b/arrow/array.go index 891697e9b..09fafe42f 100644 --- a/arrow/array.go +++ b/arrow/array.go @@ -111,7 +111,7 @@ type Array interface { ValueStr(i int) string // Get single value to be marshalled with `json.Marshal` - GetOneForMarshal(i int, nullable bool) interface{} + GetOneForMarshal(i int) interface{} Data() ArrayData @@ -128,6 +128,21 @@ type Array interface { Release() } +// NullableMarshaler is an optional interface an Array may implement to control +// whether a value is rendered as JSON null based on the field-local nullability +// from the schema, rather than solely on the array's own validity bitmap. +// +// It exists so containers (struct/union/dictionary/run-end-encoded/extension) +// and RecordToJSON can propagate each field's Nullable flag down to its values +// without a breaking change to the Array interface. Arrays that do not implement +// it fall back to GetOneForMarshal, preserving the previous validity-based behavior. +type NullableMarshaler interface { + // GetOneForMarshalNullable returns the value at i for json.Marshal. When + // nullable is false, a value is returned even if the array's validity bitmap + // marks it null (a non-nullable field must not serialize as null). + GetOneForMarshalNullable(i int, nullable bool) interface{} +} + // ValueType is a generic constraint for valid Arrow primitive types type ValueType interface { bool | FixedWidthType | string | []byte diff --git a/arrow/array/binary.go b/arrow/array/binary.go index ba9ce287b..1801d9d64 100644 --- a/arrow/array/binary.go +++ b/arrow/array/binary.go @@ -152,17 +152,21 @@ func (a *Binary) setData(data *Data) { } } -func (a *Binary) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Binary) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } return a.Value(i) } +func (a *Binary) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Binary) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string @@ -346,17 +350,21 @@ func (a *LargeBinary) setData(data *Data) { } } -func (a *LargeBinary) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *LargeBinary) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } return a.Value(i) } +func (a *LargeBinary) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *LargeBinary) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string @@ -522,17 +530,21 @@ func (a *BinaryView) ValueStr(i int) string { return base64.StdEncoding.EncodeToString(a.Value(i)) } -func (a *BinaryView) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *BinaryView) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } return a.Value(i) } +func (a *BinaryView) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *BinaryView) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } // golang marshal standard says that []byte will be marshalled // as a base64-encoded string diff --git a/arrow/array/boolean.go b/arrow/array/boolean.go index c555536f5..183d4955a 100644 --- a/arrow/array/boolean.go +++ b/arrow/array/boolean.go @@ -90,13 +90,17 @@ func (a *Boolean) setData(data *Data) { } } -func (a *Boolean) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Boolean) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.Value(i) } return nil } +func (a *Boolean) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Boolean) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { diff --git a/arrow/array/decimal.go b/arrow/array/decimal.go index 3227f66e3..269ca2904 100644 --- a/arrow/array/decimal.go +++ b/arrow/array/decimal.go @@ -55,7 +55,7 @@ func (a *baseDecimal[T]) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - return a.GetOneForMarshal(i, true).(string) + return a.GetOneForMarshal(i).(string) } func (a *baseDecimal[T]) Values() []T { return a.values } @@ -89,7 +89,7 @@ func (a *baseDecimal[T]) setData(data *Data) { } } -func (a *baseDecimal[T]) GetOneForMarshal(i int, nullable bool) any { +func (a *baseDecimal[T]) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && a.IsNull(i) { return nil } @@ -99,10 +99,14 @@ func (a *baseDecimal[T]) GetOneForMarshal(i int, nullable bool) any { return n.ToBigFloat(scale).Text('g', int(typ.GetPrecision())) } +func (a *baseDecimal[T]) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) +} + func (a *baseDecimal[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } diff --git a/arrow/array/decimal128_test.go b/arrow/array/decimal128_test.go index dcb234a63..e642d0374 100644 --- a/arrow/array/decimal128_test.go +++ b/arrow/array/decimal128_test.go @@ -279,7 +279,7 @@ func TestDecimal128GetOneForMarshal(t *testing.T) { } for i := range cases { - assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i, true), "unexpected value at index %d", i) + assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i), "unexpected value at index %d", i) } } diff --git a/arrow/array/decimal256_test.go b/arrow/array/decimal256_test.go index 0d6e9f140..b5674253e 100644 --- a/arrow/array/decimal256_test.go +++ b/arrow/array/decimal256_test.go @@ -288,6 +288,6 @@ func TestDecimal256GetOneForMarshal(t *testing.T) { } for i := range cases { - assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i, true), "unexpected value at index %d", i) + assert.Equalf(t, cases[i].want, arr.GetOneForMarshal(i), "unexpected value at index %d", i) } } diff --git a/arrow/array/dictionary.go b/arrow/array/dictionary.go index 966131047..982aa2e3e 100644 --- a/arrow/array/dictionary.go +++ b/arrow/array/dictionary.go @@ -286,19 +286,22 @@ func (d *Dictionary) GetValueIndex(i int) int { debug.Assert(false, "unreachable dictionary index") return -1 } - -func (d *Dictionary) GetOneForMarshal(i int, nullable bool) interface{} { +func (d *Dictionary) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && d.IsNull(i) { return nil } vidx := d.GetValueIndex(i) - return d.Dictionary().GetOneForMarshal(vidx, nullable) + return getOneForMarshalNullable(d.Dictionary(), vidx, nullable) +} + +func (d *Dictionary) GetOneForMarshal(i int) interface{} { + return d.GetOneForMarshalNullable(i, true) } func (d *Dictionary) MarshalJSON() ([]byte, error) { vals := make([]any, d.Len()) for i := range d.Len() { - vals[i] = d.GetOneForMarshal(i, true) + vals[i] = d.GetOneForMarshal(i) } return json.Marshal(vals) } diff --git a/arrow/array/encoded.go b/arrow/array/encoded.go index bcfeded76..441bf9ded 100644 --- a/arrow/array/encoded.go +++ b/arrow/array/encoded.go @@ -219,13 +219,13 @@ func (r *RunEndEncoded) String() string { buf.WriteByte(',') } - value := r.values.GetOneForMarshal(i, true) + value := r.values.GetOneForMarshal(i) if byts, ok := value.(json.RawMessage); ok { value = string(byts) } var runEnd int - switch e := r.ends.GetOneForMarshal(i, true).(type) { + switch e := r.ends.GetOneForMarshal(i).(type) { case int16: runEnd = int(e) - r.data.offset case int32: @@ -239,9 +239,16 @@ func (r *RunEndEncoded) String() string { buf.WriteByte(']') return buf.String() } +func (r *RunEndEncoded) GetOneForMarshalNullable(i int, nullable bool) interface{} { + // The values child may serialize as JSON null only when both the outer REE + // field and the REE values child (ValueNullable) are nullable; a non-nullable + // outer field must never emit null even if the values child is nullable. + nullable = nullable && r.data.dtype.(*arrow.RunEndEncodedType).ValueNullable + return getOneForMarshalNullable(r.values, r.GetPhysicalIndex(i), nullable) +} -func (r *RunEndEncoded) GetOneForMarshal(i int, nullable bool) interface{} { - return r.values.GetOneForMarshal(r.GetPhysicalIndex(i), nullable) +func (r *RunEndEncoded) GetOneForMarshal(i int) interface{} { + return r.GetOneForMarshalNullable(i, true) } func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { @@ -252,7 +259,7 @@ func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(r.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(r.GetOneForMarshal(i)); err != nil { return nil, err } } diff --git a/arrow/array/extension.go b/arrow/array/extension.go index 48c4f03ad..a74f6dee2 100644 --- a/arrow/array/extension.go +++ b/arrow/array/extension.go @@ -116,8 +116,13 @@ func (e *ExtensionArrayBase) String() string { return fmt.Sprintf("(%s)%s", e.data.dtype, e.storage) } -func (e *ExtensionArrayBase) GetOneForMarshal(i int, nullable bool) interface{} { - return e.storage.GetOneForMarshal(i, nullable) +// GetOneForMarshal returns the value at i from the underlying storage array. +// ExtensionArrayBase deliberately does not implement arrow.NullableMarshaler: +// doing so would promote a nullable-aware method onto every embedding extension +// array and bypass a concrete type's own GetOneForMarshal override. Field-local +// nullability for plain extension arrays is handled by the marshaling helper. +func (e *ExtensionArrayBase) GetOneForMarshal(i int) interface{} { + return e.storage.GetOneForMarshal(i) } func (e *ExtensionArrayBase) MarshalJSON() ([]byte, error) { diff --git a/arrow/array/fixed_size_list.go b/arrow/array/fixed_size_list.go index d4ddfe8f3..07fbe4411 100644 --- a/arrow/array/fixed_size_list.go +++ b/arrow/array/fixed_size_list.go @@ -51,7 +51,7 @@ func (a *FixedSizeList) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - return string(a.GetOneForMarshal(i, true).(json.RawMessage)) + return string(a.GetOneForMarshal(i).(json.RawMessage)) } func (a *FixedSizeList) String() string { @@ -124,7 +124,7 @@ func (a *FixedSizeList) Release() { a.values.Release() } -func (a *FixedSizeList) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *FixedSizeList) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -138,6 +138,10 @@ func (a *FixedSizeList) GetOneForMarshal(i int, nullable bool) interface{} { return json.RawMessage(v) } +func (a *FixedSizeList) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *FixedSizeList) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) diff --git a/arrow/array/fixedsize_binary.go b/arrow/array/fixedsize_binary.go index b9cf2bfdc..11dc76154 100644 --- a/arrow/array/fixedsize_binary.go +++ b/arrow/array/fixedsize_binary.go @@ -86,7 +86,7 @@ func (a *FixedSizeBinary) setData(data *Data) { } } -func (a *FixedSizeBinary) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *FixedSizeBinary) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -94,6 +94,10 @@ func (a *FixedSizeBinary) GetOneForMarshal(i int, nullable bool) interface{} { return a.Value(i) } +func (a *FixedSizeBinary) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *FixedSizeBinary) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { diff --git a/arrow/array/float16.go b/arrow/array/float16.go index 8536df402..8923eb438 100644 --- a/arrow/array/float16.go +++ b/arrow/array/float16.go @@ -77,13 +77,17 @@ func (a *Float16) setData(data *Data) { } } -func (a *Float16) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Float16) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.values[i].Float32() } return nil } +func (a *Float16) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Float16) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i, v := range a.values { diff --git a/arrow/array/interval.go b/arrow/array/interval.go index b2aad56f1..076a7a9d4 100644 --- a/arrow/array/interval.go +++ b/arrow/array/interval.go @@ -94,13 +94,17 @@ func (a *MonthInterval) setData(data *Data) { } } -func (a *MonthInterval) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *MonthInterval) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.values[i] } return nil } +func (a *MonthInterval) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + // MarshalJSON will create a json array out of a MonthInterval array, // each value will be an object of the form {"months": #} where // # is the numeric value of that index @@ -361,7 +365,7 @@ func (a *DayTimeInterval) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i, true)) + data, err := json.Marshal(a.GetOneForMarshal(i)) if err != nil { panic(err) } @@ -399,13 +403,17 @@ func (a *DayTimeInterval) setData(data *Data) { } } -func (a *DayTimeInterval) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *DayTimeInterval) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.values[i] } return nil } +func (a *DayTimeInterval) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + // MarshalJSON will marshal this array to JSON as an array of objects, // consisting of the form {"days": #, "milliseconds": #} for each element. func (a *DayTimeInterval) MarshalJSON() ([]byte, error) { @@ -663,7 +671,7 @@ func (a *MonthDayNanoInterval) ValueStr(i int) string { if a.IsNull(i) { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i, true)) + data, err := json.Marshal(a.GetOneForMarshal(i)) if err != nil { panic(err) } @@ -703,13 +711,17 @@ func (a *MonthDayNanoInterval) setData(data *Data) { } } -func (a *MonthDayNanoInterval) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *MonthDayNanoInterval) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.values[i] } return nil } +func (a *MonthDayNanoInterval) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + // MarshalJSON will marshal this array to a JSON array with elements // marshalled to the form {"months": #, "days": #, "nanoseconds": #} func (a *MonthDayNanoInterval) MarshalJSON() ([]byte, error) { diff --git a/arrow/array/json_reader_test.go b/arrow/array/json_reader_test.go index fc7cedea9..3d0def659 100644 --- a/arrow/array/json_reader_test.go +++ b/arrow/array/json_reader_test.go @@ -253,7 +253,7 @@ func recordBatchToNDJSON(t *testing.T, rec arrow.RecordBatch) string { defer arr.Release() for pos := range arr.Len() { - s, err := json.Marshal(arr.GetOneForMarshal(pos, true)) + s, err := json.Marshal(arr.GetOneForMarshal(pos)) assert.NoError(t, err) sb.Write(s) sb.WriteByte('\n') diff --git a/arrow/array/list.go b/arrow/array/list.go index a777eab6b..7e6348247 100644 --- a/arrow/array/list.go +++ b/arrow/array/list.go @@ -61,7 +61,7 @@ func (a *List) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i, true).(json.RawMessage)) + return string(a.GetOneForMarshal(i).(json.RawMessage)) } func (a *List) String() string { @@ -98,7 +98,7 @@ func (a *List) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *List) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *List) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -112,6 +112,10 @@ func (a *List) GetOneForMarshal(i int, nullable bool) interface{} { return json.RawMessage(v) } +func (a *List) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *List) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) @@ -121,7 +125,7 @@ func (a *List) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } @@ -194,7 +198,7 @@ func (a *LargeList) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i, true).(json.RawMessage)) + return string(a.GetOneForMarshal(i).(json.RawMessage)) } func (a *LargeList) String() string { @@ -231,7 +235,7 @@ func (a *LargeList) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *LargeList) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *LargeList) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -245,6 +249,10 @@ func (a *LargeList) GetOneForMarshal(i int, nullable bool) interface{} { return json.RawMessage(v) } +func (a *LargeList) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *LargeList) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) @@ -254,7 +262,7 @@ func (a *LargeList) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } @@ -666,7 +674,7 @@ func (a *ListView) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i, true).(json.RawMessage)) + return string(a.GetOneForMarshal(i).(json.RawMessage)) } func (a *ListView) String() string { @@ -707,7 +715,7 @@ func (a *ListView) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *ListView) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *ListView) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -721,6 +729,10 @@ func (a *ListView) GetOneForMarshal(i int, nullable bool) interface{} { return json.RawMessage(v) } +func (a *ListView) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *ListView) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) @@ -730,7 +742,7 @@ func (a *ListView) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } @@ -814,7 +826,7 @@ func (a *LargeListView) ValueStr(i int) string { if !a.IsValid(i) { return NullValueStr } - return string(a.GetOneForMarshal(i, true).(json.RawMessage)) + return string(a.GetOneForMarshal(i).(json.RawMessage)) } func (a *LargeListView) String() string { @@ -855,7 +867,7 @@ func (a *LargeListView) setData(data *Data) { a.values = MakeFromData(data.childData[0]) } -func (a *LargeListView) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *LargeListView) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -869,6 +881,10 @@ func (a *LargeListView) GetOneForMarshal(i int, nullable bool) interface{} { return json.RawMessage(v) } +func (a *LargeListView) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *LargeListView) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) @@ -878,7 +894,7 @@ func (a *LargeListView) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } diff --git a/arrow/array/null.go b/arrow/array/null.go index 01001776c..bdb4ee22b 100644 --- a/arrow/array/null.go +++ b/arrow/array/null.go @@ -80,10 +80,14 @@ func (a *Null) setData(data *Data) { a.data.nulls = a.data.length } -func (a *Null) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Null) GetOneForMarshalNullable(i int, nullable bool) interface{} { return nil } +func (a *Null) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Null) MarshalJSON() ([]byte, error) { return json.Marshal(make([]interface{}, a.Len())) } diff --git a/arrow/array/numeric_generic.go b/arrow/array/numeric_generic.go index c7616cf20..5abade6d2 100644 --- a/arrow/array/numeric_generic.go +++ b/arrow/array/numeric_generic.go @@ -83,7 +83,7 @@ func (a *numericArray[T]) ValueStr(i int) string { return fmt.Sprintf("%v", a.values[i]) } -func (a *numericArray[T]) GetOneForMarshal(i int, nullable bool) any { +func (a *numericArray[T]) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && a.IsNull(i) { return nil } @@ -91,6 +91,10 @@ func (a *numericArray[T]) GetOneForMarshal(i int, nullable bool) any { return a.values[i] } +func (a *numericArray[T]) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) +} + func (a *numericArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.Len() { @@ -107,7 +111,7 @@ type oneByteArrs[T int8 | uint8] struct { numericArray[T] } -func (a *oneByteArrs[T]) GetOneForMarshal(i int, nullable bool) any { +func (a *oneByteArrs[T]) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && a.IsNull(i) { return nil } @@ -115,6 +119,10 @@ func (a *oneByteArrs[T]) GetOneForMarshal(i int, nullable bool) any { return float64(a.values[i]) // prevent uint8/int8 from being seen as binary data } +func (a *oneByteArrs[T]) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) +} + func (a *oneByteArrs[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.Len() { @@ -141,7 +149,7 @@ func (a *floatArray[T]) ValueStr(i int) string { return strconv.FormatFloat(float64(a.Value(i)), 'g', -1, bitWidth) } -func (a *floatArray[T]) GetOneForMarshal(i int, nullable bool) any { +func (a *floatArray[T]) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && a.IsNull(i) { return nil } @@ -157,10 +165,14 @@ func (a *floatArray[T]) GetOneForMarshal(i int, nullable bool) any { } } +func (a *floatArray[T]) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) +} + func (a *floatArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } @@ -176,7 +188,7 @@ type dateArray[T interface { func (d *dateArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, d.Len()) for i := range d.values { - vals[i] = d.GetOneForMarshal(i, true) + vals[i] = d.GetOneForMarshal(i) } return json.Marshal(vals) } @@ -189,7 +201,7 @@ func (d *dateArray[T]) ValueStr(i int) string { return d.values[i].FormattedString() } -func (d *dateArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { +func (d *dateArray[T]) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && d.IsNull(i) { return nil } @@ -197,6 +209,10 @@ func (d *dateArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { return d.values[i].FormattedString() } +func (d *dateArray[T]) GetOneForMarshal(i int) interface{} { + return d.GetOneForMarshalNullable(i, true) +} + type timeType interface { TimeUnit() arrow.TimeUnit } @@ -212,7 +228,7 @@ type timeArray[T interface { func (a *timeArray[T]) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } @@ -225,7 +241,7 @@ func (a *timeArray[T]) ValueStr(i int) string { return a.values[i].FormattedString(a.DataType().(timeType).TimeUnit()) } -func (a *timeArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *timeArray[T]) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -233,6 +249,10 @@ func (a *timeArray[T]) GetOneForMarshal(i int, nullable bool) interface{} { return a.values[i].ToTime(a.DataType().(timeType).TimeUnit()).Format("15:04:05.999999999") } +func (a *timeArray[T]) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + type Duration struct { numericArray[arrow.Duration] } @@ -248,7 +268,7 @@ func (a *Duration) DurationValues() []arrow.Duration { return a.Values() } func (a *Duration) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } @@ -261,13 +281,17 @@ func (a *Duration) ValueStr(i int) string { return fmt.Sprintf("%d%s", a.values[i], a.DataType().(timeType).TimeUnit()) } -func (a *Duration) GetOneForMarshal(i int, nullable bool) any { +func (a *Duration) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && a.IsNull(i) { return nil } return fmt.Sprintf("%d%s", a.values[i], a.DataType().(timeType).TimeUnit()) } +func (a *Duration) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) +} + type Int64 struct { numericArray[int64] } diff --git a/arrow/array/string.go b/arrow/array/string.go index e3ed07236..bbeba8879 100644 --- a/arrow/array/string.go +++ b/arrow/array/string.go @@ -154,13 +154,17 @@ func (a *String) setData(data *Data) { } } -func (a *String) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *String) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.Value(i) } return nil } +func (a *String) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *String) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { @@ -362,17 +366,21 @@ func (a *LargeString) setData(data *Data) { } } -func (a *LargeString) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *LargeString) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.Value(i) } return nil } +func (a *LargeString) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *LargeString) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } @@ -526,17 +534,21 @@ func (a *StringView) ValueStr(i int) string { return a.Value(i) } -func (a *StringView) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *StringView) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } return a.Value(i) } +func (a *StringView) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *StringView) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := 0; i < a.Len(); i++ { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } diff --git a/arrow/array/struct.go b/arrow/array/struct.go index 17f6152a3..5579158d0 100644 --- a/arrow/array/struct.go +++ b/arrow/array/struct.go @@ -130,7 +130,7 @@ func (a *Struct) ValueStr(i int) string { return NullValueStr } - data, err := json.Marshal(a.GetOneForMarshal(i, true)) + data, err := json.Marshal(a.GetOneForMarshal(i)) if err != nil { panic(err) } @@ -209,7 +209,7 @@ func (a *Struct) setData(data *Data) { } } -func (a *Struct) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Struct) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } @@ -218,11 +218,15 @@ func (a *Struct) GetOneForMarshal(i int, nullable bool) interface{} { dtype := a.data.dtype.(*arrow.StructType) fieldList := dtype.Fields() for j, d := range a.fields { - tmp[fieldList[j].Name] = d.GetOneForMarshal(i, dtype.Field(j).Nullable) + tmp[fieldList[j].Name] = getOneForMarshalNullable(d, i, dtype.Field(j).Nullable) } return tmp } +func (a *Struct) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Struct) MarshalJSON() ([]byte, error) { var buf bytes.Buffer enc := json.NewEncoder(&buf) @@ -232,7 +236,7 @@ func (a *Struct) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } diff --git a/arrow/array/timestamp.go b/arrow/array/timestamp.go index 613941f7f..e2897fac5 100644 --- a/arrow/array/timestamp.go +++ b/arrow/array/timestamp.go @@ -109,17 +109,21 @@ func (a *Timestamp) ValueStr(i int) string { return toTime(a.values[i]).Format(layout) } -func (a *Timestamp) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Timestamp) GetOneForMarshalNullable(i int, nullable bool) interface{} { if val := a.ValueStr(i); !nullable || val != NullValueStr { return val } return nil } +func (a *Timestamp) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + func (a *Timestamp) MarshalJSON() ([]byte, error) { vals := make([]interface{}, a.Len()) for i := range a.values { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) diff --git a/arrow/array/union.go b/arrow/array/union.go index c10e9964d..8e94900c9 100644 --- a/arrow/array/union.go +++ b/arrow/array/union.go @@ -320,7 +320,7 @@ func (a *SparseUnion) setData(data *Data) { debug.Assert(a.data.buffers[0] == nil, "arrow/array: validity bitmap for sparse unions should be nil") } -func (a *SparseUnion) GetOneForMarshal(i int, nullable bool) any { +func (a *SparseUnion) GetOneForMarshalNullable(i int, nullable bool) any { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) @@ -331,7 +331,11 @@ func (a *SparseUnion) GetOneForMarshal(i int, nullable bool) any { return []any{typeID, nil} } - return []any{typeID, data.GetOneForMarshal(i, childNullable)} + return []any{typeID, getOneForMarshalNullable(data, i, childNullable)} +} + +func (a *SparseUnion) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) } func (a *SparseUnion) MarshalJSON() ([]byte, error) { @@ -343,7 +347,7 @@ func (a *SparseUnion) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } @@ -356,7 +360,7 @@ func (a *SparseUnion) ValueStr(i int) string { return NullValueStr } - val := a.GetOneForMarshal(i, true) + val := a.GetOneForMarshal(i) if val == nil { // child is nil return NullValueStr @@ -381,7 +385,7 @@ func (a *SparseUnion) String() string { field := fieldList[a.ChildID(i)] f := a.Field(a.ChildID(i)) - fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(i, true)) + fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(i)) } b.WriteByte(']') return b.String() @@ -615,7 +619,7 @@ func (a *DenseUnion) setData(data *Data) { } } -func (a *DenseUnion) GetOneForMarshal(i int, nullable bool) any { +func (a *DenseUnion) GetOneForMarshalNullable(i int, nullable bool) any { typeID := a.RawTypeCodes()[i] childID := a.ChildID(i) @@ -627,7 +631,11 @@ func (a *DenseUnion) GetOneForMarshal(i int, nullable bool) any { return []any{typeID, nil} } - return []any{typeID, data.GetOneForMarshal(offset, childNullable)} + return []any{typeID, getOneForMarshalNullable(data, offset, childNullable)} +} + +func (a *DenseUnion) GetOneForMarshal(i int) any { + return a.GetOneForMarshalNullable(i, true) } func (a *DenseUnion) MarshalJSON() ([]byte, error) { @@ -639,7 +647,7 @@ func (a *DenseUnion) MarshalJSON() ([]byte, error) { if i != 0 { buf.WriteByte(',') } - if err := enc.Encode(a.GetOneForMarshal(i, true)); err != nil { + if err := enc.Encode(a.GetOneForMarshal(i)); err != nil { return nil, err } } @@ -652,7 +660,7 @@ func (a *DenseUnion) ValueStr(i int) string { return NullValueStr } - val := a.GetOneForMarshal(i, true) + val := a.GetOneForMarshal(i) if val == nil { // child in nil return NullValueStr @@ -679,7 +687,7 @@ func (a *DenseUnion) String() string { field := fieldList[a.ChildID(i)] f := a.Field(a.ChildID(i)) - fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(int(offsets[i]), true)) + fmt.Fprintf(&b, "{%s=%v}", field.Name, f.GetOneForMarshal(int(offsets[i]))) } b.WriteByte(']') return b.String() diff --git a/arrow/array/util.go b/arrow/array/util.go index b3375152a..0461e1717 100644 --- a/arrow/array/util.go +++ b/arrow/array/util.go @@ -284,7 +284,7 @@ func RecordToJSON(rec arrow.RecordBatch, w io.Writer) error { cols := make(map[string]interface{}) for i := 0; int64(i) < rec.NumRows(); i++ { for j, c := range rec.Columns() { - cols[fields[j].Name] = c.GetOneForMarshal(i, rec.Schema().Field(j).Nullable) + cols[fields[j].Name] = getOneForMarshalNullable(c, i, rec.Schema().Field(j).Nullable) } if err := enc.Encode(cols); err != nil { return err @@ -293,6 +293,34 @@ func RecordToJSON(rec arrow.RecordBatch, w io.Writer) error { return nil } +// getOneForMarshalNullable dispatches to an Array's field-local-nullability-aware +// marshaler when it implements arrow.NullableMarshaler, otherwise it falls back to +// the plain GetOneForMarshal (which decides null purely from the validity bitmap). +// This lets containers propagate a field's Nullable flag without every Array having +// to implement the optional interface. +func getOneForMarshalNullable(a arrow.Array, i int, nullable bool) interface{} { + if nm, ok := a.(arrow.NullableMarshaler); ok { + return nm.GetOneForMarshalNullable(i, nullable) + } + + // ExtensionArrayBase intentionally does not implement arrow.NullableMarshaler + // (see extension.go), so extension arrays reach here. Preserve the extension's + // own logical JSON by default; the only field-local nullability we can safely + // add is that a null slot in a non-nullable field must not serialize as null. + // Ask the extension first (honoring any custom GetOneForMarshal override); only + // if it would emit null do we fall back to the storage value. Plain wrapper + // extensions (e.g. Parametric*Array) return nil at a null slot and thus round + // trip through storage, without bypassing a custom representation. + if ext, ok := a.(ExtensionArray); ok && !nullable && a.IsNull(i) { + if v := a.GetOneForMarshal(i); v != nil { + return v + } + return getOneForMarshalNullable(ext.Storage(), i, false) + } + + return a.GetOneForMarshal(i) +} + func TableFromJSON(mem memory.Allocator, sc *arrow.Schema, recJSON []string, opt ...FromJSONOption) (arrow.Table, error) { batches := make([]arrow.RecordBatch, len(recJSON)) for i, batchJSON := range recJSON { diff --git a/arrow/extensions/bool8.go b/arrow/extensions/bool8.go index aaf51f0c5..8f784a043 100644 --- a/arrow/extensions/bool8.go +++ b/arrow/extensions/bool8.go @@ -114,13 +114,17 @@ func (a *Bool8Array) MarshalJSON() ([]byte, error) { return json.Marshal(values) } -func (a *Bool8Array) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *Bool8Array) GetOneForMarshalNullable(i int, nullable bool) interface{} { if nullable && a.IsNull(i) { return nil } return a.Value(i) } +func (a *Bool8Array) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + // boolToInt8 performs the simple scalar conversion of bool to the canonical int8 // value for the Bool8Type. func boolToInt8(v bool) int8 { diff --git a/arrow/extensions/json.go b/arrow/extensions/json.go index b9cc50a73..4b7194f94 100644 --- a/arrow/extensions/json.go +++ b/arrow/extensions/json.go @@ -118,7 +118,7 @@ func (a *JSONArray) ValueBytes(i int) []byte { func (a *JSONArray) valueJSON(i int, nullable bool) json.RawMessage { var val json.RawMessage - if a.IsValid(i) { + if !nullable || a.IsValid(i) { val = json.RawMessage(a.Storage().(array.StringLike).Value(i)) } return val @@ -142,10 +142,14 @@ func (a *JSONArray) MarshalJSON() ([]byte, error) { } // GetOneForMarshal implements arrow.Array. -func (a *JSONArray) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *JSONArray) GetOneForMarshalNullable(i int, nullable bool) interface{} { return a.valueJSON(i, nullable) } +func (a *JSONArray) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + var ( _ arrow.ExtensionType = (*JSONType)(nil) _ array.ExtensionArray = (*JSONArray)(nil) diff --git a/arrow/extensions/uuid.go b/arrow/extensions/uuid.go index cbe56ecd6..cbe4d7c22 100644 --- a/arrow/extensions/uuid.go +++ b/arrow/extensions/uuid.go @@ -190,18 +190,22 @@ func (a *UUIDArray) ValueStr(i int) string { func (a *UUIDArray) MarshalJSON() ([]byte, error) { vals := make([]any, a.Len()) for i := range vals { - vals[i] = a.GetOneForMarshal(i, true) + vals[i] = a.GetOneForMarshal(i) } return json.Marshal(vals) } -func (a *UUIDArray) GetOneForMarshal(i int, nullable bool) interface{} { +func (a *UUIDArray) GetOneForMarshalNullable(i int, nullable bool) interface{} { if !nullable || a.IsValid(i) { return a.Value(i) } return nil } +func (a *UUIDArray) GetOneForMarshal(i int) interface{} { + return a.GetOneForMarshalNullable(i, true) +} + // UUIDType is a simple extension type that represents a FixedSizeBinary(16) // to be used for representing UUIDs type UUIDType struct { diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index 1f41ef076..e2123485b 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -599,7 +599,7 @@ func (v *VariantArray) MarshalJSON() ([]byte, error) { return json.Marshal(values) } -func (v *VariantArray) GetOneForMarshal(i int, nullable bool) any { +func (v *VariantArray) GetOneForMarshalNullable(i int, nullable bool) any { if nullable && v.IsNull(i) { return nil } @@ -612,6 +612,10 @@ func (v *VariantArray) GetOneForMarshal(i int, nullable bool) any { return val.Value() } +func (v *VariantArray) GetOneForMarshal(i int) any { + return v.GetOneForMarshalNullable(i, true) +} + type variantReader interface { IsNull(i int) bool Value(i int) (variant.Value, error) From a474f1faa92ba5e78b9b316dd242f991bcc99c4c Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 9 Jul 2026 16:18:58 -0400 Subject: [PATCH 17/18] Thread field-local nullability through list, union, and REE Address review findings on incomplete nullability propagation: - List marshaling (List/LargeList/ListView/LargeListView/FixedSizeList): marshal child slices element-by-element via getOneForMarshalNullable so a non-nullable element field serializes underlying values instead of JSON null at null child slots. - Exact sparse/dense union equality: thread equalOption and compare the active child with its field-local nullability (matching the approximate path) instead of the default-option public SliceEqual. - Run-end-encoded equality: compare the values child using opt.nullable && ValueNullable, matching the marshaler so JSON round trips stay consistent, instead of propagating the parent field nullability to the values child. --- arrow/array/compare.go | 4 ++-- arrow/array/encoded.go | 6 ++++-- arrow/array/fixed_size_list.go | 7 +------ arrow/array/list.go | 24 ++++-------------------- arrow/array/union.go | 14 ++++++++------ arrow/array/util.go | 16 ++++++++++++++++ 6 files changed, 35 insertions(+), 36 deletions(-) diff --git a/arrow/array/compare.go b/arrow/array/compare.go index 348df5009..119a77596 100644 --- a/arrow/array/compare.go +++ b/arrow/array/compare.go @@ -400,10 +400,10 @@ func equal(left, right arrow.Array, opt equalOption) bool { return arrayEqualDict(l, r, opt) case *SparseUnion: r := right.(*SparseUnion) - return arraySparseUnionEqual(l, r) + return arraySparseUnionEqual(l, r, opt) case *DenseUnion: r := right.(*DenseUnion) - return arrayDenseUnionEqual(l, r) + return arrayDenseUnionEqual(l, r, opt) case *RunEndEncoded: r := right.(*RunEndEncoded) return arrayRunEndEncodedEqual(l, r, opt) diff --git a/arrow/array/encoded.go b/arrow/array/encoded.go index 441bf9ded..11c2064e0 100644 --- a/arrow/array/encoded.go +++ b/arrow/array/encoded.go @@ -270,11 +270,12 @@ func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { func arrayRunEndEncodedEqual(l, r *RunEndEncoded, opt equalOption) bool { // types were already checked before getting here, so we know // the encoded types are equal + childOpt := withNullable(opt, opt.nullable && l.DataType().(*arrow.RunEndEncodedType).ValueNullable) mr := encoded.NewMergedRuns([2]arrow.Array{l, r}) for mr.Next() { lIndex := mr.IndexIntoArray(0) rIndex := mr.IndexIntoArray(1) - if !sliceEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1, opt) { + if !sliceEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1, childOpt) { return false } } @@ -284,11 +285,12 @@ func arrayRunEndEncodedEqual(l, r *RunEndEncoded, opt equalOption) bool { func arrayRunEndEncodedApproxEqual(l, r *RunEndEncoded, opt equalOption) bool { // types were already checked before getting here, so we know // the encoded types are equal + childOpt := withNullable(opt, opt.nullable && l.DataType().(*arrow.RunEndEncodedType).ValueNullable) mr := encoded.NewMergedRuns([2]arrow.Array{l, r}) for mr.Next() { lIndex := mr.IndexIntoArray(0) rIndex := mr.IndexIntoArray(1) - if !sliceApproxEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1, opt) { + if !sliceApproxEqual(l.values, lIndex, lIndex+1, r.values, rIndex, rIndex+1, childOpt) { return false } } diff --git a/arrow/array/fixed_size_list.go b/arrow/array/fixed_size_list.go index 07fbe4411..65d0654cb 100644 --- a/arrow/array/fixed_size_list.go +++ b/arrow/array/fixed_size_list.go @@ -130,12 +130,7 @@ func (a *FixedSizeList) GetOneForMarshalNullable(i int, nullable bool) interface } slice := a.newListValue(i) defer slice.Release() - v, err := json.Marshal(slice) - if err != nil { - panic(err) - } - - return json.RawMessage(v) + return marshalListElemsNullable(slice, a.DataType().(arrow.ListLikeType).ElemField().Nullable) } func (a *FixedSizeList) GetOneForMarshal(i int) interface{} { diff --git a/arrow/array/list.go b/arrow/array/list.go index 7e6348247..5ce57a341 100644 --- a/arrow/array/list.go +++ b/arrow/array/list.go @@ -105,11 +105,7 @@ func (a *List) GetOneForMarshalNullable(i int, nullable bool) interface{} { slice := a.newListValue(i) defer slice.Release() - v, err := json.Marshal(slice) - if err != nil { - panic(err) - } - return json.RawMessage(v) + return marshalListElemsNullable(slice, a.DataType().(arrow.ListLikeType).ElemField().Nullable) } func (a *List) GetOneForMarshal(i int) interface{} { @@ -242,11 +238,7 @@ func (a *LargeList) GetOneForMarshalNullable(i int, nullable bool) interface{} { slice := a.newListValue(i) defer slice.Release() - v, err := json.Marshal(slice) - if err != nil { - panic(err) - } - return json.RawMessage(v) + return marshalListElemsNullable(slice, a.DataType().(arrow.ListLikeType).ElemField().Nullable) } func (a *LargeList) GetOneForMarshal(i int) interface{} { @@ -722,11 +714,7 @@ func (a *ListView) GetOneForMarshalNullable(i int, nullable bool) interface{} { slice := a.newListValue(i) defer slice.Release() - v, err := json.Marshal(slice) - if err != nil { - panic(err) - } - return json.RawMessage(v) + return marshalListElemsNullable(slice, a.DataType().(arrow.ListLikeType).ElemField().Nullable) } func (a *ListView) GetOneForMarshal(i int) interface{} { @@ -874,11 +862,7 @@ func (a *LargeListView) GetOneForMarshalNullable(i int, nullable bool) interface slice := a.newListValue(i) defer slice.Release() - v, err := json.Marshal(slice) - if err != nil { - panic(err) - } - return json.RawMessage(v) + return marshalListElemsNullable(slice, a.DataType().(arrow.ListLikeType).ElemField().Nullable) } func (a *LargeListView) GetOneForMarshal(i int) interface{} { diff --git a/arrow/array/union.go b/arrow/array/union.go index 8e94900c9..f8e00fe74 100644 --- a/arrow/array/union.go +++ b/arrow/array/union.go @@ -445,7 +445,7 @@ func (a *SparseUnion) GetFlattenedField(mem memory.Allocator, index int) (arrow. return MakeFromData(childData), nil } -func arraySparseUnionEqual(l, r *SparseUnion) bool { +func arraySparseUnionEqual(l, r *SparseUnion, opt equalOption) bool { childIDs := l.unionType.ChildIDs() leftCodes, rightCodes := l.RawTypeCodes(), r.RawTypeCodes() @@ -456,8 +456,9 @@ func arraySparseUnionEqual(l, r *SparseUnion) bool { } childNum := childIDs[typeID] - eq := SliceEqual(l.children[childNum], int64(i), int64(i+1), - r.children[childNum], int64(i), int64(i+1)) + childOpt := withNullable(opt, l.unionType.Fields()[childNum].Nullable) + eq := sliceEqual(l.children[childNum], int64(i), int64(i+1), + r.children[childNum], int64(i), int64(i+1), childOpt) if !eq { return false } @@ -693,7 +694,7 @@ func (a *DenseUnion) String() string { return b.String() } -func arrayDenseUnionEqual(l, r *DenseUnion) bool { +func arrayDenseUnionEqual(l, r *DenseUnion, opt equalOption) bool { childIDs := l.unionType.ChildIDs() leftCodes, rightCodes := l.RawTypeCodes(), r.RawTypeCodes() leftOffsets, rightOffsets := l.RawValueOffsets(), r.RawValueOffsets() @@ -705,8 +706,9 @@ func arrayDenseUnionEqual(l, r *DenseUnion) bool { } childNum := childIDs[typeID] - eq := SliceEqual(l.children[childNum], int64(leftOffsets[i]), int64(leftOffsets[i]+1), - r.children[childNum], int64(rightOffsets[i]), int64(rightOffsets[i]+1)) + childOpt := withNullable(opt, l.unionType.Fields()[childNum].Nullable) + eq := sliceEqual(l.children[childNum], int64(leftOffsets[i]), int64(leftOffsets[i]+1), + r.children[childNum], int64(rightOffsets[i]), int64(rightOffsets[i]+1), childOpt) if !eq { return false } diff --git a/arrow/array/util.go b/arrow/array/util.go index 0461e1717..83b892bab 100644 --- a/arrow/array/util.go +++ b/arrow/array/util.go @@ -321,6 +321,22 @@ func getOneForMarshalNullable(a arrow.Array, i int, nullable bool) interface{} { return a.GetOneForMarshal(i) } +// marshalListElemsNullable marshals a single list element's child slice +// element-by-element, honoring the element field's nullability so a +// non-nullable element field serializes underlying values instead of JSON +// null at null child slots (matching struct/record field-local behavior). +func marshalListElemsNullable(slice arrow.Array, elemNullable bool) json.RawMessage { + vals := make([]interface{}, slice.Len()) + for k := 0; k < slice.Len(); k++ { + vals[k] = getOneForMarshalNullable(slice, k, elemNullable) + } + v, err := json.Marshal(vals) + if err != nil { + panic(err) + } + return json.RawMessage(v) +} + func TableFromJSON(mem memory.Allocator, sc *arrow.Schema, recJSON []string, opt ...FromJSONOption) (arrow.Table, error) { batches := make([]arrow.RecordBatch, len(recJSON)) for i, batchJSON := range recJSON { From 45fef5ef7223bd036da1fc400597c1f4cb14e4a4 Mon Sep 17 00:00:00 2001 From: Matt Topol Date: Thu, 9 Jul 2026 16:23:21 -0400 Subject: [PATCH 18/18] Make run-end-encoded value comparison symmetric in ValueNullable arrow.TypeEqual ignores RunEndEncodedType.ValueNullable, so two REE arrays with matching run-ends/values types but differing ValueNullable can reach the value comparison. Deriving childOpt from only the left array made Equal asymmetric (order-dependent). AND both sides' ValueNullable so the derivation is order-independent; when the two types match (the common case) behavior is unchanged. --- arrow/array/encoded.go | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/arrow/array/encoded.go b/arrow/array/encoded.go index 11c2064e0..2724c1c0f 100644 --- a/arrow/array/encoded.go +++ b/arrow/array/encoded.go @@ -270,7 +270,9 @@ func (r *RunEndEncoded) MarshalJSON() ([]byte, error) { func arrayRunEndEncodedEqual(l, r *RunEndEncoded, opt equalOption) bool { // types were already checked before getting here, so we know // the encoded types are equal - childOpt := withNullable(opt, opt.nullable && l.DataType().(*arrow.RunEndEncodedType).ValueNullable) + childOpt := withNullable(opt, opt.nullable && + l.DataType().(*arrow.RunEndEncodedType).ValueNullable && + r.DataType().(*arrow.RunEndEncodedType).ValueNullable) mr := encoded.NewMergedRuns([2]arrow.Array{l, r}) for mr.Next() { lIndex := mr.IndexIntoArray(0) @@ -285,7 +287,9 @@ func arrayRunEndEncodedEqual(l, r *RunEndEncoded, opt equalOption) bool { func arrayRunEndEncodedApproxEqual(l, r *RunEndEncoded, opt equalOption) bool { // types were already checked before getting here, so we know // the encoded types are equal - childOpt := withNullable(opt, opt.nullable && l.DataType().(*arrow.RunEndEncodedType).ValueNullable) + childOpt := withNullable(opt, opt.nullable && + l.DataType().(*arrow.RunEndEncodedType).ValueNullable && + r.DataType().(*arrow.RunEndEncodedType).ValueNullable) mr := encoded.NewMergedRuns([2]arrow.Array{l, r}) for mr.Next() { lIndex := mr.IndexIntoArray(0)