This is an automated email from the ASF dual-hosted git repository.

zeroshade pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-go.git


The following commit(s) were added to refs/heads/main by this push:
     new d0de6ddd fix(arrow/extensions): use canonical Variant extension name 
(#1241)
d0de6ddd is described below

commit d0de6dddd58dfc4b092e446277b6113fc6e04841
Author: Digvijay <[email protected]>
AuthorDate: Wed Sep 2 13:45:30 2026 -0500

    fix(arrow/extensions): use canonical Variant extension name (#1241)
    
    ### Rationale for this change
    
    The canonical Variant extension type name is `arrow.parquet.variant`.
    arrow-go still writes `parquet.variant`, which predates
    https://github.com/apache/arrow/pull/47456. That makes Variant IPC
    unreadable by other Arrow implementations, and vice versa.
    
    C++ already made this change in apache/arrow#49082.
    
    Fixes #1203
    
    ### What changes are included in this PR?
    
    - `VariantType.ExtensionName()` now returns `arrow.parquet.variant`
    - Keep `parquet.variant` registered so older arrow-go IPC still
    deserializes as `VariantType`
    - Parquet schema conversion recognizes both names
    
    ### Are these changes tested?
    
    - `go test ./arrow/extensions`
    - `go test ./parquet/pqarrow -run
    'TestConvertSchemaParquetVariant|TestShreddedVariantSchema'`
    
    ### Are there any user-facing changes?
    
    Yes — newly written IPC/Flight metadata uses `arrow.parquet.variant`.
    Existing `parquet.variant` data still reads.
    
    ---------
    
    Signed-off-by: Digvijay <[email protected]>
---
 arrow/extensions/extensions.go   |  30 +++++++
 arrow/extensions/variant.go      |  50 +++++++++---
 arrow/extensions/variant_test.go | 166 ++++++++++++++++++++++++++++++++++++++-
 parquet/pqarrow/schema.go        |   8 +-
 parquet/pqarrow/schema_test.go   |  28 ++++++-
 5 files changed, 269 insertions(+), 13 deletions(-)

diff --git a/arrow/extensions/extensions.go b/arrow/extensions/extensions.go
index 6f13aa64..c422a14c 100644
--- a/arrow/extensions/extensions.go
+++ b/arrow/extensions/extensions.go
@@ -17,6 +17,8 @@
 package extensions
 
 import (
+       "fmt"
+
        "github.com/apache/arrow-go/v18/arrow"
 )
 
@@ -35,4 +37,32 @@ func init() {
                        panic(err)
                }
        }
+
+       // arrow-go originally registered Variant as parquet.variant. Keep that
+       // name in the registry so older IPC still deserializes. Seed default
+       // storage so GetExtensionType("parquet.variant") is a complete type.
+       if err := arrow.RegisterExtensionType(&legacyVariantType{VariantType: 
*NewDefaultVariantType()}); err != nil {
+               panic(err)
+       }
+}
+
+// legacyVariantType is a compatibility adapter for the historical
+// parquet.variant name. Deserialize always returns a canonical VariantType;
+// newly written IPC uses VariantExtensionName.
+type legacyVariantType struct {
+       VariantType
+}
+
+func (*legacyVariantType) ExtensionName() string { return 
LegacyVariantExtensionName }
+
+func (v *legacyVariantType) String() string {
+       return fmt.Sprintf("extension<%s>", v.ExtensionName())
+}
+
+func (v *legacyVariantType) ExtensionEquals(other arrow.ExtensionType) bool {
+       return variantExtensionEquals(v.StorageType(), other)
+}
+
+func (*legacyVariantType) Deserialize(storageType arrow.DataType, _ string) 
(arrow.ExtensionType, error) {
+       return NewVariantType(storageType)
 }
diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go
index 5222d5fb..853b36ed 100644
--- a/arrow/extensions/variant.go
+++ b/arrow/extensions/variant.go
@@ -265,15 +265,47 @@ func (v *VariantType) TypedValue() arrow.Field {
        return v.StorageType().(*arrow.StructType).Field(v.typedValueFieldIdx)
 }
 
-func (*VariantType) ExtensionName() string { return "parquet.variant" }
+const (
+       // VariantExtensionName is the canonical Arrow extension type name.
+       // See 
https://arrow.apache.org/docs/format/CanonicalExtensions.html#parquet-variant
+       VariantExtensionName = "arrow.parquet.variant"
+
+       // LegacyVariantExtensionName was used by arrow-go before the canonical
+       // name landed. It is still accepted when reading IPC so older data
+       // continues to deserialize as VariantType.
+       LegacyVariantExtensionName = "parquet.variant"
+)
+
+// IsVariantExtensionName reports whether name is the canonical or historical
+// Variant extension type name.
+func IsVariantExtensionName(name string) bool {
+       return name == VariantExtensionName || name == 
LegacyVariantExtensionName
+}
+
+func (*VariantType) ExtensionName() string { return VariantExtensionName }
 
 func (v *VariantType) String() string {
        return fmt.Sprintf("extension<%s>", v.ExtensionName())
 }
 
+func variantExtensionEquals(storage arrow.DataType, other arrow.ExtensionType) 
bool {
+       return IsVariantExtensionName(other.ExtensionName()) &&
+               arrow.TypeEqual(storage, other.StorageType())
+}
+
 func (v *VariantType) ExtensionEquals(other arrow.ExtensionType) bool {
-       return v.ExtensionName() == other.ExtensionName() &&
-               arrow.TypeEqual(v.Storage, other.StorageType())
+       return variantExtensionEquals(v.Storage, other)
+}
+
+func asVariantType(dt arrow.ExtensionType) *VariantType {
+       if vt, ok := dt.(*VariantType); ok {
+               return vt
+       }
+       vt, err := NewVariantType(dt.StorageType())
+       if err != nil {
+               panic(err)
+       }
+       return vt
 }
 
 func (*VariantType) Serialize() string { return "" }
