Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions arrow/extensions/extensions.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 }
19 changes: 18 additions & 1 deletion arrow/extensions/variant.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
70 changes: 69 additions & 1 deletion arrow/extensions/variant_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
package extensions_test

import (
"bytes"
"encoding/json"
"fmt"
"testing"
Expand All @@ -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"
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion parquet/pqarrow/schema.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

arrow.GetExtensionType("parquet.variant") now returns *extensions.legacyVariantType. The new name predicate accepts that type, but this unconditional assertion to *extensions.VariantType panics. ToParquet should return an error for unsupported wrappers or normalize compatible legacy storage through NewVariantType, rather than asserting based only on the extension name. Please add a regression test using the registry-returned legacy type.

}
}
Expand Down
2 changes: 1 addition & 1 deletion parquet/pqarrow/schema_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading