From 00ae86606a7d3b58ec7080c40d00479d88975b26 Mon Sep 17 00:00:00 2001 From: Neelesh Salian Date: Sat, 15 Aug 2026 17:43:05 -0700 Subject: [PATCH 1/5] feat(extensions): add VariantGet for path extraction from variant arrays --- arrow/extensions/variant.go | 80 +-- arrow/extensions/variant_get.go | 454 ++++++++++++++++++ arrow/extensions/variant_get_internal_test.go | 119 +++++ arrow/extensions/variant_get_test.go | 293 +++++++++++ 4 files changed, 911 insertions(+), 35 deletions(-) create mode 100644 arrow/extensions/variant_get.go create mode 100644 arrow/extensions/variant_get_internal_test.go create mode 100644 arrow/extensions/variant_get_test.go diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index fee2e046a..c35cc07e0 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -1470,85 +1470,96 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { return v.Bytes() } - switch bldr := b.typedBldr.(type) { + if appendVariantToTypedBuilder(b.typedBldr, v) { + return nil + } + + b.typedBldr.AppendNull() + return v.Bytes() +} + +// appendVariantToTypedBuilder appends v to a typed primitive builder when v's type +// fits the builder, reporting whether it did. Shared by the shredding writer and VariantGet. +func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { + switch bldr := target.(type) { case *array.Int8Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Uint8Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Int16Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Uint16Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Int32Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Uint32Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Int64Builder: if appendNumericToTarget(bldr, v) { - return nil + return true } case *array.Float32Builder: switch v.Type() { case variant.Float: bldr.Append(v.Value().(float32)) - return nil + return true case variant.Double: val := v.Value().(float64) if val >= -math.MaxFloat32 && val <= math.MaxFloat32 { bldr.Append(float32(val)) - return nil + return true } } case *array.Float64Builder: switch v.Type() { case variant.Float: bldr.Append(float64(v.Value().(float32))) - return nil + return true case variant.Double: bldr.Append(v.Value().(float64)) - return nil + return true } case *array.BooleanBuilder: if v.Type() == variant.Bool { bldr.Append(v.Value().(bool)) - return nil + return true } case array.StringLikeBuilder: if v.Type() == variant.String { bldr.Append(v.Value().(string)) - return nil + return true } case array.BinaryLikeBuilder: if v.Type() == variant.Binary { bldr.Append(v.Value().([]byte)) - return nil + return true } case *array.Date32Builder: if v.Type() == variant.Date { bldr.Append(v.Value().(arrow.Date32)) - return nil + return true } case *array.Time64Builder: if v.Type() == variant.Time && bldr.Type().(*arrow.Time64Type).Unit == arrow.Microsecond { bldr.Append(v.Value().(arrow.Time64)) - return nil + return true } case *UUIDBuilder: if v.Type() == variant.UUID { bldr.Append(v.Value().(uuid.UUID)) - return nil + return true } case *array.TimestampBuilder: tsType := bldr.Type().(*arrow.TimestampType) @@ -1561,10 +1572,10 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { switch tsType.Unit { case arrow.Microsecond: bldr.Append(v.Value().(arrow.Timestamp)) - return nil + return true case arrow.Nanosecond: bldr.Append(v.Value().(arrow.Timestamp) * 1000) - return nil + return true } case variant.TimestampMicrosNTZ: if tsType.TimeZone != "" { @@ -1574,20 +1585,20 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { switch tsType.Unit { case arrow.Microsecond: bldr.Append(v.Value().(arrow.Timestamp)) - return nil + return true case arrow.Nanosecond: bldr.Append(v.Value().(arrow.Timestamp) * 1000) - return nil + return true } case variant.TimestampNanos: if tsType.TimeZone == "UTC" && tsType.Unit == arrow.Nanosecond { bldr.Append(v.Value().(arrow.Timestamp)) - return nil + return true } case variant.TimestampNanosNTZ: if tsType.TimeZone == "" && tsType.Unit == arrow.Nanosecond { bldr.Append(v.Value().(arrow.Timestamp)) - return nil + return true } } case *array.Decimal32Builder: @@ -1596,17 +1607,17 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal32)) - return nil + return true } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal32(val.Value.(decimal.Decimal64))) - return nil + return true } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal32(val.Value.(decimal.Decimal128).LowBits())) - return nil + return true } } case *array.Decimal64Builder: @@ -1615,17 +1626,17 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal64(val.Value.(decimal.Decimal32))) - return nil + return true } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal64)) - return nil + return true } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal64(val.Value.(decimal.Decimal128).LowBits())) - return nil + return true } } case *array.Decimal128Builder: @@ -1634,23 +1645,22 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(decimal128.FromI64(int64(val.Value.(decimal.Decimal32)))) - return nil + return true } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(decimal128.FromI64(int64(val.Value.(decimal.Decimal64)))) - return nil + return true } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal128)) - return nil + return true } } } - b.typedBldr.AppendNull() - return v.Bytes() + return false } type shreddedFieldBuilder struct { diff --git a/arrow/extensions/variant_get.go b/arrow/extensions/variant_get.go new file mode 100644 index 000000000..36e76b6c6 --- /dev/null +++ b/arrow/extensions/variant_get.go @@ -0,0 +1,454 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package extensions + +import ( + "fmt" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/bitutil" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/parquet/variant" +) + +// VariantPathElement is a single step of a variant path: either an object field +// name or an array index. +type VariantPathElement struct { + name string + index int + isIndex bool +} + +// VariantPathField returns a path element selecting the named object field. +func VariantPathField(name string) VariantPathElement { + return VariantPathElement{name: name} +} + +// VariantPathIndex returns a path element selecting the array element at index. +func VariantPathIndex(index int) VariantPathElement { + return VariantPathElement{index: index, isIndex: true} +} + +// VariantPath is an ordered list of path elements to extract from a variant value. +type VariantPath []VariantPathElement + +// GetOptions controls VariantGet. +type GetOptions struct { + // Path is the path to extract from each variant value. + Path VariantPath + // AsType, when nil, makes VariantGet return a VariantArray pointing at the path. + // When set, the extracted value is cast to this type. Nested (struct/list) types + // are not yet supported and yield arrow.ErrNotImplemented. + AsType arrow.DataType + // Safe makes cast failures produce null; when false a failure returns an error. + Safe bool + // Mem is the allocator for output arrays; nil uses memory.DefaultAllocator. + Mem memory.Allocator +} + +// VariantGet extracts opts.Path from each value of a VariantArray. It follows the +// shredded typed_value columns as far as the path allows, then falls back to a +// per-row walk of the residual value for the remainder. +func VariantGet(input arrow.Array, opts GetOptions) (arrow.Array, error) { + va, ok := input.(*VariantArray) + if !ok { + return nil, fmt.Errorf("%w: VariantGet input must be a VariantArray, got %T", arrow.ErrInvalid, input) + } + + if opts.Mem == nil { + opts.Mem = memory.DefaultAllocator + } + + return shreddedGetPath(va, opts) +} + +// shreddingState is a (value?, typed_value?) column pair at one level of a shredded +// variant, mirroring arrow-rs ShreddingState. +type shreddingState struct { + value arrow.TypedArray[[]byte] + typedValue arrow.Array + length int +} + +func stateFromVariant(va *VariantArray) shreddingState { + vt := va.ExtensionType().(*VariantType) + st := va.Storage().(*array.Struct) + + var value arrow.TypedArray[[]byte] + if vt.valueFieldIdx != -1 { + value = st.Field(vt.valueFieldIdx).(arrow.TypedArray[[]byte]) + } + + var typed arrow.Array + if vt.typedValueFieldIdx != -1 { + typed = st.Field(vt.typedValueFieldIdx) + } + + return shreddingState{value: value, typedValue: typed, length: va.Len()} +} + +func stateFromFieldStruct(child *array.Struct) shreddingState { + ct := child.DataType().(*arrow.StructType) + + var value arrow.TypedArray[[]byte] + if idx, ok := ct.FieldIdx("value"); ok { + value = child.Field(idx).(arrow.TypedArray[[]byte]) + } + + var typed arrow.Array + if idx, ok := ct.FieldIdx("typed_value"); ok { + typed = child.Field(idx) + } + + return shreddingState{value: value, typedValue: typed, length: child.Len()} +} + +type pathStepKind int + +const ( + stepSuccess pathStepKind = iota + stepMissing + stepNotShredded +) + +type pathStep struct { + kind pathStepKind + state shreddingState +} + +// missingStep decides whether an absent typed field means the value is provably +// missing (value column all-null) or merely not shredded (residual may hold it). +func (s shreddingState) missingStep() pathStep { + if s.value == nil || s.value.NullN() == s.value.Len() { + return pathStep{kind: stepMissing} + } + + return pathStep{kind: stepNotShredded} +} + +// followFieldElement takes one field step deeper into the shredded columns. +func followFieldElement(s shreddingState, name string) (pathStep, error) { + if s.typedValue == nil { + return s.missingStep(), nil + } + + st, ok := s.typedValue.(*array.Struct) + if !ok { + return s.missingStep(), nil + } + + idx, ok := st.DataType().(*arrow.StructType).FieldIdx(name) + if !ok { + return s.missingStep(), nil + } + + child, ok := st.Field(idx).(*array.Struct) + if !ok { + return pathStep{}, fmt.Errorf("%w: expected struct field %q while following path, got %s", + arrow.ErrInvalid, name, st.Field(idx).DataType()) + } + + return pathStep{kind: stepSuccess, state: stateFromFieldStruct(child)}, nil +} + +func shreddedGetPath(va *VariantArray, opts GetOptions) (arrow.Array, error) { + state := stateFromVariant(va) + nulls := newNullTracker(va.Len()) + nulls.apply(va.Storage()) + + // Peel the field prefix of the path through the shredded columns. Index steps + // and non-shredded fields stop the columnar walk and hand the rest to a per-row + // fallback over the fully reassembled value at the current node. + idx := 0 + for idx < len(opts.Path) { + elem := opts.Path[idx] + if elem.isIndex { + break + } + + step, err := followFieldElement(state, elem.name) + if err != nil { + return nil, err + } + + switch step.kind { + case stepSuccess: + nulls.apply(state.typedValue) + state = step.state + idx++ + + continue + case stepMissing: + return allNullResult(va, opts) + } + + break // stepNotShredded + } + + remaining := opts.Path[idx:] + target, err := buildTargetVariant(va, state, nulls, opts.Mem) + if err != nil { + return nil, err + } + defer target.Release() + + if len(remaining) == 0 { + if opts.AsType == nil { + target.Retain() + + return target, nil + } + + if shredded := tryPerfectShredding(state, nulls, opts.AsType); shredded != nil { + return shredded, nil + } + } + + return shredBasicVariant(target, remaining, opts) +} + +// shredBasicVariant walks the remaining path per row and produces either a +// VariantArray (AsType nil) or a typed array. +func shredBasicVariant(target *VariantArray, remaining VariantPath, opts GetOptions) (arrow.Array, error) { + if opts.AsType == nil { + bldr := NewVariantBuilder(opts.Mem, NewDefaultVariantType()) + defer bldr.Release() + bldr.Reserve(target.Len()) + + for i := 0; i < target.Len(); i++ { + leaf, ok, err := navigateRow(target, i, remaining) + if err != nil { + return nil, err + } + if !ok { + bldr.AppendNull() + + continue + } + bldr.Append(leaf) + } + + return bldr.NewArray(), nil + } + + if _, ok := opts.AsType.(arrow.NestedType); ok { + return nil, fmt.Errorf("%w: VariantGet cast to nested type %s", arrow.ErrNotImplemented, opts.AsType) + } + + bldr := array.NewBuilder(opts.Mem, opts.AsType) + defer bldr.Release() + bldr.Reserve(target.Len()) + + for i := 0; i < target.Len(); i++ { + leaf, ok, err := navigateRow(target, i, remaining) + if err != nil { + return nil, err + } + if !ok || leaf.Type() == variant.Null { + bldr.AppendNull() + + continue + } + + if appendVariantToTypedBuilder(bldr, leaf) { + continue + } + + if opts.Safe { + bldr.AppendNull() + + continue + } + + return nil, fmt.Errorf("%w: cannot cast variant %v to %s", arrow.ErrInvalid, leaf.Type(), opts.AsType) + } + + return bldr.NewArray(), nil +} + +// navigateRow reassembles row i of target and walks path into it. It returns +// (value, false) when the row is null or the path is absent. +func navigateRow(target *VariantArray, i int, path VariantPath) (variant.Value, bool, error) { + if target.IsNull(i) { + return variant.Value{}, false, nil + } + + v, err := target.Value(i) + if err != nil { + return variant.Value{}, false, fmt.Errorf("variant: reassembling row %d: %w", i, err) + } + + return navigateValue(v, path) +} + +// navigateValue walks path into a fully reassembled variant value. +func navigateValue(v variant.Value, path VariantPath) (variant.Value, bool, error) { + cur := v + for _, elem := range path { + if elem.isIndex { + arr, ok := cur.Value().(variant.ArrayValue) + if !ok || elem.index < 0 || uint32(elem.index) >= arr.Len() { + return variant.Value{}, false, nil + } + el, err := arr.Value(uint32(elem.index)) + if err != nil { + return variant.Value{}, false, nil + } + cur = el + + continue + } + + obj, ok := cur.Value().(variant.ObjectValue) + if !ok { + return variant.Value{}, false, nil + } + field, err := obj.ValueByKey(elem.name) + if err != nil { + return variant.Value{}, false, nil + } + cur = field.Value + } + + return cur, true, nil +} + +// tryPerfectShredding returns the typed_value column directly when the target is +// perfectly shredded to AsType. It only fires when no ancestor nulls need merging; +// otherwise the caller's per-row path produces the same values. +func tryPerfectShredding(state shreddingState, nulls *nullTracker, asType arrow.DataType) arrow.Array { + if _, ok := asType.(arrow.NestedType); ok { + return nil + } + if state.typedValue == nil || !nulls.allValid() { + return nil + } + if !arrow.TypeEqual(state.typedValue.DataType(), asType) { + return nil + } + if state.value != nil && state.value.NullN() != state.value.Len() { + return nil + } + + state.typedValue.Retain() + + return state.typedValue +} + +// buildTargetVariant wraps the current shredding state as a VariantArray, carrying +// the accumulated ancestor nulls onto the storage struct. +func buildTargetVariant(va *VariantArray, state shreddingState, nulls *nullTracker, mem memory.Allocator) (*VariantArray, error) { + // Take the raw metadata array (not va.Metadata) so dictionary- or large-binary- + // encoded metadata is preserved and decoded by the target's own reader. + srcVT := va.ExtensionType().(*VariantType) + metadata := va.Storage().(*array.Struct).Field(srcVT.metadataFieldIdx) + + fields := []arrow.Field{{Name: "metadata", Type: metadata.DataType(), Nullable: false}} + cols := []arrow.Array{metadata} + + if state.value != nil { + fields = append(fields, arrow.Field{Name: "value", Type: state.value.DataType(), Nullable: true}) + cols = append(cols, state.value) + } + if state.typedValue != nil { + fields = append(fields, arrow.Field{Name: "typed_value", Type: state.typedValue.DataType(), Nullable: true}) + cols = append(cols, state.typedValue) + } + + bitmap, nullCount := nulls.bitmap(mem) + if bitmap != nil { + defer bitmap.Release() + } + + st, err := array.NewStructArrayWithFieldsAndNulls(cols, fields, bitmap, nullCount, 0) + if err != nil { + return nil, err + } + defer st.Release() + + vt, err := NewVariantType(st.DataType()) + if err != nil { + return nil, err + } + + return array.NewExtensionArrayWithStorage(vt, st).(*VariantArray), nil +} + +// allNullResult builds the all-null output for a provably missing path. +func allNullResult(va *VariantArray, opts GetOptions) (arrow.Array, error) { + if opts.AsType != nil { + return array.MakeArrayOfNull(opts.Mem, opts.AsType, va.Len()), nil + } + + bldr := NewVariantBuilder(opts.Mem, NewDefaultVariantType()) + defer bldr.Release() + for i := 0; i < va.Len(); i++ { + bldr.AppendNull() + } + + return bldr.NewArray(), nil +} + +// nullTracker accumulates ancestor null masks encountered while walking the path. +type nullTracker struct { + length int + valid []bool // nil means all valid +} + +func newNullTracker(length int) *nullTracker { + return &nullTracker{length: length} +} + +func (n *nullTracker) apply(arr arrow.Array) { + if arr == nil || arr.NullN() == 0 { + return + } + if n.valid == nil { + n.valid = make([]bool, n.length) + for i := range n.valid { + n.valid[i] = true + } + } + for i := 0; i < n.length; i++ { + if arr.IsNull(i) { + n.valid[i] = false + } + } +} + +func (n *nullTracker) allValid() bool { return n.valid == nil } + +func (n *nullTracker) bitmap(mem memory.Allocator) (*memory.Buffer, int) { + if n.valid == nil { + return nil, 0 + } + + buf := memory.NewResizableBuffer(mem) + buf.Resize(int(bitutil.BytesForBits(int64(n.length)))) + nullCount := 0 + for i, v := range n.valid { + if v { + bitutil.SetBit(buf.Bytes(), i) + } else { + bitutil.ClearBit(buf.Bytes(), i) + nullCount++ + } + } + + return buf, nullCount +} diff --git a/arrow/extensions/variant_get_internal_test.go b/arrow/extensions/variant_get_internal_test.go new file mode 100644 index 000000000..4119aa618 --- /dev/null +++ b/arrow/extensions/variant_get_internal_test.go @@ -0,0 +1,119 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package extensions + +import ( + "testing" + + "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/parquet/variant" + "github.com/stretchr/testify/require" +) + +func mkShreddedIntObj(t *testing.T, mem memory.Allocator, v int64) *VariantArray { + t.Helper() + vt := NewShreddedVariantType(arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) + bldr := NewVariantBuilder(mem, vt) + defer bldr.Release() + var vb variant.Builder + require.NoError(t, vb.Append(map[string]any{"a": v})) + val, err := vb.Build() + require.NoError(t, err) + bldr.Append(val) + + return bldr.NewArray().(*VariantArray) +} + +// TestTryPerfectShreddingFires proves the fast path returns the typed_value column +// directly for a perfect shredding, and declines when the residual value has data. +func TestTryPerfectShreddingFires(t *testing.T) { + mem := memory.DefaultAllocator + arr := mkShreddedIntObj(t, mem, 5) + defer arr.Release() + + state := stateFromVariant(arr) + nulls := newNullTracker(arr.Len()) + nulls.apply(arr.Storage()) + + step, err := followFieldElement(state, "a") + require.NoError(t, err) + require.Equal(t, stepSuccess, step.kind) + + out := tryPerfectShredding(step.state, nulls, arrow.PrimitiveTypes.Int64) + require.NotNil(t, out, "perfect shredding must fire for a fully shredded int64 leaf") + defer out.Release() + require.Equal(t, int64(5), out.(*array.Int64).Value(0)) + + // A non-matching target type must decline. + require.Nil(t, tryPerfectShredding(step.state, nulls, arrow.PrimitiveTypes.Int32)) +} + +func TestVariantGetNoLeak(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + + arr := mkShreddedIntObj(t, mem, 9) + + perfect, err := VariantGet(arr, GetOptions{ + Path: VariantPath{VariantPathField("a")}, + AsType: arrow.PrimitiveTypes.Int64, + Mem: mem, + }) + require.NoError(t, err) + perfect.Release() + + variantOut, err := VariantGet(arr, GetOptions{ + Path: VariantPath{VariantPathField("a")}, + Mem: mem, + }) + require.NoError(t, err) + variantOut.Release() + + missing, err := VariantGet(arr, GetOptions{ + Path: VariantPath{VariantPathField("nope")}, + AsType: arrow.PrimitiveTypes.Int64, + Mem: mem, + }) + require.NoError(t, err) + missing.Release() + + arr.Release() + + // Ancestor-null case: forces nullTracker to allocate a real bitmap buffer that + // buildTargetVariant threads onto the target struct. + vt := NewShreddedVariantType(arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) + nb := NewVariantBuilder(mem, vt) + var vb variant.Builder + require.NoError(t, vb.Append(map[string]any{"a": int64(1)})) + val, err := vb.Build() + require.NoError(t, err) + nb.Append(val) + nb.AppendNull() + withNull := nb.NewArray().(*VariantArray) + nb.Release() + + got, err := VariantGet(withNull, GetOptions{ + Path: VariantPath{VariantPathField("a")}, + Mem: mem, + }) + require.NoError(t, err) + require.True(t, got.(*VariantArray).IsNull(1)) + got.Release() + withNull.Release() +} diff --git a/arrow/extensions/variant_get_test.go b/arrow/extensions/variant_get_test.go new file mode 100644 index 000000000..1f1e5a6d6 --- /dev/null +++ b/arrow/extensions/variant_get_test.go @@ -0,0 +1,293 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package extensions_test + +import ( + "testing" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/extensions" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/parquet/variant" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func mkVariant(t *testing.T, v any) variant.Value { + t.Helper() + var b variant.Builder + require.NoError(t, b.Append(v)) + val, err := b.Build() + require.NoError(t, err) + + return val +} + +// nonShreddedVariants builds a plain (metadata, value) VariantArray from Go values. +func nonShreddedVariants(t *testing.T, mem memory.Allocator, vals ...any) *extensions.VariantArray { + t.Helper() + bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) + defer bldr.Release() + for _, v := range vals { + if v == nil { + bldr.AppendNull() + + continue + } + bldr.Append(mkVariant(t, v)) + } + + return bldr.NewArray().(*extensions.VariantArray) +} + +// shreddedIntObjects builds a VariantArray shredding an object with a single int64 field "a". +func shreddedIntObjects(t *testing.T, mem memory.Allocator, vals ...int64) *extensions.VariantArray { + t.Helper() + vt := extensions.NewShreddedVariantType(arrow.StructOf( + arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) + bldr := extensions.NewVariantBuilder(mem, vt) + defer bldr.Release() + for _, v := range vals { + bldr.Append(mkVariant(t, map[string]any{"a": v})) + } + + return bldr.NewArray().(*extensions.VariantArray) +} + +func TestVariantGetTypedOutput(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, + map[string]any{"a": int64(1), "b": "x"}, + map[string]any{"a": int64(2), "b": "y"}, + nil, + ) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a")}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + ints := out.(*array.Int64) + require.Equal(t, 3, ints.Len()) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 2, ints.Value(1)) + assert.True(t, ints.IsNull(2)) +} + +func TestVariantGetVariantOutput(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, + map[string]any{"a": int64(7)}, + map[string]any{"b": int64(9)}, // no "a" -> null + ) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a")}, + }) + require.NoError(t, err) + defer out.Release() + + varr := out.(*extensions.VariantArray) + require.Equal(t, 2, varr.Len()) + + v, err := varr.Value(0) + require.NoError(t, err) + assert.EqualValues(t, 7, v.Value()) + assert.True(t, varr.IsNull(1)) +} + +func TestVariantGetNestedPath(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, map[string]any{"a": map[string]any{"b": int64(5)}}) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a"), extensions.VariantPathField("b")}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + assert.EqualValues(t, 5, out.(*array.Int64).Value(0)) +} + +func TestVariantGetIndex(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, []any{int64(10), int64(20), int64(30)}) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathIndex(1)}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + assert.EqualValues(t, 20, out.(*array.Int64).Value(0)) + + oob, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathIndex(9)}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer oob.Release() + assert.True(t, oob.(*array.Int64).IsNull(0)) +} + +func TestVariantGetPerfectShredding(t *testing.T) { + mem := memory.DefaultAllocator + arr := shreddedIntObjects(t, mem, 11, 22, 33) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a")}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + ints := out.(*array.Int64) + require.Equal(t, 3, ints.Len()) + assert.EqualValues(t, 11, ints.Value(0)) + assert.EqualValues(t, 22, ints.Value(1)) + assert.EqualValues(t, 33, ints.Value(2)) +} + +func TestVariantGetShreddedVariantOutput(t *testing.T) { + mem := memory.DefaultAllocator + arr := shreddedIntObjects(t, mem, 100, 200) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a")}, + }) + require.NoError(t, err) + defer out.Release() + + varr := out.(*extensions.VariantArray) + v, err := varr.Value(0) + require.NoError(t, err) + assert.EqualValues(t, 100, v.Value()) + v, err = varr.Value(1) + require.NoError(t, err) + assert.EqualValues(t, 200, v.Value()) +} + +func TestVariantGetMissingField(t *testing.T) { + mem := memory.DefaultAllocator + arr := shreddedIntObjects(t, mem, 1, 2) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("missing")}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + ints := out.(*array.Int64) + require.Equal(t, 2, ints.Len()) + assert.True(t, ints.IsNull(0)) + assert.True(t, ints.IsNull(1)) +} + +// TestVariantGetNotShreddedFallback covers a field present only in the residual value +// of a shredded object: the columnar walk stops and the per-row fallback recovers it. +func TestVariantGetNotShreddedFallback(t *testing.T) { + mem := memory.DefaultAllocator + vt := extensions.NewShreddedVariantType(arrow.StructOf( + arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) + bldr := extensions.NewVariantBuilder(mem, vt) + defer bldr.Release() + // "b" is not in the shredding schema, so it lands in the residual value column. + bldr.Append(mkVariant(t, map[string]any{"a": int64(1), "b": int64(42)})) + arr := bldr.NewArray().(*extensions.VariantArray) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("b")}, + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + + assert.EqualValues(t, 42, out.(*array.Int64).Value(0)) +} + +func TestVariantGetNestedTypeUnsupported(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, map[string]any{"a": int64(1)}) + defer arr.Release() + + _, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("a")}, + AsType: arrow.StructOf(arrow.Field{Name: "x", Type: arrow.PrimitiveTypes.Int64}), + }) + require.ErrorIs(t, err, arrow.ErrNotImplemented) +} + +func TestVariantGetRejectsNonVariant(t *testing.T) { + mem := memory.DefaultAllocator + bldr := array.NewInt64Builder(mem) + defer bldr.Release() + bldr.Append(1) + arr := bldr.NewArray() + defer arr.Release() + + _, err := extensions.VariantGet(arr, extensions.GetOptions{}) + require.ErrorIs(t, err, arrow.ErrInvalid) +} + +// TestVariantGetDictMetadata guards against the panic from asserting TypedArray[[]byte] +// on a dictionary-encoded metadata column, which is spec-legal. +func TestVariantGetDictMetadata(t *testing.T) { + mem := memory.DefaultAllocator + s := arrow.StructOf( + arrow.Field{Name: "metadata", Type: &arrow.DictionaryType{ + IndexType: arrow.PrimitiveTypes.Uint8, ValueType: arrow.BinaryTypes.Binary}}, + arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, + arrow.Field{Name: "typed_value", Type: arrow.StructOf( + arrow.Field{Name: "a", Type: arrow.StructOf( + arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, + arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}, + )}, + ), Nullable: true}) + + vt, err := extensions.NewVariantType(s) + require.NoError(t, err) + bldr := vt.NewBuilder(mem).(*extensions.VariantBuilder) + defer bldr.Release() + // "b" is not shredded, so extracting it takes the NotShredded -> buildTargetVariant path. + bldr.Append(mkVariant(t, map[string]any{"a": int64(5), "b": "resid"})) + arr := bldr.NewArray().(*extensions.VariantArray) + defer arr.Release() + + out, err := extensions.VariantGet(arr, extensions.GetOptions{ + Path: extensions.VariantPath{extensions.VariantPathField("b")}, + }) + require.NoError(t, err) + defer out.Release() + + v, err := out.(*extensions.VariantArray).Value(0) + require.NoError(t, err) + assert.Equal(t, "resid", v.Value()) +} From 0bb88a9d56d8728d57952edb5adebcaba24f40c9 Mon Sep 17 00:00:00 2001 From: Neelesh Salian Date: Sat, 15 Aug 2026 20:06:20 -0700 Subject: [PATCH 2/5] guard value-less VariantGet layout and align cast default with arrow-rs --- arrow/extensions/variant.go | 5 +++ arrow/extensions/variant_get.go | 13 ++++---- arrow/extensions/variant_get_test.go | 49 +++++++++++++++++++++++++++- 3 files changed, 59 insertions(+), 8 deletions(-) diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index c35cc07e0..a4f5e3b42 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -513,6 +513,11 @@ func (v *VariantArray) IsNull(i int) bool { } } + if vt.valueFieldIdx == -1 { + // No residual value column: a null typed_value means the value is missing. + return true + } + valArr := v.Storage().(*array.Struct).Field(vt.valueFieldIdx) b := valArr.(arrow.TypedArray[[]byte]).Value(i) return len(b) == 1 && b[0] == 0 // variant null diff --git a/arrow/extensions/variant_get.go b/arrow/extensions/variant_get.go index 36e76b6c6..09f8e79bd 100644 --- a/arrow/extensions/variant_get.go +++ b/arrow/extensions/variant_get.go @@ -55,8 +55,9 @@ type GetOptions struct { // When set, the extracted value is cast to this type. Nested (struct/list) types // are not yet supported and yield arrow.ErrNotImplemented. AsType arrow.DataType - // Safe makes cast failures produce null; when false a failure returns an error. - Safe bool + // Strict makes a cast failure return an error. The default (false) mirrors + // arrow-rs: a cast failure produces null. + Strict bool // Mem is the allocator for output arrays; nil uses memory.DefaultAllocator. Mem memory.Allocator } @@ -269,13 +270,11 @@ func shredBasicVariant(target *VariantArray, remaining VariantPath, opts GetOpti continue } - if opts.Safe { - bldr.AppendNull() - - continue + if opts.Strict { + return nil, fmt.Errorf("%w: cannot cast variant %v to %s", arrow.ErrInvalid, leaf.Type(), opts.AsType) } - return nil, fmt.Errorf("%w: cannot cast variant %v to %s", arrow.ErrInvalid, leaf.Type(), opts.AsType) + bldr.AppendNull() } return bldr.NewArray(), nil diff --git a/arrow/extensions/variant_get_test.go b/arrow/extensions/variant_get_test.go index 1f1e5a6d6..455c4ecf7 100644 --- a/arrow/extensions/variant_get_test.go +++ b/arrow/extensions/variant_get_test.go @@ -257,7 +257,54 @@ func TestVariantGetRejectsNonVariant(t *testing.T) { require.ErrorIs(t, err, arrow.ErrInvalid) } -// TestVariantGetDictMetadata guards against the panic from asserting TypedArray[[]byte] +// TestVariantGetStrictVsDefaultCast covers the cast-outcome axis: an uncastable +// value nulls by default (mirrors arrow-rs) and errors under Strict. +func TestVariantGetStrictVsDefaultCast(t *testing.T) { + mem := memory.DefaultAllocator + arr := nonShreddedVariants(t, mem, map[string]any{"a": "not-a-number"}) + defer arr.Release() + path := extensions.VariantPath{extensions.VariantPathField("a")} + + // Default (Strict false): uncastable string -> Int64 becomes null. + out, err := extensions.VariantGet(arr, extensions.GetOptions{Path: path, AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + assert.True(t, out.(*array.Int64).IsNull(0)) + + // Strict: the same cast returns an error. + _, err = extensions.VariantGet(arr, extensions.GetOptions{Path: path, AsType: arrow.PrimitiveTypes.Int64, Strict: true}) + require.ErrorIs(t, err, arrow.ErrInvalid) +} + +// TestVariantGetValuelessLayout guards against the Field(-1) panic on a shredded +// layout that has no residual value column. +func TestVariantGetValuelessLayout(t *testing.T) { + mem := memory.DefaultAllocator + s := arrow.StructOf( + arrow.Field{Name: "metadata", Type: arrow.BinaryTypes.Binary}, + arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}) + b := array.NewStructBuilder(mem, s) + defer b.Release() + mb := b.FieldBuilder(0).(*array.BinaryBuilder) + tv := b.FieldBuilder(1).(*array.Int64Builder) + b.Append(true) + mb.Append(variant.EmptyMetadataBytes[:]) + tv.AppendNull() + st := b.NewArray() + defer st.Release() + + vt, err := extensions.NewVariantType(s) + require.NoError(t, err) + arr := array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray) + defer arr.Release() + + // AsType mismatch (Int32 vs Int64) forces the per-row fallback over a value-less target. + out, err := extensions.VariantGet(arr, extensions.GetOptions{AsType: arrow.PrimitiveTypes.Int32}) + require.NoError(t, err) + defer out.Release() + assert.True(t, out.(*array.Int32).IsNull(0)) +} + // on a dictionary-encoded metadata column, which is spec-legal. func TestVariantGetDictMetadata(t *testing.T) { mem := memory.DefaultAllocator From 1b2818f5fdea122560a3eeeb32538cdd3473bae2 Mon Sep 17 00:00:00 2001 From: Neelesh Salian Date: Fri, 21 Aug 2026 15:26:20 -0700 Subject: [PATCH 3/5] addressing PR comments --- arrow/compute/variant_get.go | 658 ++++++++++++++++++ arrow/compute/variant_get_test.go | 349 ++++++++++ arrow/extensions/variant.go | 85 ++- arrow/extensions/variant_get.go | 453 ------------ arrow/extensions/variant_get_internal_test.go | 119 ---- arrow/extensions/variant_get_test.go | 340 --------- parquet/variant/path.go | 108 +++ parquet/variant/path_test.go | 94 +++ 8 files changed, 1249 insertions(+), 957 deletions(-) create mode 100644 arrow/compute/variant_get.go create mode 100644 arrow/compute/variant_get_test.go delete mode 100644 arrow/extensions/variant_get.go delete mode 100644 arrow/extensions/variant_get_internal_test.go delete mode 100644 arrow/extensions/variant_get_test.go create mode 100644 parquet/variant/path.go create mode 100644 parquet/variant/path_test.go diff --git a/arrow/compute/variant_get.go b/arrow/compute/variant_get.go new file mode 100644 index 000000000..d23f4ece1 --- /dev/null +++ b/arrow/compute/variant_get.go @@ -0,0 +1,658 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package compute + +import ( + "context" + "fmt" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/bitutil" + "github.com/apache/arrow-go/v18/arrow/decimal" + "github.com/apache/arrow-go/v18/arrow/decimal128" + "github.com/apache/arrow-go/v18/arrow/extensions" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/parquet/variant" + "github.com/google/uuid" +) + +// VariantGetOptions controls VariantGet. +type VariantGetOptions struct { + // Path is the path to extract from every variant value. + Path variant.VariantPath + // AsType, when nil, makes VariantGet return a VariantArray pointing at the path; + // when set, the extracted values are cast to it via the cast kernels. + AsType arrow.DataType + // Strict makes a lossy cast fail; the default allows overflow and truncation via + // the cast kernels. Unlike arrow-rs safe mode there is no null-on-failure: an + // impossible cast always errors, since arrow-go's cast kernels have no safe flag. + Strict bool +} + +// VariantGet extracts opts.Path from every value of input. It follows the shredded +// typed_value columns as far as the path allows - stepping into struct fields +// directly and gathering list elements with the take kernel - then reassembles only +// the residual for any remaining path. With AsType nil it returns a VariantArray of +// the extracted values; otherwise it casts them to AsType with the cast kernels. +func VariantGet(ctx context.Context, input *extensions.VariantArray, opts VariantGetOptions) (arrow.Array, error) { + if input == nil { + return nil, fmt.Errorf("%w: VariantGet requires a non-nil VariantArray", arrow.ErrInvalid) + } + + // Empty path, no cast: the values are returned unchanged. + if opts.Path.Len() == 0 && opts.AsType == nil { + input.Retain() + + return input, nil + } + + return shreddedGetPath(ctx, input, opts) +} + +// shreddingState is a (value?, typed_value?) column pair at one level of a shredded +// variant, mirroring arrow-rs ShreddingState. +type shreddingState struct { + value arrow.TypedArray[[]byte] + typedValue arrow.Array + length int +} + +func stateFromInput(input *extensions.VariantArray) shreddingState { + return shreddingState{ + value: input.UntypedValues(), + typedValue: input.Shredded(), + length: input.Len(), + } +} + +func stateFromFieldStruct(child *array.Struct) shreddingState { + ct := child.DataType().(*arrow.StructType) + + var value arrow.TypedArray[[]byte] + if idx, ok := ct.FieldIdx("value"); ok { + value = child.Field(idx).(arrow.TypedArray[[]byte]) + } + + var typed arrow.Array + if idx, ok := ct.FieldIdx("typed_value"); ok { + typed = child.Field(idx) + } + + return shreddingState{value: value, typedValue: typed, length: child.Len()} +} + +type pathStepKind int + +const ( + stepSuccess pathStepKind = iota + stepMissing + stepNotShredded +) + +type pathStep struct { + kind pathStepKind + state shreddingState + owned []arrow.Array // intermediate take results the caller must release +} + +// missingStep reports whether an absent typed field is provably missing (the value +// column is all-null) or merely not shredded (a residual may hold it). +func (s shreddingState) missingStep() pathStep { + if s.value == nil || s.value.NullN() == s.value.Len() { + return pathStep{kind: stepMissing} + } + + return pathStep{kind: stepNotShredded} +} + +func fieldStep(s shreddingState, name string) (pathStep, error) { + if s.typedValue == nil { + return s.missingStep(), nil + } + st, ok := s.typedValue.(*array.Struct) + if !ok { + return s.missingStep(), nil + } + idx, ok := st.DataType().(*arrow.StructType).FieldIdx(name) + if !ok { + return s.missingStep(), nil + } + child, ok := st.Field(idx).(*array.Struct) + if !ok { + return pathStep{}, fmt.Errorf("%w: expected struct field %q while following path, got %s", + arrow.ErrInvalid, name, st.Field(idx).DataType()) + } + + return pathStep{kind: stepSuccess, state: stateFromFieldStruct(child)}, nil +} + +// indexStep gathers element index from every row of a shredded list with the take +// kernel, producing the shredding state one level deeper. +func indexStep(ctx context.Context, mem memory.Allocator, s shreddingState, index int) (pathStep, error) { + if s.typedValue == nil { + return s.missingStep(), nil + } + list, ok := s.typedValue.(array.ListLike) + if !ok { + return s.missingStep(), nil + } + elems, ok := list.ListValues().(*array.Struct) + if !ok { + return s.missingStep(), nil + } + + ib := array.NewUint64Builder(mem) + defer ib.Release() + ib.Reserve(s.length) + for row := 0; row < s.length; row++ { + start, end := list.ValueOffsets(row) + if list.IsValid(row) && index >= 0 && int64(index) < end-start { + ib.Append(uint64(start + int64(index))) + } else { + ib.AppendNull() + } + } + indices := ib.NewArray() + defer indices.Release() + + et := elems.DataType().(*arrow.StructType) + var owned []arrow.Array + var next shreddingState + next.length = s.length + + if vi, ok := et.FieldIdx("value"); ok { + taken, err := TakeArray(ctx, elems.Field(vi), indices) + if err != nil { + return pathStep{}, err + } + owned = append(owned, taken) + next.value = taken.(arrow.TypedArray[[]byte]) + } + if ti, ok := et.FieldIdx("typed_value"); ok { + taken, err := TakeArray(ctx, elems.Field(ti), indices) + if err != nil { + releaseAll(owned) + + return pathStep{}, err + } + owned = append(owned, taken) + next.typedValue = taken + } + + return pathStep{kind: stepSuccess, state: next, owned: owned}, nil +} + +func releaseAll(arrs []arrow.Array) { + for _, a := range arrs { + a.Release() + } +} + +func shreddedGetPath(ctx context.Context, input *extensions.VariantArray, opts VariantGetOptions) (arrow.Array, error) { + mem := GetAllocator(ctx) + state := stateFromInput(input) + nulls := newNullTracker(input.Len(), mem) + defer nulls.release() + nulls.merge(input.Storage()) + + var owned []arrow.Array + defer func() { releaseAll(owned) }() + + idx := 0 + for idx < opts.Path.Len() { + name, index := opts.Path.StepAt(idx) + var ( + step pathStep + err error + ) + if name != "" { + step, err = fieldStep(state, name) + } else { + step, err = indexStep(ctx, mem, state, index) + } + if err != nil { + return nil, err + } + + if step.kind == stepSuccess { + nulls.merge(state.typedValue) + state = step.state + owned = append(owned, step.owned...) + idx++ + + continue + } + if step.kind == stepMissing { + return allNullResult(mem, input.Len(), opts.AsType), nil + } + + break // stepNotShredded + } + + remaining := subPath(opts.Path, idx) + + // Try to return the typed column directly before building the target array, + // so a perfect shredding does not allocate a struct and bitmap it discards. + if remaining.Len() == 0 && opts.AsType != nil { + if col := perfectShredded(state, nulls, opts.AsType); col != nil { + defer col.Release() + + return CastArray(ctx, col, NewCastOptions(opts.AsType, opts.Strict)) + } + } + + target, err := buildTargetVariant(input, state, nulls, mem) + if err != nil { + return nil, err + } + defer target.Release() + + if remaining.Len() == 0 && opts.AsType == nil { + target.Retain() + + return target, nil + } + + leaves, err := extractLeaves(target, remaining) + if err != nil { + return nil, err + } + if opts.AsType == nil { + return buildLeafVariantArray(mem, leaves), nil + } + + src := buildNaturalArray(mem, leaves) + if src == nil { + return allNullResult(mem, len(leaves), opts.AsType), nil + } + defer src.Release() + + return CastArray(ctx, src, NewCastOptions(opts.AsType, opts.Strict)) +} + +// perfectShredded returns the typed_value column when the path landed on a fully +// shredded value of exactly AsType and no ancestor nulls need merging; otherwise +// the caller's reassembly path produces the same values. +func perfectShredded(s shreddingState, nulls *nullTracker, asType arrow.DataType) arrow.Array { + if _, ok := asType.(arrow.NestedType); ok { + return nil + } + if s.typedValue == nil || !nulls.allValid() { + return nil + } + if !arrow.TypeEqual(s.typedValue.DataType(), asType) { + return nil + } + if s.value != nil && s.value.NullN() != s.value.Len() { + return nil + } + + s.typedValue.Retain() + + return s.typedValue +} + +func buildTargetVariant(input *extensions.VariantArray, s shreddingState, nulls *nullTracker, mem memory.Allocator) (*extensions.VariantArray, error) { + // Read the raw metadata column rather than input.Metadata(), which asserts plain + // binary and panics on dictionary-encoded metadata; the raw column preserves + // dictionary/large-binary encoding and is decoded by the target's own reader. + storage := input.Storage().(*array.Struct) + mdIdx, ok := storage.DataType().(*arrow.StructType).FieldIdx("metadata") + if !ok { + return nil, fmt.Errorf("%w: variant storage is missing its metadata field", arrow.ErrInvalid) + } + metadata := storage.Field(mdIdx) + + fields := []arrow.Field{{Name: "metadata", Type: metadata.DataType(), Nullable: false}} + cols := []arrow.Array{metadata} + if s.value != nil { + fields = append(fields, arrow.Field{Name: "value", Type: s.value.DataType(), Nullable: true}) + cols = append(cols, s.value) + } + if s.typedValue != nil { + fields = append(fields, arrow.Field{Name: "typed_value", Type: s.typedValue.DataType(), Nullable: true}) + cols = append(cols, s.typedValue) + } + + bitmap, nullCount := nulls.validityBitmap() + st, err := array.NewStructArrayWithFieldsAndNulls(cols, fields, bitmap, nullCount, 0) + if err != nil { + return nil, err + } + defer st.Release() + + vt, err := extensions.NewVariantType(st.DataType()) + if err != nil { + return nil, err + } + + return array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray), nil +} + +// subPath returns the suffix of p starting at from, rebuilt through the opaque API. +func subPath(p variant.VariantPath, from int) variant.VariantPath { + var out variant.VariantPath + for i := from; i < p.Len(); i++ { + if name, index := p.StepAt(i); name != "" { + out = out.Field(name) + } else { + out = out.Index(index) + } + } + + return out +} + +// variantLeaf is one row's extracted value; present is false when the path is +// absent for that row (or the row is null). +type variantLeaf struct { + value variant.Value + present bool +} + +func extractLeaves(target *extensions.VariantArray, path variant.VariantPath) ([]variantLeaf, error) { + leaves := make([]variantLeaf, target.Len()) + for i := range leaves { + if target.IsNull(i) { + continue + } + v, err := target.Value(i) + if err != nil { + return nil, fmt.Errorf("variant: reassembling row %d: %w", i, err) + } + leaf, found, err := v.GetByPath(path) + if err != nil { + return nil, err + } + leaves[i] = variantLeaf{value: leaf, present: found} + } + + return leaves, nil +} + +func buildLeafVariantArray(mem memory.Allocator, leaves []variantLeaf) arrow.Array { + bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) + defer bldr.Release() + bldr.Reserve(len(leaves)) + for _, l := range leaves { + if !l.present { + bldr.AppendNull() + + continue + } + bldr.Append(l.value) + } + + return bldr.NewArray() +} + +// buildNaturalArray materializes the leaves as an array of the first present leaf's +// natural Arrow type so the cast kernels can convert it. Rows whose value does not +// match that natural type become null. Returns nil when no leaf is present. +func buildNaturalArray(mem memory.Allocator, leaves []variantLeaf) arrow.Array { + var natural arrow.DataType + for _, l := range leaves { + if l.present && l.value.Type() != variant.Null { + natural = naturalArrowType(l.value) + + break + } + } + if natural == nil { + return nil + } + + bldr := array.NewBuilder(mem, natural) + defer bldr.Release() + bldr.Reserve(len(leaves)) + for _, l := range leaves { + if !l.present || !appendNatural(bldr, l.value) { + bldr.AppendNull() + } + } + + return bldr.NewArray() +} + +// allNullResult builds the all-null output for a provably missing path. +func allNullResult(mem memory.Allocator, n int, asType arrow.DataType) arrow.Array { + if asType != nil { + return array.MakeArrayOfNull(mem, asType, n) + } + + // MakeArrayOfNull cannot build the variant extension type (its storage struct's + // metadata/value are non-nullable), so append encoded variant nulls instead. + bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) + defer bldr.Release() + for range n { + bldr.AppendNull() + } + + return bldr.NewArray() +} + +// nullTracker accumulates ancestor validity bitmaps with a bitmap AND. +type nullTracker struct { + length int + mem memory.Allocator + buf *memory.Buffer // validity bitmap (1 = valid); nil means all valid +} + +func newNullTracker(length int, mem memory.Allocator) *nullTracker { + return &nullTracker{length: length, mem: mem} +} + +// merge folds arr's validity into the accumulated mask. Arrow validity bits are +// 1=valid, so accumulating ancestor nulls is a bitmap AND (a row is null in the +// result when it is null at any level) - the validity-space equivalent of OR-ing +// null masks. +func (n *nullTracker) merge(arr arrow.Array) { + if arr == nil { + return + } + vb := arr.Data().Buffers()[0] + if vb == nil { + return // all valid + } + off := int64(arr.Data().Offset()) + if n.buf == nil { + n.buf = bitutil.BitmapAndAlloc(n.mem, vb.Bytes(), vb.Bytes(), off, off, int64(n.length), 0) + + return + } + merged := bitutil.BitmapAndAlloc(n.mem, n.buf.Bytes(), vb.Bytes(), 0, off, int64(n.length), 0) + n.buf.Release() + n.buf = merged +} + +func (n *nullTracker) allValid() bool { return n.buf == nil } + +func (n *nullTracker) validityBitmap() (*memory.Buffer, int) { + if n.buf == nil { + return nil, 0 + } + + return n.buf, n.length - bitutil.CountSetBits(n.buf.Bytes(), 0, n.length) +} + +func (n *nullTracker) release() { + if n.buf != nil { + n.buf.Release() + n.buf = nil + } +} + +func naturalArrowType(v variant.Value) arrow.DataType { + switch v.Type() { + case variant.Bool: + return arrow.FixedWidthTypes.Boolean + case variant.Int8: + return arrow.PrimitiveTypes.Int8 + case variant.Int16: + return arrow.PrimitiveTypes.Int16 + case variant.Int32: + return arrow.PrimitiveTypes.Int32 + case variant.Int64: + return arrow.PrimitiveTypes.Int64 + case variant.Float: + return arrow.PrimitiveTypes.Float32 + case variant.Double: + return arrow.PrimitiveTypes.Float64 + case variant.String: + return arrow.BinaryTypes.String + case variant.Binary: + return arrow.BinaryTypes.Binary + case variant.Date: + return arrow.FixedWidthTypes.Date32 + case variant.Time: + return arrow.FixedWidthTypes.Time64us + case variant.TimestampMicros: + return &arrow.TimestampType{Unit: arrow.Microsecond, TimeZone: "UTC"} + case variant.TimestampMicrosNTZ: + return &arrow.TimestampType{Unit: arrow.Microsecond} + case variant.TimestampNanos: + return &arrow.TimestampType{Unit: arrow.Nanosecond, TimeZone: "UTC"} + case variant.TimestampNanosNTZ: + return &arrow.TimestampType{Unit: arrow.Nanosecond} + case variant.UUID: + return extensions.NewUUIDType() + case variant.Decimal4, variant.Decimal8, variant.Decimal16: + return &arrow.Decimal128Type{Precision: 38, Scale: int32(decimalScale(v))} + } + + return nil +} + +// appendNatural appends v to bldr when v matches bldr's natural type, reporting +// whether it did; a non-matching value is left for the caller to null. +func appendNatural(bldr array.Builder, v variant.Value) bool { + switch b := bldr.(type) { + case *array.BooleanBuilder: + if x, ok := v.Value().(bool); ok { + b.Append(x) + + return true + } + case *array.Int8Builder: + if x, ok := v.Value().(int8); ok { + b.Append(x) + + return true + } + case *array.Int16Builder: + if x, ok := v.Value().(int16); ok { + b.Append(x) + + return true + } + case *array.Int32Builder: + if x, ok := v.Value().(int32); ok { + b.Append(x) + + return true + } + case *array.Int64Builder: + if x, ok := v.Value().(int64); ok { + b.Append(x) + + return true + } + case *array.Float32Builder: + if x, ok := v.Value().(float32); ok { + b.Append(x) + + return true + } + case *array.Float64Builder: + if x, ok := v.Value().(float64); ok { + b.Append(x) + + return true + } + case *array.StringBuilder: + if x, ok := v.Value().(string); ok { + b.Append(x) + + return true + } + case *array.BinaryBuilder: + if x, ok := v.Value().([]byte); ok { + b.Append(x) + + return true + } + case *array.Date32Builder: + if x, ok := v.Value().(arrow.Date32); ok { + b.Append(x) + + return true + } + case *array.Time64Builder: + if x, ok := v.Value().(arrow.Time64); ok { + b.Append(x) + + return true + } + case *array.TimestampBuilder: + if x, ok := v.Value().(arrow.Timestamp); ok { + b.Append(x) + + return true + } + case *extensions.UUIDBuilder: + if x, ok := v.Value().(uuid.UUID); ok { + b.Append(x) + + return true + } + case *array.Decimal128Builder: + if num, ok := decimalAsNum128(v); ok && int32(decimalScale(v)) == b.Type().(*arrow.Decimal128Type).Scale { + b.Append(num) + + return true + } + } + + return false +} + +func decimalScale(v variant.Value) uint8 { + switch d := v.Value().(type) { + case variant.DecimalValue[decimal.Decimal32]: + return d.Scale + case variant.DecimalValue[decimal.Decimal64]: + return d.Scale + case variant.DecimalValue[decimal.Decimal128]: + return d.Scale + } + + return 0 +} + +func decimalAsNum128(v variant.Value) (decimal128.Num, bool) { + switch d := v.Value().(type) { + case variant.DecimalValue[decimal.Decimal32]: + return decimal128.FromI64(int64(d.Value.(decimal.Decimal32))), true + case variant.DecimalValue[decimal.Decimal64]: + return decimal128.FromI64(int64(d.Value.(decimal.Decimal64))), true + case variant.DecimalValue[decimal.Decimal128]: + return d.Value.(decimal.Decimal128), true + } + + return decimal128.Num{}, false +} diff --git a/arrow/compute/variant_get_test.go b/arrow/compute/variant_get_test.go new file mode 100644 index 000000000..993463af9 --- /dev/null +++ b/arrow/compute/variant_get_test.go @@ -0,0 +1,349 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package compute_test + +import ( + "context" + "testing" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/compute" + "github.com/apache/arrow-go/v18/arrow/compute/exec" + "github.com/apache/arrow-go/v18/arrow/extensions" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/parquet/variant" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func vgVariant(t *testing.T, v any) variant.Value { + t.Helper() + var b variant.Builder + require.NoError(t, b.Append(v)) + val, err := b.Build() + require.NoError(t, err) + + return val +} + +func vgNonShredded(t *testing.T, mem memory.Allocator, vals ...any) *extensions.VariantArray { + t.Helper() + bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) + defer bldr.Release() + for _, v := range vals { + if v == nil { + bldr.AppendNull() + + continue + } + bldr.Append(vgVariant(t, v)) + } + + return bldr.NewArray().(*extensions.VariantArray) +} + +func vgShreddedInt(t *testing.T, mem memory.Allocator, vals ...int64) *extensions.VariantArray { + t.Helper() + vt := extensions.NewShreddedVariantType(arrow.PrimitiveTypes.Int64) + bldr := extensions.NewVariantBuilder(mem, vt) + defer bldr.Release() + for _, v := range vals { + bldr.Append(vgVariant(t, v)) + } + + return bldr.NewArray().(*extensions.VariantArray) +} + +func field(name string) variant.VariantPath { return variant.VariantPath{}.Field(name) } + +func TestVariantGetTyped(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, + map[string]any{"a": int64(1)}, + map[string]any{"a": int64(2)}, + nil, + ) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + + ints := out.(*array.Int64) + require.Equal(t, 3, ints.Len()) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 2, ints.Value(1)) + assert.True(t, ints.IsNull(2)) +} + +func TestVariantGetVariantOutput(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, map[string]any{"a": int64(7)}, map[string]any{"b": int64(9)}) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a")}) + require.NoError(t, err) + defer out.Release() + + varr := out.(*extensions.VariantArray) + v, err := varr.Value(0) + require.NoError(t, err) + assert.EqualValues(t, 7, v.Value()) + assert.True(t, varr.IsNull(1)) +} + +func TestVariantGetNestedAndIndex(t *testing.T) { + mem := memory.DefaultAllocator + nested := vgNonShredded(t, mem, map[string]any{"a": map[string]any{"b": int64(5)}}) + defer nested.Release() + out, err := compute.VariantGet(context.Background(), nested, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Field("a").Field("b"), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + assert.EqualValues(t, 5, out.(*array.Int64).Value(0)) + + arrs := vgNonShredded(t, mem, []any{int64(10), int64(20), int64(30)}) + defer arrs.Release() + got, err := compute.VariantGet(context.Background(), arrs, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(1), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer got.Release() + assert.EqualValues(t, 20, got.(*array.Int64).Value(0)) + + oob, err := compute.VariantGet(context.Background(), arrs, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(9), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer oob.Release() + assert.True(t, oob.(*array.Int64).IsNull(0)) +} + +// TestVariantGetMixedShreddedRows reproduces zeroshade's [1,2] case: row 0 is in +// typed_value, row 1 is in the residual value. Both must come back, not [1,null]. +func TestVariantGetMixedShreddedRows(t *testing.T) { + mem := memory.DefaultAllocator + s := arrow.StructOf( + arrow.Field{Name: "metadata", Type: arrow.BinaryTypes.Binary}, + arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, + arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}) + b := array.NewStructBuilder(mem, s) + defer b.Release() + mb := b.FieldBuilder(0).(*array.BinaryBuilder) + vb := b.FieldBuilder(1).(*array.BinaryBuilder) + tb := b.FieldBuilder(2).(*array.Int64Builder) + + b.Append(true) + mb.Append(variant.EmptyMetadataBytes[:]) + vb.AppendNull() + tb.Append(1) + + b.Append(true) + mb.Append(variant.EmptyMetadataBytes[:]) + enc, err := variant.Encode(int64(2)) + require.NoError(t, err) + vb.Append(enc) + tb.AppendNull() + + st := b.NewArray() + defer st.Release() + vt, err := extensions.NewVariantType(s) + require.NoError(t, err) + arr := array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + require.Equal(t, 2, ints.Len()) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 2, ints.Value(1), "residual-value row must be reconstructed, not null") + assert.False(t, ints.IsNull(1)) +} + +// TestVariantGetLenientCast covers zeroshade's :269 examples: ordinary widening +// casts succeed under the default (non-strict) mode. +func TestVariantGetLenientCast(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgShreddedInt(t, mem, 3, 5) + defer arr.Release() + + f64, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64}) + require.NoError(t, err) + defer f64.Release() + assert.EqualValues(t, 3, f64.(*array.Float64).Value(0)) + assert.EqualValues(t, 5, f64.(*array.Float64).Value(1)) + + dec, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: &arrow.Decimal128Type{Precision: 10, Scale: 0}}) + require.NoError(t, err) + defer dec.Release() + assert.Equal(t, 2, dec.Len()) +} + +func TestVariantGetStrictCastErrors(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgShreddedInt(t, mem, 5_000_000_000) // overflows int8 + defer arr.Release() + + _, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + AsType: arrow.PrimitiveTypes.Int8, Strict: true, + }) + require.Error(t, err, "strict cast of an overflowing value must error") +} + +// TestVariantGetFieldOnScalarErrors covers :323: a field step into a scalar is a +// type error, not a silent null. +func TestVariantGetFieldOnScalarErrors(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, int64(1)) + defer arr.Release() + + _, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a")}) + require.ErrorIs(t, err, arrow.ErrInvalid) +} + +// TestVariantGetHugeIndex covers :307: a huge index must not wrap to a valid one. +func TestVariantGetHugeIndex(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, []any{int64(10), int64(20)}) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(1 << 40), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + assert.True(t, out.(*array.Int64).IsNull(0)) +} + +func TestVariantGetEmptyPath(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, int64(1), int64(2)) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{}) + require.NoError(t, err) + defer out.Release() + assert.Equal(t, 2, out.Len()) +} + +// TestVariantGetShreddedFieldPushdown drives nested field steps through the shredded +// typed_value columns and the perfect-shredding fast path. +func TestVariantGetShreddedFieldPushdown(t *testing.T) { + mem := memory.DefaultAllocator + vt := extensions.NewShreddedVariantType(arrow.StructOf( + arrow.Field{Name: "a", Type: arrow.StructOf( + arrow.Field{Name: "b", Type: arrow.PrimitiveTypes.Int64})})) + bldr := extensions.NewVariantBuilder(mem, vt) + defer bldr.Release() + bldr.Append(vgVariant(t, map[string]any{"a": map[string]any{"b": int64(5)}})) + bldr.Append(vgVariant(t, map[string]any{"a": map[string]any{"b": int64(6)}})) + arr := bldr.NewArray().(*extensions.VariantArray) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Field("a").Field("b"), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + assert.EqualValues(t, 5, ints.Value(0)) + assert.EqualValues(t, 6, ints.Value(1)) +} + +// TestVariantGetShreddedListIndex drives an index step over a shredded list, which +// gathers elements with the take kernel. +func TestVariantGetShreddedListIndex(t *testing.T) { + mem := memory.DefaultAllocator + vt := extensions.NewShreddedVariantType(arrow.ListOf(arrow.PrimitiveTypes.Int64)) + bldr := extensions.NewVariantBuilder(mem, vt) + defer bldr.Release() + bldr.Append(vgVariant(t, []any{int64(10), int64(20), int64(30)})) + bldr.Append(vgVariant(t, []any{int64(40), int64(50)})) + arr := bldr.NewArray().(*extensions.VariantArray) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(1), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + assert.EqualValues(t, 20, ints.Value(0)) + assert.EqualValues(t, 50, ints.Value(1)) +} + +func TestVariantGetNoLeak(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := exec.WithAllocator(context.Background(), mem) + + vt := extensions.NewShreddedVariantType(arrow.ListOf(arrow.PrimitiveTypes.Int64)) + bldr := extensions.NewVariantBuilder(mem, vt) + bldr.Append(vgVariant(t, []any{int64(10), int64(20)})) + bldr.AppendNull() + arr := bldr.NewArray().(*extensions.VariantArray) + bldr.Release() + + idx, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(0), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + idx.Release() + + missing, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Field("nope"), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + missing.Release() + + arr.Release() +} + +// TestVariantGetDictMetadata guards against a panic when the metadata column is +// dictionary-encoded (spec-legal): buildTargetVariant must read the raw column, +// not the plain-binary accessor. +func TestVariantGetDictMetadata(t *testing.T) { + mem := memory.DefaultAllocator + s := arrow.StructOf( + arrow.Field{Name: "metadata", Type: &arrow.DictionaryType{IndexType: arrow.PrimitiveTypes.Uint8, ValueType: arrow.BinaryTypes.Binary}}, + arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, + arrow.Field{Name: "typed_value", Type: arrow.StructOf( + arrow.Field{Name: "a", Type: arrow.StructOf( + arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, + arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}, + )}, + ), Nullable: true}) + vt, err := extensions.NewVariantType(s) + require.NoError(t, err) + bldr := vt.NewBuilder(mem).(*extensions.VariantBuilder) + defer bldr.Release() + bldr.Append(vgVariant(t, map[string]any{"a": int64(5), "b": "resid"})) + arr := bldr.NewArray().(*extensions.VariantArray) + defer arr.Release() + + // "b" is not shredded, so this takes the NotShredded -> buildTargetVariant path. + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: variant.VariantPath{}.Field("b")}) + require.NoError(t, err) + defer out.Release() + v, err := out.(*extensions.VariantArray).Value(0) + require.NoError(t, err) + assert.Equal(t, "resid", v.Value()) +} diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index a4f5e3b42..a0cc74edc 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -458,6 +458,11 @@ func (v *VariantArray) IsShredded() bool { return v.ExtensionType().(*VariantType).typedValueFieldIdx != -1 } +// VariantType returns the array's extension type without the ExtensionType cast. +func (v *VariantArray) VariantType() *VariantType { + return v.ExtensionType().(*VariantType) +} + // UnshredVariant returns an equivalent VariantArray in the non-shredded layout // (a struct of metadata and value), reassembling each row's value from the // shredded typed_value and value columns. If the array is already non-shredded @@ -1475,96 +1480,85 @@ func (b *shreddedPrimitiveBuilder) tryTyped(v variant.Value) (residual []byte) { return v.Bytes() } - if appendVariantToTypedBuilder(b.typedBldr, v) { - return nil - } - - b.typedBldr.AppendNull() - return v.Bytes() -} - -// appendVariantToTypedBuilder appends v to a typed primitive builder when v's type -// fits the builder, reporting whether it did. Shared by the shredding writer and VariantGet. -func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { - switch bldr := target.(type) { + switch bldr := b.typedBldr.(type) { case *array.Int8Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Uint8Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Int16Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Uint16Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Int32Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Uint32Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Int64Builder: if appendNumericToTarget(bldr, v) { - return true + return nil } case *array.Float32Builder: switch v.Type() { case variant.Float: bldr.Append(v.Value().(float32)) - return true + return nil case variant.Double: val := v.Value().(float64) if val >= -math.MaxFloat32 && val <= math.MaxFloat32 { bldr.Append(float32(val)) - return true + return nil } } case *array.Float64Builder: switch v.Type() { case variant.Float: bldr.Append(float64(v.Value().(float32))) - return true + return nil case variant.Double: bldr.Append(v.Value().(float64)) - return true + return nil } case *array.BooleanBuilder: if v.Type() == variant.Bool { bldr.Append(v.Value().(bool)) - return true + return nil } case array.StringLikeBuilder: if v.Type() == variant.String { bldr.Append(v.Value().(string)) - return true + return nil } case array.BinaryLikeBuilder: if v.Type() == variant.Binary { bldr.Append(v.Value().([]byte)) - return true + return nil } case *array.Date32Builder: if v.Type() == variant.Date { bldr.Append(v.Value().(arrow.Date32)) - return true + return nil } case *array.Time64Builder: if v.Type() == variant.Time && bldr.Type().(*arrow.Time64Type).Unit == arrow.Microsecond { bldr.Append(v.Value().(arrow.Time64)) - return true + return nil } case *UUIDBuilder: if v.Type() == variant.UUID { bldr.Append(v.Value().(uuid.UUID)) - return true + return nil } case *array.TimestampBuilder: tsType := bldr.Type().(*arrow.TimestampType) @@ -1577,10 +1571,10 @@ func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { switch tsType.Unit { case arrow.Microsecond: bldr.Append(v.Value().(arrow.Timestamp)) - return true + return nil case arrow.Nanosecond: bldr.Append(v.Value().(arrow.Timestamp) * 1000) - return true + return nil } case variant.TimestampMicrosNTZ: if tsType.TimeZone != "" { @@ -1590,20 +1584,20 @@ func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { switch tsType.Unit { case arrow.Microsecond: bldr.Append(v.Value().(arrow.Timestamp)) - return true + return nil case arrow.Nanosecond: bldr.Append(v.Value().(arrow.Timestamp) * 1000) - return true + return nil } case variant.TimestampNanos: if tsType.TimeZone == "UTC" && tsType.Unit == arrow.Nanosecond { bldr.Append(v.Value().(arrow.Timestamp)) - return true + return nil } case variant.TimestampNanosNTZ: if tsType.TimeZone == "" && tsType.Unit == arrow.Nanosecond { bldr.Append(v.Value().(arrow.Timestamp)) - return true + return nil } } case *array.Decimal32Builder: @@ -1612,17 +1606,17 @@ func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal32)) - return true + return nil } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal32(val.Value.(decimal.Decimal64))) - return true + return nil } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal32(val.Value.(decimal.Decimal128).LowBits())) - return true + return nil } } case *array.Decimal64Builder: @@ -1631,17 +1625,17 @@ func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal64(val.Value.(decimal.Decimal32))) - return true + return nil } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal64)) - return true + return nil } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(decimal.Decimal64(val.Value.(decimal.Decimal128).LowBits())) - return true + return nil } } case *array.Decimal128Builder: @@ -1650,22 +1644,23 @@ func appendVariantToTypedBuilder(target array.Builder, v variant.Value) bool { case variant.DecimalValue[decimal.Decimal32]: if decimalCanFit(dt, val) { bldr.Append(decimal128.FromI64(int64(val.Value.(decimal.Decimal32)))) - return true + return nil } case variant.DecimalValue[decimal.Decimal64]: if decimalCanFit(dt, val) { bldr.Append(decimal128.FromI64(int64(val.Value.(decimal.Decimal64)))) - return true + return nil } case variant.DecimalValue[decimal.Decimal128]: if decimalCanFit(dt, val) { bldr.Append(val.Value.(decimal.Decimal128)) - return true + return nil } } } - return false + b.typedBldr.AppendNull() + return v.Bytes() } type shreddedFieldBuilder struct { diff --git a/arrow/extensions/variant_get.go b/arrow/extensions/variant_get.go deleted file mode 100644 index 09f8e79bd..000000000 --- a/arrow/extensions/variant_get.go +++ /dev/null @@ -1,453 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package extensions - -import ( - "fmt" - - "github.com/apache/arrow-go/v18/arrow" - "github.com/apache/arrow-go/v18/arrow/array" - "github.com/apache/arrow-go/v18/arrow/bitutil" - "github.com/apache/arrow-go/v18/arrow/memory" - "github.com/apache/arrow-go/v18/parquet/variant" -) - -// VariantPathElement is a single step of a variant path: either an object field -// name or an array index. -type VariantPathElement struct { - name string - index int - isIndex bool -} - -// VariantPathField returns a path element selecting the named object field. -func VariantPathField(name string) VariantPathElement { - return VariantPathElement{name: name} -} - -// VariantPathIndex returns a path element selecting the array element at index. -func VariantPathIndex(index int) VariantPathElement { - return VariantPathElement{index: index, isIndex: true} -} - -// VariantPath is an ordered list of path elements to extract from a variant value. -type VariantPath []VariantPathElement - -// GetOptions controls VariantGet. -type GetOptions struct { - // Path is the path to extract from each variant value. - Path VariantPath - // AsType, when nil, makes VariantGet return a VariantArray pointing at the path. - // When set, the extracted value is cast to this type. Nested (struct/list) types - // are not yet supported and yield arrow.ErrNotImplemented. - AsType arrow.DataType - // Strict makes a cast failure return an error. The default (false) mirrors - // arrow-rs: a cast failure produces null. - Strict bool - // Mem is the allocator for output arrays; nil uses memory.DefaultAllocator. - Mem memory.Allocator -} - -// VariantGet extracts opts.Path from each value of a VariantArray. It follows the -// shredded typed_value columns as far as the path allows, then falls back to a -// per-row walk of the residual value for the remainder. -func VariantGet(input arrow.Array, opts GetOptions) (arrow.Array, error) { - va, ok := input.(*VariantArray) - if !ok { - return nil, fmt.Errorf("%w: VariantGet input must be a VariantArray, got %T", arrow.ErrInvalid, input) - } - - if opts.Mem == nil { - opts.Mem = memory.DefaultAllocator - } - - return shreddedGetPath(va, opts) -} - -// shreddingState is a (value?, typed_value?) column pair at one level of a shredded -// variant, mirroring arrow-rs ShreddingState. -type shreddingState struct { - value arrow.TypedArray[[]byte] - typedValue arrow.Array - length int -} - -func stateFromVariant(va *VariantArray) shreddingState { - vt := va.ExtensionType().(*VariantType) - st := va.Storage().(*array.Struct) - - var value arrow.TypedArray[[]byte] - if vt.valueFieldIdx != -1 { - value = st.Field(vt.valueFieldIdx).(arrow.TypedArray[[]byte]) - } - - var typed arrow.Array - if vt.typedValueFieldIdx != -1 { - typed = st.Field(vt.typedValueFieldIdx) - } - - return shreddingState{value: value, typedValue: typed, length: va.Len()} -} - -func stateFromFieldStruct(child *array.Struct) shreddingState { - ct := child.DataType().(*arrow.StructType) - - var value arrow.TypedArray[[]byte] - if idx, ok := ct.FieldIdx("value"); ok { - value = child.Field(idx).(arrow.TypedArray[[]byte]) - } - - var typed arrow.Array - if idx, ok := ct.FieldIdx("typed_value"); ok { - typed = child.Field(idx) - } - - return shreddingState{value: value, typedValue: typed, length: child.Len()} -} - -type pathStepKind int - -const ( - stepSuccess pathStepKind = iota - stepMissing - stepNotShredded -) - -type pathStep struct { - kind pathStepKind - state shreddingState -} - -// missingStep decides whether an absent typed field means the value is provably -// missing (value column all-null) or merely not shredded (residual may hold it). -func (s shreddingState) missingStep() pathStep { - if s.value == nil || s.value.NullN() == s.value.Len() { - return pathStep{kind: stepMissing} - } - - return pathStep{kind: stepNotShredded} -} - -// followFieldElement takes one field step deeper into the shredded columns. -func followFieldElement(s shreddingState, name string) (pathStep, error) { - if s.typedValue == nil { - return s.missingStep(), nil - } - - st, ok := s.typedValue.(*array.Struct) - if !ok { - return s.missingStep(), nil - } - - idx, ok := st.DataType().(*arrow.StructType).FieldIdx(name) - if !ok { - return s.missingStep(), nil - } - - child, ok := st.Field(idx).(*array.Struct) - if !ok { - return pathStep{}, fmt.Errorf("%w: expected struct field %q while following path, got %s", - arrow.ErrInvalid, name, st.Field(idx).DataType()) - } - - return pathStep{kind: stepSuccess, state: stateFromFieldStruct(child)}, nil -} - -func shreddedGetPath(va *VariantArray, opts GetOptions) (arrow.Array, error) { - state := stateFromVariant(va) - nulls := newNullTracker(va.Len()) - nulls.apply(va.Storage()) - - // Peel the field prefix of the path through the shredded columns. Index steps - // and non-shredded fields stop the columnar walk and hand the rest to a per-row - // fallback over the fully reassembled value at the current node. - idx := 0 - for idx < len(opts.Path) { - elem := opts.Path[idx] - if elem.isIndex { - break - } - - step, err := followFieldElement(state, elem.name) - if err != nil { - return nil, err - } - - switch step.kind { - case stepSuccess: - nulls.apply(state.typedValue) - state = step.state - idx++ - - continue - case stepMissing: - return allNullResult(va, opts) - } - - break // stepNotShredded - } - - remaining := opts.Path[idx:] - target, err := buildTargetVariant(va, state, nulls, opts.Mem) - if err != nil { - return nil, err - } - defer target.Release() - - if len(remaining) == 0 { - if opts.AsType == nil { - target.Retain() - - return target, nil - } - - if shredded := tryPerfectShredding(state, nulls, opts.AsType); shredded != nil { - return shredded, nil - } - } - - return shredBasicVariant(target, remaining, opts) -} - -// shredBasicVariant walks the remaining path per row and produces either a -// VariantArray (AsType nil) or a typed array. -func shredBasicVariant(target *VariantArray, remaining VariantPath, opts GetOptions) (arrow.Array, error) { - if opts.AsType == nil { - bldr := NewVariantBuilder(opts.Mem, NewDefaultVariantType()) - defer bldr.Release() - bldr.Reserve(target.Len()) - - for i := 0; i < target.Len(); i++ { - leaf, ok, err := navigateRow(target, i, remaining) - if err != nil { - return nil, err - } - if !ok { - bldr.AppendNull() - - continue - } - bldr.Append(leaf) - } - - return bldr.NewArray(), nil - } - - if _, ok := opts.AsType.(arrow.NestedType); ok { - return nil, fmt.Errorf("%w: VariantGet cast to nested type %s", arrow.ErrNotImplemented, opts.AsType) - } - - bldr := array.NewBuilder(opts.Mem, opts.AsType) - defer bldr.Release() - bldr.Reserve(target.Len()) - - for i := 0; i < target.Len(); i++ { - leaf, ok, err := navigateRow(target, i, remaining) - if err != nil { - return nil, err - } - if !ok || leaf.Type() == variant.Null { - bldr.AppendNull() - - continue - } - - if appendVariantToTypedBuilder(bldr, leaf) { - continue - } - - if opts.Strict { - return nil, fmt.Errorf("%w: cannot cast variant %v to %s", arrow.ErrInvalid, leaf.Type(), opts.AsType) - } - - bldr.AppendNull() - } - - return bldr.NewArray(), nil -} - -// navigateRow reassembles row i of target and walks path into it. It returns -// (value, false) when the row is null or the path is absent. -func navigateRow(target *VariantArray, i int, path VariantPath) (variant.Value, bool, error) { - if target.IsNull(i) { - return variant.Value{}, false, nil - } - - v, err := target.Value(i) - if err != nil { - return variant.Value{}, false, fmt.Errorf("variant: reassembling row %d: %w", i, err) - } - - return navigateValue(v, path) -} - -// navigateValue walks path into a fully reassembled variant value. -func navigateValue(v variant.Value, path VariantPath) (variant.Value, bool, error) { - cur := v - for _, elem := range path { - if elem.isIndex { - arr, ok := cur.Value().(variant.ArrayValue) - if !ok || elem.index < 0 || uint32(elem.index) >= arr.Len() { - return variant.Value{}, false, nil - } - el, err := arr.Value(uint32(elem.index)) - if err != nil { - return variant.Value{}, false, nil - } - cur = el - - continue - } - - obj, ok := cur.Value().(variant.ObjectValue) - if !ok { - return variant.Value{}, false, nil - } - field, err := obj.ValueByKey(elem.name) - if err != nil { - return variant.Value{}, false, nil - } - cur = field.Value - } - - return cur, true, nil -} - -// tryPerfectShredding returns the typed_value column directly when the target is -// perfectly shredded to AsType. It only fires when no ancestor nulls need merging; -// otherwise the caller's per-row path produces the same values. -func tryPerfectShredding(state shreddingState, nulls *nullTracker, asType arrow.DataType) arrow.Array { - if _, ok := asType.(arrow.NestedType); ok { - return nil - } - if state.typedValue == nil || !nulls.allValid() { - return nil - } - if !arrow.TypeEqual(state.typedValue.DataType(), asType) { - return nil - } - if state.value != nil && state.value.NullN() != state.value.Len() { - return nil - } - - state.typedValue.Retain() - - return state.typedValue -} - -// buildTargetVariant wraps the current shredding state as a VariantArray, carrying -// the accumulated ancestor nulls onto the storage struct. -func buildTargetVariant(va *VariantArray, state shreddingState, nulls *nullTracker, mem memory.Allocator) (*VariantArray, error) { - // Take the raw metadata array (not va.Metadata) so dictionary- or large-binary- - // encoded metadata is preserved and decoded by the target's own reader. - srcVT := va.ExtensionType().(*VariantType) - metadata := va.Storage().(*array.Struct).Field(srcVT.metadataFieldIdx) - - fields := []arrow.Field{{Name: "metadata", Type: metadata.DataType(), Nullable: false}} - cols := []arrow.Array{metadata} - - if state.value != nil { - fields = append(fields, arrow.Field{Name: "value", Type: state.value.DataType(), Nullable: true}) - cols = append(cols, state.value) - } - if state.typedValue != nil { - fields = append(fields, arrow.Field{Name: "typed_value", Type: state.typedValue.DataType(), Nullable: true}) - cols = append(cols, state.typedValue) - } - - bitmap, nullCount := nulls.bitmap(mem) - if bitmap != nil { - defer bitmap.Release() - } - - st, err := array.NewStructArrayWithFieldsAndNulls(cols, fields, bitmap, nullCount, 0) - if err != nil { - return nil, err - } - defer st.Release() - - vt, err := NewVariantType(st.DataType()) - if err != nil { - return nil, err - } - - return array.NewExtensionArrayWithStorage(vt, st).(*VariantArray), nil -} - -// allNullResult builds the all-null output for a provably missing path. -func allNullResult(va *VariantArray, opts GetOptions) (arrow.Array, error) { - if opts.AsType != nil { - return array.MakeArrayOfNull(opts.Mem, opts.AsType, va.Len()), nil - } - - bldr := NewVariantBuilder(opts.Mem, NewDefaultVariantType()) - defer bldr.Release() - for i := 0; i < va.Len(); i++ { - bldr.AppendNull() - } - - return bldr.NewArray(), nil -} - -// nullTracker accumulates ancestor null masks encountered while walking the path. -type nullTracker struct { - length int - valid []bool // nil means all valid -} - -func newNullTracker(length int) *nullTracker { - return &nullTracker{length: length} -} - -func (n *nullTracker) apply(arr arrow.Array) { - if arr == nil || arr.NullN() == 0 { - return - } - if n.valid == nil { - n.valid = make([]bool, n.length) - for i := range n.valid { - n.valid[i] = true - } - } - for i := 0; i < n.length; i++ { - if arr.IsNull(i) { - n.valid[i] = false - } - } -} - -func (n *nullTracker) allValid() bool { return n.valid == nil } - -func (n *nullTracker) bitmap(mem memory.Allocator) (*memory.Buffer, int) { - if n.valid == nil { - return nil, 0 - } - - buf := memory.NewResizableBuffer(mem) - buf.Resize(int(bitutil.BytesForBits(int64(n.length)))) - nullCount := 0 - for i, v := range n.valid { - if v { - bitutil.SetBit(buf.Bytes(), i) - } else { - bitutil.ClearBit(buf.Bytes(), i) - nullCount++ - } - } - - return buf, nullCount -} diff --git a/arrow/extensions/variant_get_internal_test.go b/arrow/extensions/variant_get_internal_test.go deleted file mode 100644 index 4119aa618..000000000 --- a/arrow/extensions/variant_get_internal_test.go +++ /dev/null @@ -1,119 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package extensions - -import ( - "testing" - - "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/parquet/variant" - "github.com/stretchr/testify/require" -) - -func mkShreddedIntObj(t *testing.T, mem memory.Allocator, v int64) *VariantArray { - t.Helper() - vt := NewShreddedVariantType(arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) - bldr := NewVariantBuilder(mem, vt) - defer bldr.Release() - var vb variant.Builder - require.NoError(t, vb.Append(map[string]any{"a": v})) - val, err := vb.Build() - require.NoError(t, err) - bldr.Append(val) - - return bldr.NewArray().(*VariantArray) -} - -// TestTryPerfectShreddingFires proves the fast path returns the typed_value column -// directly for a perfect shredding, and declines when the residual value has data. -func TestTryPerfectShreddingFires(t *testing.T) { - mem := memory.DefaultAllocator - arr := mkShreddedIntObj(t, mem, 5) - defer arr.Release() - - state := stateFromVariant(arr) - nulls := newNullTracker(arr.Len()) - nulls.apply(arr.Storage()) - - step, err := followFieldElement(state, "a") - require.NoError(t, err) - require.Equal(t, stepSuccess, step.kind) - - out := tryPerfectShredding(step.state, nulls, arrow.PrimitiveTypes.Int64) - require.NotNil(t, out, "perfect shredding must fire for a fully shredded int64 leaf") - defer out.Release() - require.Equal(t, int64(5), out.(*array.Int64).Value(0)) - - // A non-matching target type must decline. - require.Nil(t, tryPerfectShredding(step.state, nulls, arrow.PrimitiveTypes.Int32)) -} - -func TestVariantGetNoLeak(t *testing.T) { - mem := memory.NewCheckedAllocator(memory.DefaultAllocator) - defer mem.AssertSize(t, 0) - - arr := mkShreddedIntObj(t, mem, 9) - - perfect, err := VariantGet(arr, GetOptions{ - Path: VariantPath{VariantPathField("a")}, - AsType: arrow.PrimitiveTypes.Int64, - Mem: mem, - }) - require.NoError(t, err) - perfect.Release() - - variantOut, err := VariantGet(arr, GetOptions{ - Path: VariantPath{VariantPathField("a")}, - Mem: mem, - }) - require.NoError(t, err) - variantOut.Release() - - missing, err := VariantGet(arr, GetOptions{ - Path: VariantPath{VariantPathField("nope")}, - AsType: arrow.PrimitiveTypes.Int64, - Mem: mem, - }) - require.NoError(t, err) - missing.Release() - - arr.Release() - - // Ancestor-null case: forces nullTracker to allocate a real bitmap buffer that - // buildTargetVariant threads onto the target struct. - vt := NewShreddedVariantType(arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) - nb := NewVariantBuilder(mem, vt) - var vb variant.Builder - require.NoError(t, vb.Append(map[string]any{"a": int64(1)})) - val, err := vb.Build() - require.NoError(t, err) - nb.Append(val) - nb.AppendNull() - withNull := nb.NewArray().(*VariantArray) - nb.Release() - - got, err := VariantGet(withNull, GetOptions{ - Path: VariantPath{VariantPathField("a")}, - Mem: mem, - }) - require.NoError(t, err) - require.True(t, got.(*VariantArray).IsNull(1)) - got.Release() - withNull.Release() -} diff --git a/arrow/extensions/variant_get_test.go b/arrow/extensions/variant_get_test.go deleted file mode 100644 index 455c4ecf7..000000000 --- a/arrow/extensions/variant_get_test.go +++ /dev/null @@ -1,340 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package extensions_test - -import ( - "testing" - - "github.com/apache/arrow-go/v18/arrow" - "github.com/apache/arrow-go/v18/arrow/array" - "github.com/apache/arrow-go/v18/arrow/extensions" - "github.com/apache/arrow-go/v18/arrow/memory" - "github.com/apache/arrow-go/v18/parquet/variant" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func mkVariant(t *testing.T, v any) variant.Value { - t.Helper() - var b variant.Builder - require.NoError(t, b.Append(v)) - val, err := b.Build() - require.NoError(t, err) - - return val -} - -// nonShreddedVariants builds a plain (metadata, value) VariantArray from Go values. -func nonShreddedVariants(t *testing.T, mem memory.Allocator, vals ...any) *extensions.VariantArray { - t.Helper() - bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) - defer bldr.Release() - for _, v := range vals { - if v == nil { - bldr.AppendNull() - - continue - } - bldr.Append(mkVariant(t, v)) - } - - return bldr.NewArray().(*extensions.VariantArray) -} - -// shreddedIntObjects builds a VariantArray shredding an object with a single int64 field "a". -func shreddedIntObjects(t *testing.T, mem memory.Allocator, vals ...int64) *extensions.VariantArray { - t.Helper() - vt := extensions.NewShreddedVariantType(arrow.StructOf( - arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) - bldr := extensions.NewVariantBuilder(mem, vt) - defer bldr.Release() - for _, v := range vals { - bldr.Append(mkVariant(t, map[string]any{"a": v})) - } - - return bldr.NewArray().(*extensions.VariantArray) -} - -func TestVariantGetTypedOutput(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, - map[string]any{"a": int64(1), "b": "x"}, - map[string]any{"a": int64(2), "b": "y"}, - nil, - ) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a")}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - ints := out.(*array.Int64) - require.Equal(t, 3, ints.Len()) - assert.EqualValues(t, 1, ints.Value(0)) - assert.EqualValues(t, 2, ints.Value(1)) - assert.True(t, ints.IsNull(2)) -} - -func TestVariantGetVariantOutput(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, - map[string]any{"a": int64(7)}, - map[string]any{"b": int64(9)}, // no "a" -> null - ) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a")}, - }) - require.NoError(t, err) - defer out.Release() - - varr := out.(*extensions.VariantArray) - require.Equal(t, 2, varr.Len()) - - v, err := varr.Value(0) - require.NoError(t, err) - assert.EqualValues(t, 7, v.Value()) - assert.True(t, varr.IsNull(1)) -} - -func TestVariantGetNestedPath(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, map[string]any{"a": map[string]any{"b": int64(5)}}) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a"), extensions.VariantPathField("b")}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - assert.EqualValues(t, 5, out.(*array.Int64).Value(0)) -} - -func TestVariantGetIndex(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, []any{int64(10), int64(20), int64(30)}) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathIndex(1)}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - assert.EqualValues(t, 20, out.(*array.Int64).Value(0)) - - oob, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathIndex(9)}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer oob.Release() - assert.True(t, oob.(*array.Int64).IsNull(0)) -} - -func TestVariantGetPerfectShredding(t *testing.T) { - mem := memory.DefaultAllocator - arr := shreddedIntObjects(t, mem, 11, 22, 33) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a")}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - ints := out.(*array.Int64) - require.Equal(t, 3, ints.Len()) - assert.EqualValues(t, 11, ints.Value(0)) - assert.EqualValues(t, 22, ints.Value(1)) - assert.EqualValues(t, 33, ints.Value(2)) -} - -func TestVariantGetShreddedVariantOutput(t *testing.T) { - mem := memory.DefaultAllocator - arr := shreddedIntObjects(t, mem, 100, 200) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a")}, - }) - require.NoError(t, err) - defer out.Release() - - varr := out.(*extensions.VariantArray) - v, err := varr.Value(0) - require.NoError(t, err) - assert.EqualValues(t, 100, v.Value()) - v, err = varr.Value(1) - require.NoError(t, err) - assert.EqualValues(t, 200, v.Value()) -} - -func TestVariantGetMissingField(t *testing.T) { - mem := memory.DefaultAllocator - arr := shreddedIntObjects(t, mem, 1, 2) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("missing")}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - ints := out.(*array.Int64) - require.Equal(t, 2, ints.Len()) - assert.True(t, ints.IsNull(0)) - assert.True(t, ints.IsNull(1)) -} - -// TestVariantGetNotShreddedFallback covers a field present only in the residual value -// of a shredded object: the columnar walk stops and the per-row fallback recovers it. -func TestVariantGetNotShreddedFallback(t *testing.T) { - mem := memory.DefaultAllocator - vt := extensions.NewShreddedVariantType(arrow.StructOf( - arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) - bldr := extensions.NewVariantBuilder(mem, vt) - defer bldr.Release() - // "b" is not in the shredding schema, so it lands in the residual value column. - bldr.Append(mkVariant(t, map[string]any{"a": int64(1), "b": int64(42)})) - arr := bldr.NewArray().(*extensions.VariantArray) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("b")}, - AsType: arrow.PrimitiveTypes.Int64, - }) - require.NoError(t, err) - defer out.Release() - - assert.EqualValues(t, 42, out.(*array.Int64).Value(0)) -} - -func TestVariantGetNestedTypeUnsupported(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, map[string]any{"a": int64(1)}) - defer arr.Release() - - _, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("a")}, - AsType: arrow.StructOf(arrow.Field{Name: "x", Type: arrow.PrimitiveTypes.Int64}), - }) - require.ErrorIs(t, err, arrow.ErrNotImplemented) -} - -func TestVariantGetRejectsNonVariant(t *testing.T) { - mem := memory.DefaultAllocator - bldr := array.NewInt64Builder(mem) - defer bldr.Release() - bldr.Append(1) - arr := bldr.NewArray() - defer arr.Release() - - _, err := extensions.VariantGet(arr, extensions.GetOptions{}) - require.ErrorIs(t, err, arrow.ErrInvalid) -} - -// TestVariantGetStrictVsDefaultCast covers the cast-outcome axis: an uncastable -// value nulls by default (mirrors arrow-rs) and errors under Strict. -func TestVariantGetStrictVsDefaultCast(t *testing.T) { - mem := memory.DefaultAllocator - arr := nonShreddedVariants(t, mem, map[string]any{"a": "not-a-number"}) - defer arr.Release() - path := extensions.VariantPath{extensions.VariantPathField("a")} - - // Default (Strict false): uncastable string -> Int64 becomes null. - out, err := extensions.VariantGet(arr, extensions.GetOptions{Path: path, AsType: arrow.PrimitiveTypes.Int64}) - require.NoError(t, err) - defer out.Release() - assert.True(t, out.(*array.Int64).IsNull(0)) - - // Strict: the same cast returns an error. - _, err = extensions.VariantGet(arr, extensions.GetOptions{Path: path, AsType: arrow.PrimitiveTypes.Int64, Strict: true}) - require.ErrorIs(t, err, arrow.ErrInvalid) -} - -// TestVariantGetValuelessLayout guards against the Field(-1) panic on a shredded -// layout that has no residual value column. -func TestVariantGetValuelessLayout(t *testing.T) { - mem := memory.DefaultAllocator - s := arrow.StructOf( - arrow.Field{Name: "metadata", Type: arrow.BinaryTypes.Binary}, - arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}) - b := array.NewStructBuilder(mem, s) - defer b.Release() - mb := b.FieldBuilder(0).(*array.BinaryBuilder) - tv := b.FieldBuilder(1).(*array.Int64Builder) - b.Append(true) - mb.Append(variant.EmptyMetadataBytes[:]) - tv.AppendNull() - st := b.NewArray() - defer st.Release() - - vt, err := extensions.NewVariantType(s) - require.NoError(t, err) - arr := array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray) - defer arr.Release() - - // AsType mismatch (Int32 vs Int64) forces the per-row fallback over a value-less target. - out, err := extensions.VariantGet(arr, extensions.GetOptions{AsType: arrow.PrimitiveTypes.Int32}) - require.NoError(t, err) - defer out.Release() - assert.True(t, out.(*array.Int32).IsNull(0)) -} - -// on a dictionary-encoded metadata column, which is spec-legal. -func TestVariantGetDictMetadata(t *testing.T) { - mem := memory.DefaultAllocator - s := arrow.StructOf( - arrow.Field{Name: "metadata", Type: &arrow.DictionaryType{ - IndexType: arrow.PrimitiveTypes.Uint8, ValueType: arrow.BinaryTypes.Binary}}, - arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, - arrow.Field{Name: "typed_value", Type: arrow.StructOf( - arrow.Field{Name: "a", Type: arrow.StructOf( - arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, Nullable: true}, - arrow.Field{Name: "typed_value", Type: arrow.PrimitiveTypes.Int64, Nullable: true}, - )}, - ), Nullable: true}) - - vt, err := extensions.NewVariantType(s) - require.NoError(t, err) - bldr := vt.NewBuilder(mem).(*extensions.VariantBuilder) - defer bldr.Release() - // "b" is not shredded, so extracting it takes the NotShredded -> buildTargetVariant path. - bldr.Append(mkVariant(t, map[string]any{"a": int64(5), "b": "resid"})) - arr := bldr.NewArray().(*extensions.VariantArray) - defer arr.Release() - - out, err := extensions.VariantGet(arr, extensions.GetOptions{ - Path: extensions.VariantPath{extensions.VariantPathField("b")}, - }) - require.NoError(t, err) - defer out.Release() - - v, err := out.(*extensions.VariantArray).Value(0) - require.NoError(t, err) - assert.Equal(t, "resid", v.Value()) -} diff --git a/parquet/variant/path.go b/parquet/variant/path.go new file mode 100644 index 000000000..437e08529 --- /dev/null +++ b/parquet/variant/path.go @@ -0,0 +1,108 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package variant + +import ( + "errors" + "fmt" + + "github.com/apache/arrow-go/v18/arrow" +) + +// pathElem is one step of a VariantPath: an object field when name != "", else an +// array index. +type pathElem struct { + name string + index int +} + +// VariantPath is an ordered list of steps to navigate into a variant value. The +// zero value is the root path; extend it with Field and Index. +type VariantPath struct { + elems []pathElem +} + +// Field returns a copy of the path with an object-field step appended. +func (p VariantPath) Field(name string) VariantPath { + return VariantPath{elems: append(p.grow(), pathElem{name: name})} +} + +// Index returns a copy of the path with an array-index step appended. +func (p VariantPath) Index(i int) VariantPath { + return VariantPath{elems: append(p.grow(), pathElem{index: i})} +} + +// Join returns a copy of the path with other's steps appended. +func (p VariantPath) Join(other VariantPath) VariantPath { + return VariantPath{elems: append(p.grow(), other.elems...)} +} + +func (p VariantPath) grow() []pathElem { + return append(make([]pathElem, 0, len(p.elems)+1), p.elems...) +} + +// Len returns the number of steps in the path. +func (p VariantPath) Len() int { return len(p.elems) } + +// StepAt returns the i-th step. When name != "" it is an object-field step; +// otherwise it is an array-index step selecting index. +func (p VariantPath) StepAt(i int) (name string, index int) { + return p.elems[i].name, p.elems[i].index +} + +// GetByPath navigates path into v and returns the leaf value. found is false when +// the path is cleanly absent (a missing object field, or an out-of-range or +// non-array index). It returns an error for a type error (a field step into a +// non-object) or corrupt data (a field id not present in the metadata). +func (v Value) GetByPath(path VariantPath) (leaf Value, found bool, err error) { + cur := v + for _, e := range path.elems { + if e.name != "" { + obj, ok := cur.Value().(ObjectValue) + if !ok { + return Value{}, false, fmt.Errorf("%w: variant path field %q applied to non-object", arrow.ErrInvalid, e.name) + } + field, ferr := obj.ValueByKey(e.name) + if ferr != nil { + if errors.Is(ferr, arrow.ErrNotFound) { + return Value{}, false, nil + } + + return Value{}, false, ferr + } + cur = field.Value + + continue + } + + arr, ok := cur.Value().(ArrayValue) + if !ok { + return Value{}, false, nil + } + if e.index < 0 || uint64(e.index) >= uint64(arr.Len()) { + return Value{}, false, nil + } + el, aerr := arr.Value(uint32(e.index)) + if aerr != nil { + return Value{}, false, nil + } + cur = el + } + + return cur, true, nil +} diff --git a/parquet/variant/path_test.go b/parquet/variant/path_test.go new file mode 100644 index 000000000..6a8cfaaea --- /dev/null +++ b/parquet/variant/path_test.go @@ -0,0 +1,94 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package variant_test + +import ( + "testing" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/parquet/variant" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func buildVar(t *testing.T, v any) variant.Value { + t.Helper() + var b variant.Builder + require.NoError(t, b.Append(v)) + val, err := b.Build() + require.NoError(t, err) + + return val +} + +func TestGetByPathFieldAndIndex(t *testing.T) { + v := buildVar(t, map[string]any{"a": map[string]any{"b": int64(5)}, "arr": []any{int64(10), int64(20)}}) + + leaf, found, err := v.GetByPath(variant.VariantPath{}.Field("a").Field("b")) + require.NoError(t, err) + require.True(t, found) + assert.EqualValues(t, 5, leaf.Value()) + + leaf, found, err = v.GetByPath(variant.VariantPath{}.Field("arr").Index(1)) + require.NoError(t, err) + require.True(t, found) + assert.EqualValues(t, 20, leaf.Value()) +} + +func TestGetByPathAbsent(t *testing.T) { + v := buildVar(t, map[string]any{"a": int64(1), "arr": []any{int64(10)}}) + + for _, p := range []variant.VariantPath{ + variant.VariantPath{}.Field("missing"), // absent object field + variant.VariantPath{}.Field("arr").Index(9), // out-of-range index + variant.VariantPath{}.Field("a").Index(0), // index into a scalar + variant.VariantPath{}.Index(0), // index into an object + } { + _, found, err := v.GetByPath(p) + require.NoError(t, err) + assert.False(t, found) + } +} + +// TestGetByPathFieldOnScalarErrors: a field step into a non-object is a type error. +func TestGetByPathFieldOnScalarErrors(t *testing.T) { + v := buildVar(t, int64(1)) + _, _, err := v.GetByPath(variant.VariantPath{}.Field("a")) + require.ErrorIs(t, err, arrow.ErrInvalid) +} + +// TestGetByPathHugeIndex: a huge index must not wrap; it is simply absent. +func TestGetByPathHugeIndex(t *testing.T) { + v := buildVar(t, []any{int64(10), int64(20)}) + _, found, err := v.GetByPath(variant.VariantPath{}.Index(1 << 40)) + require.NoError(t, err) + assert.False(t, found) +} + +func TestVariantPathJoinAndStepAt(t *testing.T) { + p := variant.VariantPath{}.Field("a").Join(variant.VariantPath{}.Index(2).Field("b")) + require.Equal(t, 3, p.Len()) + + name, _ := p.StepAt(0) + assert.Equal(t, "a", name) + name, idx := p.StepAt(1) + assert.Equal(t, "", name) + assert.Equal(t, 2, idx) + name, _ = p.StepAt(2) + assert.Equal(t, "b", name) +} From 708d8e8b5e4f8acee9d2eb2335d02704e5e48ab4 Mon Sep 17 00:00:00 2001 From: Neelesh Salian Date: Fri, 28 Aug 2026 16:31:41 -0700 Subject: [PATCH 4/5] Fixing pending issues --- arrow/compute/variant_get.go | 161 +++++++--- arrow/compute/variant_get_test.go | 489 +++++++++++++++++++++++++++++- parquet/variant/path.go | 21 +- parquet/variant/path_test.go | 28 +- 4 files changed, 641 insertions(+), 58 deletions(-) diff --git a/arrow/compute/variant_get.go b/arrow/compute/variant_get.go index d23f4ece1..44ed15871 100644 --- a/arrow/compute/variant_get.go +++ b/arrow/compute/variant_get.go @@ -41,6 +41,7 @@ type VariantGetOptions struct { // Strict makes a lossy cast fail; the default allows overflow and truncation via // the cast kernels. Unlike arrow-rs safe mode there is no null-on-failure: an // impossible cast always errors, since arrow-go's cast kernels have no safe flag. + // Non-strict nulls a whole natural-type group if any value in it is inconvertible. Strict bool } @@ -54,6 +55,12 @@ func VariantGet(ctx context.Context, input *extensions.VariantArray, opts Varian return nil, fmt.Errorf("%w: VariantGet requires a non-nil VariantArray", arrow.ErrInvalid) } + // Nested target types are not yet supported; reject up front rather than + // silently producing an all-null array from the leaf cast. + if _, ok := opts.AsType.(arrow.NestedType); ok { + return nil, fmt.Errorf("%w: VariantGet cast to nested type %s", arrow.ErrNotImplemented, opts.AsType) + } + // Empty path, no cast: the values are returned unchanged. if opts.Path.Len() == 0 && opts.AsType == nil { input.Retain() @@ -101,7 +108,6 @@ type pathStepKind int const ( stepSuccess pathStepKind = iota stepMissing - stepNotShredded ) type pathStep struct { @@ -110,14 +116,15 @@ type pathStep struct { owned []arrow.Array // intermediate take results the caller must release } -// missingStep reports whether an absent typed field is provably missing (the value -// column is all-null) or merely not shredded (a residual may hold it). +// missingStep marks a path step whose typed field is absent. The descent loop breaks +// on any residual before stepping, so a step that reaches here is provably missing. func (s shreddingState) missingStep() pathStep { - if s.value == nil || s.value.NullN() == s.value.Len() { - return pathStep{kind: stepMissing} - } + return pathStep{kind: stepMissing} +} - return pathStep{kind: stepNotShredded} +// hasResidual reports whether any row carries a value in this level's value column. +func (s shreddingState) hasResidual() bool { + return s.value != nil && s.value.NullN() != s.value.Len() } func fieldStep(s shreddingState, name string) (pathStep, error) { @@ -126,7 +133,10 @@ func fieldStep(s shreddingState, name string) (pathStep, error) { } st, ok := s.typedValue.(*array.Struct) if !ok { - return s.missingStep(), nil + // A field step into a non-object shredded value is a type error, matching the + // per-row GetByPath path. Any residual was already diverted before this runs. + return pathStep{}, fmt.Errorf("%w: variant path field %q applied to non-object %s", + arrow.ErrInvalid, name, s.typedValue.DataType()) } idx, ok := st.DataType().(*arrow.StructType).FieldIdx(name) if !ok { @@ -215,12 +225,17 @@ func shreddedGetPath(ctx context.Context, input *extensions.VariantArray, opts V idx := 0 for idx < opts.Path.Len() { - name, index := opts.Path.StepAt(idx) + // residual-backed rows live in the value column, unreachable by the typed_value descent; reassemble per-row. + if state.hasResidual() { + break + } + + name, index, isField := opts.Path.StepAt(idx) var ( step pathStep err error ) - if name != "" { + if isField { step, err = fieldStep(state, name) } else { step, err = indexStep(ctx, mem, state, index) @@ -237,11 +252,9 @@ func shreddedGetPath(ctx context.Context, input *extensions.VariantArray, opts V continue } - if step.kind == stepMissing { - return allNullResult(mem, input.Len(), opts.AsType), nil - } - break // stepNotShredded + // stepMissing: the typed field is provably absent (no residual, checked above). + return allNullResult(mem, input.Len(), opts.AsType), nil } remaining := subPath(opts.Path, idx) @@ -276,22 +289,13 @@ func shreddedGetPath(ctx context.Context, input *extensions.VariantArray, opts V return buildLeafVariantArray(mem, leaves), nil } - src := buildNaturalArray(mem, leaves) - if src == nil { - return allNullResult(mem, len(leaves), opts.AsType), nil - } - defer src.Release() - - return CastArray(ctx, src, NewCastOptions(opts.AsType, opts.Strict)) + return castLeaves(ctx, mem, leaves, opts.AsType, opts.Strict) } // perfectShredded returns the typed_value column when the path landed on a fully // shredded value of exactly AsType and no ancestor nulls need merging; otherwise // the caller's reassembly path produces the same values. func perfectShredded(s shreddingState, nulls *nullTracker, asType arrow.DataType) arrow.Array { - if _, ok := asType.(arrow.NestedType); ok { - return nil - } if s.typedValue == nil || !nulls.allValid() { return nil } @@ -348,7 +352,7 @@ func buildTargetVariant(input *extensions.VariantArray, s shreddingState, nulls func subPath(p variant.VariantPath, from int) variant.VariantPath { var out variant.VariantPath for i := from; i < p.Len(); i++ { - if name, index := p.StepAt(i); name != "" { + if name, index, isField := p.StepAt(i); isField { out = out.Field(name) } else { out = out.Index(index) @@ -401,27 +405,100 @@ func buildLeafVariantArray(mem memory.Allocator, leaves []variantLeaf) arrow.Arr return bldr.NewArray() } -// buildNaturalArray materializes the leaves as an array of the first present leaf's -// natural Arrow type so the cast kernels can convert it. Rows whose value does not -// match that natural type become null. Returns nil when no leaf is present. -func buildNaturalArray(mem memory.Allocator, leaves []variantLeaf) arrow.Array { - var natural arrow.DataType - for _, l := range leaves { - if l.present && l.value.Type() != variant.Null { - natural = naturalArrowType(l.value) +// castLeaves converts each leaf to asType (arrow-rs variant_get parity): leaves are +// grouped by natural type, each group cast with the cast kernels, then scattered back +// so the result is order-independent. Strict errors on a lossy cast, else null. +func castLeaves(ctx context.Context, mem memory.Allocator, leaves []variantLeaf, asType arrow.DataType, strict bool) (arrow.Array, error) { + type leafGroup struct { + dt arrow.DataType + rows []int + } + groups := make(map[string]*leafGroup) + var order []string + for i, l := range leaves { + if !l.present || l.value.Type() == variant.Null { + continue + } + dt := naturalArrowType(l.value) + if dt == nil { + // Object/array leaves have no primitive natural type. Under Strict this is an + // impossible cast (errors like any other); otherwise the row stays null. + if strict { + return nil, fmt.Errorf("%w: cannot cast non-primitive variant leaf to %s", arrow.ErrInvalid, asType) + } - break + continue + } + key := dt.String() + g := groups[key] + if g == nil { + g = &leafGroup{dt: dt} + groups[key] = g + order = append(order, key) } + g.rows = append(g.rows, i) } - if natural == nil { - return nil + + perm := make([]uint64, len(leaves)) + valid := make([]bool, len(leaves)) + var casted []arrow.Array + defer func() { releaseAll(casted) }() + + var pos uint64 + for _, key := range order { + g := groups[key] + col := buildTypedColumn(mem, g.dt, leaves, g.rows) + cast, err := CastArray(ctx, col, NewCastOptions(asType, strict)) + col.Release() + if err != nil { + if strict { + return nil, err + } + // Non-strict: this natural type cannot convert to asType; its rows stay null. + continue + } + casted = append(casted, cast) + for _, row := range g.rows { + perm[row] = pos + valid[row] = true + pos++ + } + } + + if len(casted) == 0 { + return allNullResult(mem, len(leaves), asType), nil + } + + // One natural type covering every row already sits in original order. + if len(casted) == 1 && pos == uint64(len(leaves)) { + casted[0].Retain() + + return casted[0], nil + } + + combined, err := array.Concatenate(casted, mem) + if err != nil { + return nil, err } + defer combined.Release() - bldr := array.NewBuilder(mem, natural) + ib := array.NewUint64Builder(mem) + defer ib.Release() + ib.AppendValues(perm, valid) + indices := ib.NewArray() + defer indices.Release() + + return TakeArray(ctx, combined, indices) +} + +// buildTypedColumn materializes the given leaf rows, all of natural type dt, into a +// homogeneous Arrow array the cast kernels can consume. +func buildTypedColumn(mem memory.Allocator, dt arrow.DataType, leaves []variantLeaf, rows []int) arrow.Array { + bldr := array.NewBuilder(mem, dt) defer bldr.Release() - bldr.Reserve(len(leaves)) - for _, l := range leaves { - if !l.present || !appendNatural(bldr, l.value) { + bldr.Reserve(len(rows)) + for _, row := range rows { + if !appendNatural(bldr, leaves[row].value) { bldr.AppendNull() } } @@ -538,8 +615,8 @@ func naturalArrowType(v variant.Value) arrow.DataType { return nil } -// appendNatural appends v to bldr when v matches bldr's natural type, reporting -// whether it did; a non-matching value is left for the caller to null. +// appendNatural appends v to bldr when v's value matches bldr's element type, +// reporting whether it did; callers group leaves by natural type first, so it matches. func appendNatural(bldr array.Builder, v variant.Value) bool { switch b := bldr.(type) { case *array.BooleanBuilder: diff --git a/arrow/compute/variant_get_test.go b/arrow/compute/variant_get_test.go index 993463af9..07f832d54 100644 --- a/arrow/compute/variant_get_test.go +++ b/arrow/compute/variant_get_test.go @@ -24,6 +24,7 @@ import ( "github.com/apache/arrow-go/v18/arrow/array" "github.com/apache/arrow-go/v18/arrow/compute" "github.com/apache/arrow-go/v18/arrow/compute/exec" + "github.com/apache/arrow-go/v18/arrow/decimal" "github.com/apache/arrow-go/v18/arrow/extensions" "github.com/apache/arrow-go/v18/arrow/memory" "github.com/apache/arrow-go/v18/parquet/variant" @@ -307,14 +308,21 @@ func TestVariantGetNoLeak(t *testing.T) { }) require.NoError(t, err) idx.Release() + arr.Release() + + // A missing field on an object-shredded variant exercises the all-null path. + objVT := extensions.NewShreddedVariantType(arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64})) + ob := extensions.NewVariantBuilder(mem, objVT) + ob.Append(vgVariant(t, map[string]any{"a": int64(1)})) + objArr := ob.NewArray().(*extensions.VariantArray) + ob.Release() - missing, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{ + missing, err := compute.VariantGet(ctx, objArr, compute.VariantGetOptions{ Path: variant.VariantPath{}.Field("nope"), AsType: arrow.PrimitiveTypes.Int64, }) require.NoError(t, err) missing.Release() - - arr.Release() + objArr.Release() } // TestVariantGetDictMetadata guards against a panic when the metadata column is @@ -347,3 +355,478 @@ func TestVariantGetDictMetadata(t *testing.T) { require.NoError(t, err) assert.Equal(t, "resid", v.Value()) } + +// vgMixedResidual builds a two-row shredded array: row 0 shredded (fill populates typed_value, top value null), +// row 1 residual (top typed_value null, top value = resid). Storage comes from NewShreddedVariantType. +func vgMixedResidual(t *testing.T, mem memory.Allocator, shredType arrow.DataType, fill func(b array.Builder), resid variant.Value) *extensions.VariantArray { + t.Helper() + vt := extensions.NewShreddedVariantType(shredType) + s := vt.StorageType().(*arrow.StructType) + mIdx, _ := s.FieldIdx("metadata") + vIdx, _ := s.FieldIdx("value") + tIdx, _ := s.FieldIdx("typed_value") + + b := array.NewStructBuilder(mem, s) + defer b.Release() + mb := b.FieldBuilder(mIdx).(*array.BinaryBuilder) + vb := b.FieldBuilder(vIdx).(*array.BinaryBuilder) + + b.Append(true) + mb.Append(variant.EmptyMetadataBytes[:]) + vb.AppendNull() + fill(b.FieldBuilder(tIdx)) + + b.Append(true) + mb.Append(resid.Metadata().Bytes()) + vb.Append(resid.Bytes()) + b.FieldBuilder(tIdx).AppendNull() + + st := b.NewArray() + defer st.Release() + + return array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray) +} + +// TestVariantGetResidualBackedField covers zeroshade's blocking :233 case for a root +// field path: row 1's {"a":2} lives in the top residual, so $.a must return 2 not null. +func TestVariantGetResidualBackedField(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgMixedResidual(t, mem, arrow.StructOf(arrow.Field{Name: "a", Type: arrow.PrimitiveTypes.Int64}), + func(b array.Builder) { + tv := b.(*array.StructBuilder) + tv.Append(true) + a := tv.FieldBuilder(0).(*array.StructBuilder) + a.Append(true) + a.FieldBuilder(0).(*array.BinaryBuilder).AppendNull() + a.FieldBuilder(1).(*array.Int64Builder).Append(1) + }, vgVariant(t, map[string]any{"a": int64(2)})) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + require.Equal(t, 2, ints.Len()) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 2, ints.Value(1), "residual-backed row must be reassembled, not nulled") +} + +// TestVariantGetResidualBackedNestedField covers the nested-field case: $.a.b where +// row 1's whole {"a":{"b":6}} lives in the top residual. +func TestVariantGetResidualBackedNestedField(t *testing.T) { + mem := memory.DefaultAllocator + shred := arrow.StructOf(arrow.Field{Name: "a", Type: arrow.StructOf(arrow.Field{Name: "b", Type: arrow.PrimitiveTypes.Int64})}) + arr := vgMixedResidual(t, mem, shred, + func(b array.Builder) { + tv := b.(*array.StructBuilder) + tv.Append(true) + a := tv.FieldBuilder(0).(*array.StructBuilder) + a.Append(true) + a.FieldBuilder(0).(*array.BinaryBuilder).AppendNull() + aTV := a.FieldBuilder(1).(*array.StructBuilder) + aTV.Append(true) + bf := aTV.FieldBuilder(0).(*array.StructBuilder) + bf.Append(true) + bf.FieldBuilder(0).(*array.BinaryBuilder).AppendNull() + bf.FieldBuilder(1).(*array.Int64Builder).Append(5) + }, vgVariant(t, map[string]any{"a": map[string]any{"b": int64(6)}})) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Field("a").Field("b"), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + assert.EqualValues(t, 5, ints.Value(0)) + assert.EqualValues(t, 6, ints.Value(1), "residual-backed row must be reassembled, not nulled") +} + +// TestVariantGetResidualBackedListIndex covers the list-index case: [0] where row 1's +// whole [30,40] lives in the top residual. +func TestVariantGetResidualBackedListIndex(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgMixedResidual(t, mem, arrow.ListOf(arrow.PrimitiveTypes.Int64), + func(b array.Builder) { + lb := b.(*array.ListBuilder) + lb.Append(true) + el := lb.ValueBuilder().(*array.StructBuilder) + for _, v := range []int64{10, 20} { + el.Append(true) + el.FieldBuilder(0).(*array.BinaryBuilder).AppendNull() + el.FieldBuilder(1).(*array.Int64Builder).Append(v) + } + }, vgVariant(t, []any{int64(30), int64(40)})) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(0), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + assert.EqualValues(t, 10, ints.Value(0)) + assert.EqualValues(t, 30, ints.Value(1), "residual-backed row must be reassembled, not nulled") +} + +// TestVariantGetMixedWidthIntegers covers zeroshade :411: variant ints encode at their +// natural width, so a wider AsType must not drop rows whose leaf shredded narrower. +func TestVariantGetMixedWidthIntegers(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, + map[string]any{"a": int64(1)}, // int8 + map[string]any{"a": int64(1000)}, // int16 + map[string]any{"a": int64(5_000_000_000)}, // int64 + map[string]any{"a": int64(9007199254740993)}) // int64 > 2^53, must stay exact (no float intermediate) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + require.Equal(t, 4, ints.Len()) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 1000, ints.Value(1), "narrower-width leaf must not be dropped") + assert.EqualValues(t, 5_000_000_000, ints.Value(2)) + assert.EqualValues(t, 9007199254740993, ints.Value(3), "value > 2^53 must be exact, not routed through float64") +} + +// TestVariantGetHeterogeneousLeaves (arrow-rs parity): a valid int64 survives a +// narrower-typed sibling; only a non-numeric string nulls. +func TestVariantGetHeterogeneousLeaves(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, + map[string]any{"a": int64(1)}, // int8 encoding + map[string]any{"a": int64(5_000_000_000)}, // int64, does not fit int8 + map[string]any{"a": "x"}) // non-numeric -> null + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + assert.EqualValues(t, 1, ints.Value(0)) + assert.EqualValues(t, 5_000_000_000, ints.Value(1), "valid int64 must survive a narrower-typed sibling") + assert.True(t, ints.IsNull(2), "non-numeric string must be null") +} + +// TestVariantGetEmptyKey covers zeroshade's blocking path.go:42 case: an empty-string +// object key is a field step, not array index 0. +func TestVariantGetEmptyKey(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, map[string]any{"": int64(42)}) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field(""), AsType: arrow.PrimitiveTypes.Int64}) + require.NoError(t, err) + defer out.Release() + assert.EqualValues(t, 42, out.(*array.Int64).Value(0)) +} + +// TestVariantGetResidualBackedDeepField exercises the residual break after a successful +// columnar descent: row 0 is shredded through a.b, row 1's {"b":6} sits in a's residual. +func TestVariantGetResidualBackedDeepField(t *testing.T) { + mem := memory.DefaultAllocator + shred := arrow.StructOf(arrow.Field{Name: "a", Type: arrow.StructOf(arrow.Field{Name: "b", Type: arrow.PrimitiveTypes.Int64})}) + vt := extensions.NewShreddedVariantType(shred) + s := vt.StorageType().(*arrow.StructType) + mIdx, _ := s.FieldIdx("metadata") + vIdx, _ := s.FieldIdx("value") + tIdx, _ := s.FieldIdx("typed_value") + + b := array.NewStructBuilder(mem, s) + defer b.Release() + mb := b.FieldBuilder(mIdx).(*array.BinaryBuilder) + vb := b.FieldBuilder(vIdx).(*array.BinaryBuilder) + tvb := b.FieldBuilder(tIdx).(*array.StructBuilder) // struct{a} + aField := tvb.FieldBuilder(0).(*array.StructBuilder) + aVal := aField.FieldBuilder(0).(*array.BinaryBuilder) + aTyped := aField.FieldBuilder(1).(*array.StructBuilder) // struct{b} + bField := aTyped.FieldBuilder(0).(*array.StructBuilder) + bVal := bField.FieldBuilder(0).(*array.BinaryBuilder) + bTyped := bField.FieldBuilder(1).(*array.Int64Builder) + + // row 0: fully shredded a.b = 5 + b.Append(true) + mb.Append(variant.EmptyMetadataBytes[:]) + vb.AppendNull() + tvb.Append(true) + aField.Append(true) + aVal.AppendNull() + aTyped.Append(true) + bField.Append(true) + bVal.AppendNull() + bTyped.Append(5) + + // row 1: a is residual-backed with {"b":6} (top typed_value present, a.value set) + resid := vgVariant(t, map[string]any{"b": int64(6)}) + b.Append(true) + mb.Append(resid.Metadata().Bytes()) + vb.AppendNull() + tvb.Append(true) + aField.Append(true) + aVal.Append(resid.Bytes()) + aTyped.AppendNull() // a.typed_value null -> recurses bField/bVal/bTyped null + + st := b.NewArray() + defer st.Release() + arr := array.NewExtensionArrayWithStorage(vt, st).(*extensions.VariantArray) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Field("a").Field("b"), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + defer out.Release() + ints := out.(*array.Int64) + require.Equal(t, 2, ints.Len()) + assert.EqualValues(t, 5, ints.Value(0)) + assert.EqualValues(t, 6, ints.Value(1), "mid-level residual row must be reassembled, not nulled") +} + +// TestVariantGetResidualNoLeak guards the residual break path (buildTargetVariant + +// per-row reassembly) against leaks, which TestVariantGetNoLeak does not reach. +func TestVariantGetResidualNoLeak(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := exec.WithAllocator(context.Background(), mem) + + arr := vgMixedResidual(t, mem, arrow.ListOf(arrow.PrimitiveTypes.Int64), + func(bld array.Builder) { + lb := bld.(*array.ListBuilder) + lb.Append(true) + el := lb.ValueBuilder().(*array.StructBuilder) + for _, v := range []int64{10, 20} { + el.Append(true) + el.FieldBuilder(0).(*array.BinaryBuilder).AppendNull() + el.FieldBuilder(1).(*array.Int64Builder).Append(v) + } + }, vgVariant(t, []any{int64(30), int64(40)})) + + out, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{ + Path: variant.VariantPath{}.Index(0), AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + out.Release() + arr.Release() +} + +// vgNonShreddedVals builds a non-shredded array from pre-built values, so a test can +// control per-value encoding (e.g. timestamp unit) that vgNonShredded cannot. +func vgNonShreddedVals(t *testing.T, mem memory.Allocator, vals ...variant.Value) *extensions.VariantArray { + t.Helper() + bldr := extensions.NewVariantBuilder(mem, extensions.NewDefaultVariantType()) + defer bldr.Release() + for _, v := range vals { + bldr.Append(v) + } + + return bldr.NewArray().(*extensions.VariantArray) +} + +func vgTimestamp(t *testing.T, ts arrow.Timestamp, nano bool) variant.Value { + t.Helper() + var b variant.Builder + opts := []variant.AppendOpt{variant.OptTimestampUTC} + if nano { + opts = append(opts, variant.OptTimestampNano) + } + require.NoError(t, b.Append(ts, opts...)) + val, err := b.Build() + require.NoError(t, err) + + return val +} + +// TestVariantGetMixedFloatWidths (:411, floats): a Float(32) and a Double(64) leaf +// both survive a Float64 request, independent of order. +func TestVariantGetMixedFloatWidths(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, float32(1.5), float64(2.5)) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64}) + require.NoError(t, err) + defer out.Release() + f := out.(*array.Float64) + require.Equal(t, 2, f.Len()) + assert.InDelta(t, 1.5, f.Value(0), 1e-9) + assert.InDelta(t, 2.5, f.Value(1), 1e-9, "Double leaf must not be dropped by a Float first leaf") +} + +// TestVariantGetIntPlusFloat: a mixed int/float column cast to Float64 widens the int +// through the cast kernels rather than nulling it. +func TestVariantGetIntPlusFloat(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, int64(3), float64(2.5)) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64}) + require.NoError(t, err) + defer out.Release() + f := out.(*array.Float64) + assert.InDelta(t, 3.0, f.Value(0), 1e-9) + assert.InDelta(t, 2.5, f.Value(1), 1e-9) +} + +// TestVariantGetTypeOrderIndependent (:411): the same two values give the same result +// regardless of which row comes first. +func TestVariantGetTypeOrderIndependent(t *testing.T) { + mem := memory.DefaultAllocator + forward := vgNonShredded(t, mem, int64(3), float64(2.5)) + defer forward.Release() + reverse := vgNonShredded(t, mem, float64(2.5), int64(3)) + defer reverse.Release() + + optsF := compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64} + fwd, err := compute.VariantGet(context.Background(), forward, optsF) + require.NoError(t, err) + defer fwd.Release() + rev, err := compute.VariantGet(context.Background(), reverse, optsF) + require.NoError(t, err) + defer rev.Release() + + fa, ra := fwd.(*array.Float64), rev.(*array.Float64) + assert.InDelta(t, fa.Value(0), ra.Value(1), 1e-9) + assert.InDelta(t, fa.Value(1), ra.Value(0), 1e-9) + assert.False(t, fa.IsNull(0) || fa.IsNull(1) || ra.IsNull(0) || ra.IsNull(1), "no leaf dropped in either order") +} + +// TestVariantGetMixedTimestampUnits: a micros leaf and a nanos leaf of the same instant +// both land on it when cast to a nanos target (units are converted, not reinterpreted). +func TestVariantGetMixedTimestampUnits(t *testing.T) { + mem := memory.DefaultAllocator + const micros = arrow.Timestamp(1_600_000_000_000_000) + const nanos = arrow.Timestamp(1_600_000_000_000_000_000) + arr := vgNonShreddedVals(t, mem, vgTimestamp(t, micros, false), vgTimestamp(t, nanos, true)) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + AsType: &arrow.TimestampType{Unit: arrow.Nanosecond, TimeZone: "UTC"}, + }) + require.NoError(t, err) + defer out.Release() + ts := out.(*array.Timestamp) + require.Equal(t, 2, ts.Len()) + assert.EqualValues(t, nanos, ts.Value(0), "micros leaf must be scaled to nanos, not copied raw") + assert.EqualValues(t, nanos, ts.Value(1)) +} + +// TestVariantGetMixedDecimalScales: leaves shredded at different scales both rescale to +// the requested target scale. +func TestVariantGetMixedDecimalScales(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShreddedVals(t, mem, + vgVariant(t, variant.DecimalValue[decimal.Decimal32]{Scale: 1, Value: decimal.Decimal32(15)}), // 1.5 + vgVariant(t, variant.DecimalValue[decimal.Decimal32]{Scale: 2, Value: decimal.Decimal32(225)})) // 2.25 + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + AsType: &arrow.Decimal128Type{Precision: 38, Scale: 2}, + }) + require.NoError(t, err) + defer out.Release() + d := out.(*array.Decimal128) + require.Equal(t, 2, d.Len()) + assert.InDelta(t, 1.5, d.Value(0).ToFloat64(2), 1e-9, "scale-1 leaf must rescale to scale-2, not drop") + assert.InDelta(t, 2.25, d.Value(1).ToFloat64(2), 1e-9) +} + +// TestVariantGetStrictSlowPathErrors (:411/:269): on the reassembly path a lossy cast +// errors under Strict rather than nulling. +func TestVariantGetStrictSlowPathErrors(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, map[string]any{"a": int64(5_000_000_000)}) // overflows int32 + defer arr.Release() + + _, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: field("a"), AsType: arrow.PrimitiveTypes.Int32, Strict: true, + }) + require.Error(t, err, "Strict must error on an overflowing cast, not null it") +} + +// TestVariantGetMixedTypeNoLeak guards the multi-group scatter path (cast + Concatenate +// + Take), which the single-type leak tests do not reach. +// TestVariantGetNestedTypeNotImplemented pins that a nested AsType is rejected with +// ErrNotImplemented rather than silently producing an all-null array. +func TestVariantGetNestedTypeNotImplemented(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, map[string]any{"a": map[string]any{"x": int64(1)}}) + defer arr.Release() + + for _, nested := range []arrow.DataType{ + arrow.StructOf(arrow.Field{Name: "x", Type: arrow.PrimitiveTypes.Int64}), + arrow.ListOf(arrow.PrimitiveTypes.Int64), + } { + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: nested}) + if out != nil { + out.Release() + } + require.ErrorIs(t, err, arrow.ErrNotImplemented, "nested AsType %s must error, not null", nested) + } +} + +func TestVariantGetMixedTypeNoLeak(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := exec.WithAllocator(context.Background(), mem) + + arr := vgNonShredded(t, mem, int64(3), float64(2.5), nil, "x") + out, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64}) + require.NoError(t, err) + out.Release() + arr.Release() +} + +// TestVariantGetInterleavedScatter exercises the multi-group scatter (Concatenate + +// TakeArray) with a NON-IDENTITY permutation: two same-typed leaves straddle a +// differently-typed one, so group order [3,7,2.5] must scatter back to row order +// [3,2.5,7]. Every other multi-group test lands in identity order. +func TestVariantGetInterleavedScatter(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgNonShredded(t, mem, int64(3), float64(2.5), int64(7)) // Int8{0,2}, Float64{1} + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{AsType: arrow.PrimitiveTypes.Float64}) + require.NoError(t, err) + defer out.Release() + f := out.(*array.Float64) + require.Equal(t, 3, f.Len()) + assert.InDelta(t, 3.0, f.Value(0), 1e-9) + assert.InDelta(t, 2.5, f.Value(1), 1e-9, "interleaved leaf must scatter back to its row, not stay in group order") + assert.InDelta(t, 7.0, f.Value(2), 1e-9) +} + +// TestVariantGetStrictObjectLeafErrors pins that under Strict an object/array leaf cast +// to a primitive errors (impossible cast) rather than silently nulling. +func TestVariantGetStrictObjectLeafErrors(t *testing.T) { + mem := memory.DefaultAllocator + for _, v := range []any{ + map[string]any{"a": map[string]any{"x": int64(1)}}, // object leaf at $.a + map[string]any{"a": []any{int64(1), int64(2)}}, // array leaf at $.a + } { + arr := vgNonShredded(t, mem, v) + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: field("a"), AsType: arrow.PrimitiveTypes.Int64, Strict: true, + }) + if out != nil { + out.Release() + } + arr.Release() + require.ErrorIs(t, err, arrow.ErrInvalid, "strict cast of a non-primitive leaf to Int64 must error") + } +} + +// TestVariantGetShreddedFieldOnScalarErrors pins that a field step into a shredded +// scalar column errors on the columnar path, matching the per-row GetByPath path. +func TestVariantGetShreddedFieldOnScalarErrors(t *testing.T) { + mem := memory.DefaultAllocator + arr := vgShreddedInt(t, mem, 1, 2, 3) // typed_value is a scalar Int64 column + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{Path: field("a"), AsType: arrow.PrimitiveTypes.Int64}) + if out != nil { + out.Release() + } + require.ErrorIs(t, err, arrow.ErrInvalid, "field access on a shredded scalar must error, not return null") +} diff --git a/parquet/variant/path.go b/parquet/variant/path.go index 437e08529..bc46da77d 100644 --- a/parquet/variant/path.go +++ b/parquet/variant/path.go @@ -24,11 +24,11 @@ import ( "github.com/apache/arrow-go/v18/arrow" ) -// pathElem is one step of a VariantPath: an object field when name != "", else an -// array index. +// pathElem is one step of a VariantPath: an object field when isField is set, else an array index. type pathElem struct { - name string - index int + name string + index int + isField bool } // VariantPath is an ordered list of steps to navigate into a variant value. The @@ -39,7 +39,7 @@ type VariantPath struct { // Field returns a copy of the path with an object-field step appended. func (p VariantPath) Field(name string) VariantPath { - return VariantPath{elems: append(p.grow(), pathElem{name: name})} + return VariantPath{elems: append(p.grow(), pathElem{name: name, isField: true})} } // Index returns a copy of the path with an array-index step appended. @@ -59,10 +59,11 @@ func (p VariantPath) grow() []pathElem { // Len returns the number of steps in the path. func (p VariantPath) Len() int { return len(p.elems) } -// StepAt returns the i-th step. When name != "" it is an object-field step; -// otherwise it is an array-index step selecting index. -func (p VariantPath) StepAt(i int) (name string, index int) { - return p.elems[i].name, p.elems[i].index +// StepAt returns the i-th step's name and index, with isField true for an object-field step. +func (p VariantPath) StepAt(i int) (name string, index int, isField bool) { + e := p.elems[i] + + return e.name, e.index, e.isField } // GetByPath navigates path into v and returns the leaf value. found is false when @@ -72,7 +73,7 @@ func (p VariantPath) StepAt(i int) (name string, index int) { func (v Value) GetByPath(path VariantPath) (leaf Value, found bool, err error) { cur := v for _, e := range path.elems { - if e.name != "" { + if e.isField { obj, ok := cur.Value().(ObjectValue) if !ok { return Value{}, false, fmt.Errorf("%w: variant path field %q applied to non-object", arrow.ErrInvalid, e.name) diff --git a/parquet/variant/path_test.go b/parquet/variant/path_test.go index 6a8cfaaea..9de8b84ad 100644 --- a/parquet/variant/path_test.go +++ b/parquet/variant/path_test.go @@ -84,11 +84,33 @@ func TestVariantPathJoinAndStepAt(t *testing.T) { p := variant.VariantPath{}.Field("a").Join(variant.VariantPath{}.Index(2).Field("b")) require.Equal(t, 3, p.Len()) - name, _ := p.StepAt(0) + name, _, isField := p.StepAt(0) assert.Equal(t, "a", name) - name, idx := p.StepAt(1) + assert.True(t, isField) + name, idx, isField := p.StepAt(1) assert.Equal(t, "", name) assert.Equal(t, 2, idx) - name, _ = p.StepAt(2) + assert.False(t, isField) + name, _, isField = p.StepAt(2) assert.Equal(t, "b", name) + assert.True(t, isField) +} + +// TestGetByPathEmptyKey covers the empty-string object key: Field("") is a field +// step distinct from Index(0), so it must resolve the "" key rather than index 0. +func TestGetByPathEmptyKey(t *testing.T) { + v := buildVar(t, map[string]any{"": int64(42)}) + + _, _, isField := variant.VariantPath{}.Field("").StepAt(0) + assert.True(t, isField, `Field("") must be a field step, not an index step`) + + leaf, found, err := v.GetByPath(variant.VariantPath{}.Field("")) + require.NoError(t, err) + require.True(t, found) + assert.EqualValues(t, 42, leaf.Value()) + + // Index(0) on the object must not match the "" key. + _, found, err = v.GetByPath(variant.VariantPath{}.Index(0)) + require.NoError(t, err) + assert.False(t, found) } From 420b64ebfb05419abe8edc93d1b834eefb4e3c7a Mon Sep 17 00:00:00 2001 From: Neelesh Salian Date: Mon, 31 Aug 2026 14:54:51 -0700 Subject: [PATCH 5/5] UUID fix --- arrow/compute/variant_get.go | 39 +++++++++++++++++++--------- arrow/compute/variant_get_test.go | 43 +++++++++++++++++++++++++++++++ 2 files changed, 70 insertions(+), 12 deletions(-) diff --git a/arrow/compute/variant_get.go b/arrow/compute/variant_get.go index 44ed15871..8f5a2e2af 100644 --- a/arrow/compute/variant_get.go +++ b/arrow/compute/variant_get.go @@ -38,10 +38,8 @@ type VariantGetOptions struct { // AsType, when nil, makes VariantGet return a VariantArray pointing at the path; // when set, the extracted values are cast to it via the cast kernels. AsType arrow.DataType - // Strict makes a lossy cast fail; the default allows overflow and truncation via - // the cast kernels. Unlike arrow-rs safe mode there is no null-on-failure: an - // impossible cast always errors, since arrow-go's cast kernels have no safe flag. - // Non-strict nulls a whole natural-type group if any value in it is inconvertible. + // Strict makes a lossy cast fail; otherwise a value that cannot convert to AsType + // nulls only that row, not the rest of its natural-type group. Strict bool } @@ -55,9 +53,13 @@ func VariantGet(ctx context.Context, input *extensions.VariantArray, opts Varian return nil, fmt.Errorf("%w: VariantGet requires a non-nil VariantArray", arrow.ErrInvalid) } - // Nested target types are not yet supported; reject up front rather than - // silently producing an all-null array from the leaf cast. - if _, ok := opts.AsType.(arrow.NestedType); ok { + // Reject nested target storage. Unwrap extension types first: they embed + // ExtensionBase and so satisfy arrow.NestedType even when backed by e.g. UUID. + nestedCheck := opts.AsType + if ext, ok := nestedCheck.(arrow.ExtensionType); ok { + nestedCheck = ext.StorageType() + } + if _, ok := nestedCheck.(arrow.NestedType); ok { return nil, fmt.Errorf("%w: VariantGet cast to nested type %s", arrow.ErrNotImplemented, opts.AsType) } @@ -450,15 +452,28 @@ func castLeaves(ctx context.Context, mem memory.Allocator, leaves []variantLeaf, col := buildTypedColumn(mem, g.dt, leaves, g.rows) cast, err := CastArray(ctx, col, NewCastOptions(asType, strict)) col.Release() - if err != nil { - if strict { - return nil, err + if err == nil { + casted = append(casted, cast) + for _, row := range g.rows { + perm[row] = pos + valid[row] = true + pos++ } - // Non-strict: this natural type cannot convert to asType; its rows stay null. + continue } - casted = append(casted, cast) + if strict { + return nil, err + } + // Non-strict: retry each row alone so only the inconvertible rows null, not the whole group. for _, row := range g.rows { + rc := buildTypedColumn(mem, g.dt, leaves, []int{row}) + one, cerr := CastArray(ctx, rc, NewCastOptions(asType, strict)) + rc.Release() + if cerr != nil { + continue // this row alone is inconvertible; it stays null + } + casted = append(casted, one) perm[row] = pos valid[row] = true pos++ diff --git a/arrow/compute/variant_get_test.go b/arrow/compute/variant_get_test.go index 07f832d54..cf841ba28 100644 --- a/arrow/compute/variant_get_test.go +++ b/arrow/compute/variant_get_test.go @@ -28,6 +28,7 @@ import ( "github.com/apache/arrow-go/v18/arrow/extensions" "github.com/apache/arrow-go/v18/arrow/memory" "github.com/apache/arrow-go/v18/parquet/variant" + "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -830,3 +831,45 @@ func TestVariantGetShreddedFieldOnScalarErrors(t *testing.T) { } require.ErrorIs(t, err, arrow.ErrInvalid, "field access on a shredded scalar must error, not return null") } + +// TestVariantGetUUIDTarget pins that a UUID AsType reaches the UUID cast, not the nested-type reject. +func TestVariantGetUUIDTarget(t *testing.T) { + mem := memory.DefaultAllocator + u := uuid.MustParse("00112233-4455-6677-8899-aabbccddeeff") + arr := vgNonShredded(t, mem, map[string]any{"id": u}) + defer arr.Release() + + out, err := compute.VariantGet(context.Background(), arr, compute.VariantGetOptions{ + Path: field("id"), AsType: extensions.NewUUIDType(), + }) + require.NoError(t, err, "UUID target must not be rejected as a nested type") + defer out.Release() + + uarr, ok := out.(*extensions.UUIDArray) + require.True(t, ok, "expected *extensions.UUIDArray, got %T", out) + require.Equal(t, 1, uarr.Len()) + require.False(t, uarr.IsNull(0)) + assert.Equal(t, u, uarr.Value(0)) +} + +// TestVariantGetPartialCastFailure pins that non-strict nulls only the inconvertible row: ["1","bad"]->int64 is [1,null]. +func TestVariantGetPartialCastFailure(t *testing.T) { + mem := memory.NewCheckedAllocator(memory.DefaultAllocator) + defer mem.AssertSize(t, 0) + ctx := exec.WithAllocator(context.Background(), mem) + arr := vgNonShredded(t, mem, "1", "bad") + + out, err := compute.VariantGet(ctx, arr, compute.VariantGetOptions{ + AsType: arrow.PrimitiveTypes.Int64, + }) + require.NoError(t, err) + + ints := out.(*array.Int64) + require.Equal(t, 2, ints.Len()) + assert.False(t, ints.IsNull(0), "valid row must survive a sibling row's cast failure") + assert.EqualValues(t, 1, ints.Value(0)) + assert.True(t, ints.IsNull(1), "only the inconvertible row is null") + + out.Release() + arr.Release() +}