@@ -398,7 +430,7 @@ func (v *VariantArray) initReader() {
        // initialize a reader that coalesces shredded fields back into a 
variant
        // or just returns the basic variants if the array is not shredded.
        v.initRdr.Do(func() {
-               vt := v.ExtensionType().(*VariantType)
+               vt := asVariantType(v.ExtensionType())
                st := v.Storage().(*array.Struct)
                metaField := st.Field(vt.metadataFieldIdx)
                metadata, ok := metaField.(arrow.TypedArray[[]byte])
@@ -439,7 +471,7 @@ func (v *VariantArray) initReader() {
 // Metadata returns the metadata column of the variant array, containing the
 // metadata for each variant value.
 func (v *VariantArray) Metadata() arrow.TypedArray[[]byte] {
-       vt := v.ExtensionType().(*VariantType)
+       vt := asVariantType(v.ExtensionType())
        return 
v.Storage().(*array.Struct).Field(vt.metadataFieldIdx).(arrow.TypedArray[[]byte])
 }
 
@@ -458,7 +490,7 @@ func (v *VariantArray) Metadata() arrow.TypedArray[[]byte] {
 // it means that the value is missing entirely (as opposed to existing and 
having a
 // value of null).
 func (v *VariantArray) UntypedValues() arrow.TypedArray[[]byte] {
-       vt := v.ExtensionType().(*VariantType)
+       vt := asVariantType(v.ExtensionType())
        if vt.valueFieldIdx == -1 {
                return nil
        }
@@ -472,7 +504,7 @@ func (v *VariantArray) UntypedValues() 
arrow.TypedArray[[]byte] {
 // The reason for exposing this is to allow users to quickly access one of the 
shredded
 // fields without having to decode the entire variant value.
 func (v *VariantArray) Shredded() arrow.Array {
-       vt := v.ExtensionType().(*VariantType)
+       vt := asVariantType(v.ExtensionType())
        if vt.typedValueFieldIdx == -1 {
                return nil
        }
@@ -482,7 +514,7 @@ func (v *VariantArray) Shredded() arrow.Array {
 
 // IsShredded returns true if the variant has shredded columns.
 func (v *VariantArray) IsShredded() bool {
-       return v.ExtensionType().(*VariantType).typedValueFieldIdx != -1
+       return asVariantType(v.ExtensionType()).typedValueFieldIdx != -1
 }
 
 // VariantType returns the array's extension type without the ExtensionType 
cast.
@@ -537,7 +569,7 @@ func (v *VariantArray) IsNull(i int) bool {
                return true
        }
 
-       vt := v.ExtensionType().(*VariantType)
+       vt := asVariantType(v.ExtensionType())
        if vt.typedValueFieldIdx != -1 {
                typedArr := 
v.Storage().(*array.Struct).Field(vt.typedValueFieldIdx)
                if !typedArr.IsNull(i) {
diff --git a/arrow/extensions/variant_test.go b/arrow/extensions/variant_test.go
index 905fc5d9..babf065f 100644
--- a/arrow/extensions/variant_test.go
+++ b/arrow/extensions/variant_test.go
@@ -17,6 +17,7 @@
 package extensions_test
 
 import (
+       "bytes"
        "encoding/json"
        "fmt"
        "testing"
@@ -27,6 +28,7 @@ import (
        "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/ipc"
        "github.com/apache/arrow-go/v18/arrow/memory"
        "github.com/apache/arrow-go/v18/parquet/variant"
        "github.com/google/uuid"
@@ -44,7 +46,8 @@ func TestVariantExtensionType(t *testing.T) {
                arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, 
Nullable: false}))
        require.NoError(t, err)
 
-       assert.Equal(t, "extension<parquet.variant>", variant1.String())
+       assert.Equal(t, "arrow.parquet.variant", variant1.ExtensionName())
+       assert.Equal(t, "extension<arrow.parquet.variant>", variant1.String())
        assert.True(t, arrow.TypeEqual(variant1, variant2))
 
        // can be provided in either order
@@ -56,6 +59,10 @@ func TestVariantExtensionType(t *testing.T) {
        assert.Equal(t, "metadata", variantFieldsFlipped.Metadata().Name)
        assert.Equal(t, "value", variantFieldsFlipped.Value().Name)
 
+       assert.True(t, 
extensions.IsVariantExtensionName(extensions.VariantExtensionName))
+       assert.True(t, 
extensions.IsVariantExtensionName(extensions.LegacyVariantExtensionName))
+       assert.False(t, extensions.IsVariantExtensionName("arrow.uuid"))
+
        tests := []struct {
                dt          arrow.DataType
                expectedErr string
@@ -113,6 +120,163 @@ func TestVariantExtensionType(t *testing.T) {
        }
 }
 
+func TestVariantExtensionNameCanonicalAndLegacy(t *testing.T) {
+       storage := arrow.StructOf(
+               arrow.Field{Name: "metadata", Type: arrow.BinaryTypes.Binary, 
Nullable: false},
+               arrow.Field{Name: "value", Type: arrow.BinaryTypes.Binary, 
Nullable: false})
+       want, err := extensions.NewVariantType(storage)
+       require.NoError(t, err)
+
+       for _, name := range []string{
+               extensions.VariantExtensionName,
+               extensions.LegacyVariantExtensionName,
+       } {
+               t.Run(name, func(t *testing.T) {
+                       ext := arrow.GetExtensionType(name)
+                       require.NotNil(t, ext)
+                       got, err := ext.Deserialize(storage, "")
+                       require.NoError(t, err)
+                       assert.Equal(t, extensions.VariantExtensionName, 
got.ExtensionName())
+                       assert.True(t, arrow.TypeEqual(want, got))
+               })
+       }
+}
+
+func TestVariantTypeBatchIPCRoundTrip(t *testing.T) {
+       typ := extensions.NewDefaultVariantType()
+       bldr := extensions.NewVariantBuilder(memory.DefaultAllocator, typ)
+       defer bldr.Release()
+
+       var b variant.Builder
+       require.NoError(t, b.Append("hello"))
+       v, err := b.Build()
+       require.NoError(t, err)
+       bldr.Append(v)
+       bldr.AppendNull()
+
+       arr := bldr.NewArray()
+       defer arr.Release()
+
+       batch := array.NewRecordBatch(arrow.NewSchema([]arrow.Field{{Name: 
"field", Type: typ, Nullable: true}}, nil),
+               []arrow.Array{arr}, -1)
+       defer batch.Release()
+
+       var buf bytes.Buffer
+       wr := ipc.NewWriter(&buf, ipc.WithSchema(batch.Schema()))
+       require.NoError(t, wr.Write(batch))
+       require.NoError(t, wr.Close())
+
+       var written arrow.RecordBatch
+       {
+               rdr, err := ipc.NewReader(&buf)
+               require.NoError(t, err)
+               written, err = rdr.Read()
+               require.NoError(t, err)
+               written.Retain()
+               defer written.Release()
+               rdr.Release()
+       }
+
+       assert.Equal(t, extensions.VariantExtensionName, 
written.Schema().Field(0).Type.(arrow.ExtensionType).ExtensionName())
+       assert.Truef(t, batch.Schema().Equal(written.Schema()), "expected: %s, 
got: %s",
+               batch.Schema(), written.Schema())
+       assert.Truef(t, array.RecordEqual(batch, written), "expected: %s, got: 
%s",
+               batch, written)
+}
+
+func TestLegacyVariantTypeAdapter(t *testing.T) {
+       canonical := extensions.NewDefaultVariantType()
+       legacy := arrow.GetExtensionType(extensions.LegacyVariantExtensionName)
+       require.NotNil(t, legacy)
+       _, isCanonical := legacy.(*extensions.VariantType)
+       require.False(t, isCanonical)
+
+       assert.Equal(t, extensions.LegacyVariantExtensionName, 
legacy.ExtensionName())
+       assert.Equal(t, "extension<"+extensions.LegacyVariantExtensionName+">", 
legacy.String())
+
+       assert.True(t, canonical.ExtensionEquals(legacy))
+       assert.True(t, legacy.ExtensionEquals(canonical))
+       assert.True(t, arrow.TypeEqual(canonical, legacy))
+       assert.True(t, arrow.TypeEqual(legacy, canonical))
+}
+
+func TestLegacyVariantArrayAccessors(t *testing.T) {
+       canonical := extensions.NewDefaultVariantType()
+       bldr := extensions.NewVariantBuilder(memory.DefaultAllocator, canonical)
+       defer bldr.Release()
+
+       var b variant.Builder
+       require.NoError(t, b.Append("hello"))
+       v, err := b.Build()
+       require.NoError(t, err)
+       bldr.Append(v)
+
+       src := bldr.NewArray().(*extensions.VariantArray)
+       defer src.Release()
+
+       legacy := arrow.GetExtensionType(extensions.LegacyVariantExtensionName)
+       arr := array.NewExtensionArrayWithStorage(legacy, src.Storage())
+       defer arr.Release()
+
+       varr := arr.(*extensions.VariantArray)
+       assert.Equal(t, extensions.LegacyVariantExtensionName, 
varr.ExtensionType().ExtensionName())
+       assert.False(t, varr.IsShredded())
+       assert.NotNil(t, varr.Metadata())
+       assert.NotNil(t, varr.UntypedValues())
+       assert.Nil(t, varr.Shredded())
+       assert.False(t, varr.IsNull(0))
+
+       got, err := varr.Value(0)
+       require.NoError(t, err)
+       assert.Equal(t, "hello", got.Value())
+}
+
+func TestLegacyVariantIPCCanonicalizes(t *testing.T) {
+       canonical := extensions.NewDefaultVariantType()
+       bldr := extensions.NewVariantBuilder(memory.DefaultAllocator, canonical)
+       defer bldr.Release()
+
+       var b variant.Builder
+       require.NoError(t, b.Append("hello"))
+       v, err := b.Build()
+       require.NoError(t, err)
+       bldr.Append(v)
+
+       src := bldr.NewArray().(*extensions.VariantArray)
+       defer src.Release()
+
+       legacy := arrow.GetExtensionType(extensions.LegacyVariantExtensionName)
+       arr := array.NewExtensionArrayWithStorage(legacy, src.Storage())
+       defer arr.Release()
+
+       batch := array.NewRecordBatch(arrow.NewSchema([]arrow.Field{{Name: 
"field", Type: legacy, Nullable: true}}, nil),
+               []arrow.Array{arr}, -1)
+       defer batch.Release()
+
+       var buf bytes.Buffer
+       wr := ipc.NewWriter(&buf, ipc.WithSchema(batch.Schema()))
+       require.NoError(t, wr.Write(batch))
+       require.NoError(t, wr.Close())
+
+       rdr, err := ipc.NewReader(&buf)
+       require.NoError(t, err)
+       written, err := rdr.Read()
+       require.NoError(t, err)
+       written.Retain()
+       defer written.Release()
+       rdr.Release()
+
+       got := written.Schema().Field(0).Type.(arrow.ExtensionType)
+       assert.Equal(t, extensions.VariantExtensionName, got.ExtensionName())
+       _, isCanonical := got.(*extensions.VariantType)
+       assert.True(t, isCanonical)
+
+       want := array.NewRecordBatch(arrow.NewSchema([]arrow.Field{{Name: 
"field", Type: canonical, Nullable: true}}, nil),
+               []arrow.Array{src}, -1)
+       defer want.Release()
+       assert.Truef(t, array.RecordEqual(want, written), "expected: %s, got: 
%s", want, written)
+}
+
 func TestVariantExtensionBadNestedTypes(t *testing.T) {
        tests := []struct {
                name string
diff --git a/parquet/pqarrow/schema.go b/parquet/pqarrow/schema.go
index 2a13ea33..eeee4dbb 100644
--- a/parquet/pqarrow/schema.go
+++ b/parquet/pqarrow/schema.go
@@ -356,8 +356,12 @@ func fieldToNode(name string, field arrow.Field, props 
*parquet.WriterProperties
                return schema.MapOf(field.Name, keyNode, valueNode, 
repFromNullable(field.Nullable), fieldIDFromMeta(field.Metadata))
        case arrow.EXTENSION:
                extType := field.Type.(arrow.ExtensionType)
-               if extType.ExtensionName() == "parquet.variant" {
-                       return variantToNode(extType.(*extensions.VariantType), 
field, props, arrprops)
+               if extensions.IsVariantExtensionName(extType.ExtensionName()) {
+                       vt, err := 
extensions.NewVariantType(extType.StorageType())
+                       if err != nil {
+                               return nil, err
+                       }
+                       return variantToNode(vt, field, props, arrprops)
                }
        }
 
diff --git a/parquet/pqarrow/schema_test.go b/parquet/pqarrow/schema_test.go
index 6d9465cb..1144d73c 100644
--- a/parquet/pqarrow/schema_test.go
+++ b/parquet/pqarrow/schema_test.go
@@ -1224,13 +1224,39 @@ func TestConvertSchemaParquetVariant(t *testing.T) {
        assert.Equal(t, "variant_unshredded", outSchema.Field(0).Name)
        assert.Equal(t, arrow.EXTENSION, outSchema.Field(0).Type.ID())
 
-       assert.Equal(t, "parquet.variant", 
outSchema.Field(0).Type.(arrow.ExtensionType).ExtensionName())
+       assert.Equal(t, "arrow.parquet.variant", 
outSchema.Field(0).Type.(arrow.ExtensionType).ExtensionName())
 
        sc, err := pqarrow.ToParquet(outSchema, nil, 
pqarrow.DefaultWriterProps())
        require.NoError(t, err)
        assert.True(t, pqschema.Equals(sc), pqschema.String(), sc.String())
 }
 
+func TestToParquetLegacyRegistryVariantType(t *testing.T) {
+       ext := arrow.GetExtensionType(extensions.LegacyVariantExtensionName)
+       require.NotNil(t, ext)
+       require.Equal(t, extensions.LegacyVariantExtensionName, 
ext.ExtensionName())
+       _, isCanonical := ext.(*extensions.VariantType)
+       require.False(t, isCanonical, "registry type must be the legacy 
wrapper, not *VariantType")
+
+       legacySchema := arrow.NewSchema([]arrow.Field{{
+               Name:     "variant_col",
+               Type:     ext,
+               Nullable: true,
+       }}, nil)
+       canonicalSchema := arrow.NewSchema([]arrow.Field{{
+               Name:     "variant_col",
+               Type:     extensions.NewDefaultVariantType(),
+               Nullable: true,
+       }}, nil)
+
+       got, err := pqarrow.ToParquet(legacySchema, nil, 
pqarrow.DefaultWriterProps())
+       require.NoError(t, err)
+
+       want, err := pqarrow.ToParquet(canonicalSchema, nil, 
pqarrow.DefaultWriterProps())
+       require.NoError(t, err)
+       assert.True(t, want.Equals(got), "expected: %s\ngot: %s", want, got)
+}
+
 func TestShreddedVariantSchema(t *testing.T) {
        metaNoFieldID := arrow.NewMetadata([]string{"PARQUET:field_id"}, 
[]string{"-1"})
 

Reply via email to