From 06fb240edd6660e4f9d862f7164303ac088355d8 Mon Sep 17 00:00:00 2001 From: Digvijay Date: Fri, 28 Aug 2026 00:11:24 -0500 Subject: [PATCH] fix(arrow/extensions): use canonical Variant extension name Arrow requires arrow.parquet.variant. Keep parquet.variant registered so older IPC still deserializes. Fixes #1203 Signed-off-by: Digvijay --- arrow/extensions/extensions.go | 15 +++++++ arrow/extensions/variant.go | 19 ++++++++- arrow/extensions/variant_test.go | 70 +++++++++++++++++++++++++++++++- parquet/pqarrow/schema.go | 2 +- parquet/pqarrow/schema_test.go | 2 +- 5 files changed, 104 insertions(+), 4 deletions(-) diff --git a/arrow/extensions/extensions.go b/arrow/extensions/extensions.go index 6f13aa646..7c67e263a 100644 --- a/arrow/extensions/extensions.go +++ b/arrow/extensions/extensions.go @@ -35,4 +35,19 @@ func init() { panic(err) } } + + // arrow-go originally registered Variant as parquet.variant. Keep that + // name in the registry so older IPC still deserializes. + if err := arrow.RegisterExtensionType(&legacyVariantType{}); err != nil { + panic(err) + } } + +// legacyVariantType exists only so GetExtensionType("parquet.variant") can +// still reconstruct a VariantType. In-memory and newly written IPC use +// VariantExtensionName. +type legacyVariantType struct { + VariantType +} + +func (*legacyVariantType) ExtensionName() string { return LegacyVariantExtensionName } diff --git a/arrow/extensions/variant.go b/arrow/extensions/variant.go index fee2e046a..0dc0a506b 100644 --- a/arrow/extensions/variant.go +++ b/arrow/extensions/variant.go @@ -262,7 +262,24 @@ 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()) diff --git a/arrow/extensions/variant_test.go b/arrow/extensions/variant_test.go index a39fd5131..26af352f1 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", variant1.String()) + assert.Equal(t, "arrow.parquet.variant", variant1.ExtensionName()) + assert.Equal(t, "extension", 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 @@ -92,6 +99,67 @@ 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()) + + rdr, err := ipc.NewReader(&buf) + require.NoError(t, err) + defer rdr.Release() + + written, err := rdr.Read() + require.NoError(t, err) + defer written.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 TestVariantExtensionBadNestedTypes(t *testing.T) { tests := []struct { name string diff --git a/parquet/pqarrow/schema.go b/parquet/pqarrow/schema.go index d2a975593..a6207d672 100644 --- a/parquet/pqarrow/schema.go +++ b/parquet/pqarrow/schema.go @@ -353,7 +353,7 @@ 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" { + if extensions.IsVariantExtensionName(extType.ExtensionName()) { return variantToNode(extType.(*extensions.VariantType), field, props, arrprops) } } diff --git a/parquet/pqarrow/schema_test.go b/parquet/pqarrow/schema_test.go index db72c22a0..d22a502d7 100644 --- a/parquet/pqarrow/schema_test.go +++ b/parquet/pqarrow/schema_test.go @@ -1224,7 +1224,7 @@ 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)