Skip to content

Commit 63e8a1a

Browse files
committed
add unshred and shredded reader and roundtrip
1 parent 29f2603 commit 63e8a1a

34 files changed

Lines changed: 3655 additions & 1492 deletions

‎cpp/src/arrow/extension/parquet_variant.cc‎

Lines changed: 80 additions & 57 deletions
Original file line numberDiff line numberDiff line change
@@ -29,44 +29,6 @@
2929

3030
namespace arrow::extension {
3131

32-
VariantExtensionType::VariantExtensionType(const std::shared_ptr<DataType>& storage_type)
33-
: ExtensionType(storage_type) {
34-
// IsSupportedStorageType should have been called already, asserting that
35-
// metadata is present and at least one of value / typed_value is present.
36-
for (const auto& field : storage_type->fields()) {
37-
if (field->name() == "metadata") {
38-
metadata_ = field;
39-
} else if (field->name() == "value") {
40-
value_ = field;
41-
} else if (field->name() == "typed_value") {
42-
typed_value_ = field;
43-
}
44-
}
45-
}
46-
47-
bool VariantExtensionType::ExtensionEquals(const ExtensionType& other) const {
48-
return other.extension_name() == this->extension_name() &&
49-
other.storage_type()->Equals(this->storage_type());
50-
}
51-
52-
Result<std::shared_ptr<DataType>> VariantExtensionType::Deserialize(
53-
std::shared_ptr<DataType> storage_type, const std::string& serialized) const {
54-
if (!serialized.empty()) {
55-
return Status::Invalid("Unexpected serialized metadata: '", serialized, "'");
56-
}
57-
return VariantExtensionType::Make(std::move(storage_type));
58-
}
59-
60-
std::string VariantExtensionType::Serialize() const { return ""; }
61-
62-
std::shared_ptr<Array> VariantExtensionType::MakeArray(
63-
std::shared_ptr<ArrayData> data) const {
64-
DCHECK_EQ(data->type->id(), Type::EXTENSION);
65-
DCHECK_EQ(kVariantExtensionName,
66-
internal::checked_cast<const ExtensionType&>(*data->type).extension_name());
67-
return std::make_shared<VariantArray>(data);
68-
}
69-
7032
namespace {
7133

7234
bool IsSupportedPrimitiveTypedValue(const std::shared_ptr<DataType>& type) {
@@ -109,8 +71,10 @@ bool IsSupportedPrimitiveTypedValue(const std::shared_ptr<DataType>& type) {
10971
}
11072
}
11173

74+
template <bool strict>
11275
bool IsSupportedTypedValue(const std::shared_ptr<Field>& field);
11376

77+
template <bool strict, bool in_object>
11478
bool IsVariantFieldGroup(const std::shared_ptr<DataType>& type) {
11579
if (type->id() != Type::STRUCT) {
11680
return false;
@@ -126,44 +90,51 @@ bool IsVariantFieldGroup(const std::shared_ptr<DataType>& type) {
12690
}
12791
value = field;
12892
} else if (field->name() == "typed_value") {
129-
if (typed_value != nullptr || !IsSupportedTypedValue(field)) {
93+
if (typed_value != nullptr || !IsSupportedTypedValue<strict>(field)) {
13094
return false;
13195
}
13296
typed_value = field;
13397
} else {
13498
return false;
13599
}
136100
}
137-
return value != nullptr || typed_value != nullptr;
101+
102+
if constexpr (strict && in_object) {
103+
return value != nullptr;
104+
} else {
105+
return value != nullptr || typed_value != nullptr;
106+
}
138107
}
139108

109+
template <bool strict>
140110
bool IsSupportedTypedValue(const std::shared_ptr<Field>& field) {
141111
if (!field->nullable()) {
142112
return false;
143113
}
144-
auto is_variant_field_group = [](const auto& field) {
145-
return !field->nullable() && IsVariantFieldGroup(field->type());
146-
};
147114

148115
switch (field->type()->id()) {
149116
case Type::STRUCT:
150117
return field->type()->num_fields() > 0 &&
151-
std::ranges::all_of(field->type()->fields(), is_variant_field_group);
118+
std::ranges::all_of(field->type()->fields(), [](const auto& field) {
119+
return (!strict || !field->nullable()) &&
120+
IsVariantFieldGroup<strict, /*in_object=*/true>(field->type());
121+
});
152122
case Type::LIST:
153123
case Type::LARGE_LIST:
154124
case Type::LIST_VIEW:
155125
case Type::LARGE_LIST_VIEW:
156-
case Type::FIXED_SIZE_LIST:
157-
return is_variant_field_group(field->type()->field(0));
126+
case Type::FIXED_SIZE_LIST: {
127+
auto& inner_field = field->type()->field(0);
128+
return !inner_field->nullable() &&
129+
IsVariantFieldGroup<strict, /*in_object=*/false>(inner_field->type());
130+
}
158131
default:
159132
return IsSupportedPrimitiveTypedValue(field->type());
160133
}
161134
}
162135

163-
} // namespace
164-
165-
bool VariantExtensionType::IsSupportedStorageType(
166-
const std::shared_ptr<DataType>& storage_type) {
136+
template <bool strict>
137+
bool IsSupportedStorageTypeImpl(const std::shared_ptr<DataType>& storage_type) {
167138
if (storage_type->id() != Type::STRUCT) {
168139
return false;
169140
}
@@ -185,7 +156,7 @@ bool VariantExtensionType::IsSupportedStorageType(
185156
}
186157
value = field;
187158
} else if (field->name() == "typed_value") {
188-
if (typed_value != nullptr || !IsSupportedTypedValue(field)) {
159+
if (typed_value != nullptr || !IsSupportedTypedValue<strict>(field)) {
189160
return false;
190161
}
191162
typed_value = field;
@@ -194,16 +165,68 @@ bool VariantExtensionType::IsSupportedStorageType(
194165
}
195166
}
196167

197-
if (metadata == nullptr || (value == nullptr && typed_value == nullptr)) {
168+
if (metadata == nullptr) {
198169
return false;
199170
}
200-
if (value == nullptr) {
201-
return true;
171+
172+
if constexpr (strict) {
173+
if (value == nullptr) {
174+
return false;
175+
}
176+
return (typed_value == nullptr) != value->nullable();
177+
} else {
178+
bool value_nullable = (value == nullptr) || value->nullable();
179+
return (typed_value == nullptr) != value_nullable;
202180
}
203-
if (typed_value == nullptr) {
204-
return !value->nullable();
181+
}
182+
183+
} // namespace
184+
185+
VariantExtensionType::VariantExtensionType(const std::shared_ptr<DataType>& storage_type)
186+
: ExtensionType(storage_type) {
187+
// IsSupportedStorageType should have been called already, asserting that
188+
// metadata is present and at least one of value / typed_value is present.
189+
for (const auto& field : storage_type->fields()) {
190+
if (field->name() == "metadata") {
191+
metadata_ = field;
192+
} else if (field->name() == "value") {
193+
value_ = field;
194+
} else if (field->name() == "typed_value") {
195+
typed_value_ = field;
196+
}
197+
}
198+
}
199+
200+
bool VariantExtensionType::ExtensionEquals(const ExtensionType& other) const {
201+
return other.extension_name() == this->extension_name() &&
202+
other.storage_type()->Equals(this->storage_type());
203+
}
204+
205+
Result<std::shared_ptr<DataType>> VariantExtensionType::Deserialize(
206+
std::shared_ptr<DataType> storage_type, const std::string& serialized) const {
207+
if (!serialized.empty()) {
208+
return Status::Invalid("Unexpected serialized metadata: '", serialized, "'");
209+
}
210+
if (!IsSupportedStorageTypeImpl</*strict=*/false>(storage_type)) {
211+
return Status::Invalid("Invalid storage type for VariantExtensionType: ",
212+
storage_type->ToString());
205213
}
206-
return value->nullable();
214+
return std::make_shared<VariantExtensionType>(std::move(storage_type));
215+
}
216+
217+
std::string VariantExtensionType::Serialize() const { return ""; }
218+
219+
std::shared_ptr<Array> VariantExtensionType::MakeArray(
220+
std::shared_ptr<ArrayData> data) const {
221+
DCHECK_EQ(data->type->id(), Type::EXTENSION);
222+
DCHECK_EQ(kVariantExtensionName,
223+
internal::checked_cast<const ExtensionType&>(*data->type).extension_name());
224+
return std::make_shared<VariantArray>(data);
225+
}
226+
227+
bool VariantExtensionType::IsSupportedStorageType(
228+
const std::shared_ptr<DataType>& storage_type) {
229+
return IsSupportedStorageTypeImpl</*strict=*/true>(storage_type);
207230
}
208231

209232
Result<std::shared_ptr<DataType>> VariantExtensionType::Make(

‎cpp/src/parquet/CMakeLists.txt‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -192,7 +192,10 @@ set(PARQUET_SRCS
192192
stream_writer.cc
193193
types.cc
194194
variant/builder.cc
195-
variant/encoding.cc
195+
variant/decoding.cc
196+
variant/array_internal.cc
197+
variant/slot_internal.cc
198+
variant/unshred.cc
196199
variant/validate.cc
197200
xxhasher.cc)
198201

‎cpp/src/parquet/arrow/arrow_schema_test.cc‎

Lines changed: 39 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1092,8 +1092,35 @@ TEST_F(TestConvertParquetSchema, ParquetVariantShreddedWithoutValue) {
10921092
::arrow::struct_({::arrow::field("metadata", ::arrow::binary(), false),
10931093
::arrow::field("typed_value", ::arrow::int64())});
10941094

1095-
ASSERT_NO_FATAL_FAILURE(CheckParquetVariantSchema(
1096-
"variant_shredded", {metadata, typed_value}, storage_type));
1095+
{
1096+
auto arrow_schema = ::arrow::schema({::arrow::field("variant_shredded", storage_type)});
1097+
ASSERT_OK(ConvertSchema({GroupNode::Make("variant_shredded", Repetition::OPTIONAL,
1098+
{metadata, typed_value},
1099+
LogicalType::Variant())}));
1100+
ASSERT_NO_FATAL_FAILURE(CheckFlatSchema(arrow_schema, /*check_metadata=*/true));
1101+
}
1102+
1103+
auto registered_storage_type =
1104+
::arrow::struct_({::arrow::field("metadata", ::arrow::binary(), false),
1105+
::arrow::field("value", ::arrow::binary(), false)});
1106+
auto registered_variant = ::arrow::extension::variant(registered_storage_type);
1107+
::arrow::ExtensionTypeGuard guard({registered_variant});
1108+
1109+
ArrowReaderProperties props;
1110+
props.set_arrow_extensions_enabled(true);
1111+
ASSERT_OK(ConvertSchema({GroupNode::Make("variant_shredded", Repetition::OPTIONAL,
1112+
{metadata, typed_value},
1113+
LogicalType::Variant())},
1114+
/*metadata=*/nullptr, props));
1115+
1116+
auto field = result_schema_->field(0);
1117+
ASSERT_EQ(::arrow::Type::EXTENSION, field->type()->id());
1118+
auto variant_type =
1119+
std::dynamic_pointer_cast<::arrow::extension::VariantExtensionType>(field->type());
1120+
ASSERT_NE(nullptr, variant_type);
1121+
ASSERT_EQ(nullptr, variant_type->value());
1122+
ASSERT_NE(nullptr, variant_type->typed_value());
1123+
ASSERT_TRUE(variant_type->storage_type()->Equals(storage_type));
10971124
}
10981125

10991126
TEST_F(TestConvertParquetSchema, ParquetSchemaArrowJsonExtension) {
@@ -1594,7 +1621,16 @@ TEST_F(TestConvertArrowSchema, ParquetVariantShreddedWithoutValue) {
15941621
auto storage_type =
15951622
::arrow::struct_({::arrow::field("metadata", ::arrow::binary(), false),
15961623
::arrow::field("typed_value", ::arrow::int64())});
1597-
auto variant_type = ::arrow::extension::variant(storage_type);
1624+
auto registered_storage_type =
1625+
::arrow::struct_({::arrow::field("metadata", ::arrow::binary(), false),
1626+
::arrow::field("value", ::arrow::binary(), false)});
1627+
ASSERT_OK_AND_ASSIGN(auto registered_variant_type,
1628+
::arrow::extension::VariantExtensionType::Make(
1629+
registered_storage_type));
1630+
auto registered_variant =
1631+
::arrow::internal::checked_pointer_cast<::arrow::extension::VariantExtensionType>(
1632+
registered_variant_type);
1633+
ASSERT_OK_AND_ASSIGN(auto variant_type, registered_variant->Deserialize(storage_type, ""));
15981634

15991635
ASSERT_RAISES(Invalid, ConvertSchema({::arrow::field("variant", variant_type)}));
16001636
}

‎cpp/src/parquet/arrow/reader.cc‎

Lines changed: 29 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727

2828
#include "arrow/array.h"
2929
#include "arrow/buffer.h"
30+
#include "arrow/extension/parquet_variant.h"
3031
#include "arrow/extension_type.h"
3132
#include "arrow/io/memory.h"
3233
#include "arrow/memory_pool.h"
@@ -54,6 +55,7 @@
5455
#include "parquet/page_index.h"
5556
#include "parquet/properties.h"
5657
#include "parquet/schema.h"
58+
#include "parquet/variant/validate.h"
5759

5860
using arrow::Array;
5961
using arrow::ArrayData;
@@ -583,6 +585,26 @@ class ExtensionReader : public ColumnReaderImpl {
583585
std::unique_ptr<ColumnReaderImpl> storage_reader_;
584586
};
585587

588+
class VariantReader : public ExtensionReader {
589+
public:
590+
VariantReader(std::shared_ptr<ReaderContext> ctx, std::shared_ptr<Field> field,
591+
std::unique_ptr<ColumnReaderImpl> storage_reader)
592+
: ExtensionReader(std::move(field), std::move(storage_reader)),
593+
ctx_(std::move(ctx)) {}
594+
595+
Status BuildArray(int64_t length_upper_bound,
596+
std::shared_ptr<ChunkedArray>* out) override {
597+
RETURN_NOT_OK(ExtensionReader::BuildArray(length_upper_bound, out));
598+
if (ctx_->reader_properties->variant_validation_enabled()) {
599+
PARQUET_CATCH_NOT_OK(variant::ValidateVariants<false>(**out, ctx_->pool));
600+
}
601+
return Status::OK();
602+
}
603+
604+
private:
605+
std::shared_ptr<ReaderContext> ctx_;
606+
};
607+
586608
template <typename IndexType>
587609
class ListReader : public ColumnReaderImpl {
588610
public:
@@ -892,8 +914,8 @@ Status GetReader(const SchemaField& field, const std::shared_ptr<Field>& arrow_f
892914
auto type_id = arrow_field->type()->id();
893915

894916
if (type_id == ::arrow::Type::EXTENSION) {
895-
auto storage_field = arrow_field->WithType(
896-
checked_cast<const ExtensionType&>(*arrow_field->type()).storage_type());
917+
const auto& ext_type = checked_cast<const ExtensionType&>(*arrow_field->type());
918+
auto storage_field = arrow_field->WithType(ext_type.storage_type());
897919
RETURN_NOT_OK(GetReader(field, storage_field, ctx, out));
898920
if (*out) {
899921
auto storage_type = (*out)->field()->type();
@@ -902,7 +924,11 @@ Status GetReader(const SchemaField& field, const std::shared_ptr<Field>& arrow_f
902924
"Due to column pruning only part of an extension's storage type was loaded. "
903925
"An extension type cannot be created without all of its fields");
904926
}
905-
*out = std::make_unique<ExtensionReader>(arrow_field, std::move(*out));
927+
if (ext_type.extension_name() == ::arrow::extension::kVariantExtensionName) {
928+
*out = std::make_unique<VariantReader>(ctx, arrow_field, std::move(*out));
929+
} else {
930+
*out = std::make_unique<ExtensionReader>(arrow_field, std::move(*out));
931+
}
906932
}
907933
return Status::OK();
908934
}

0 commit comments

Comments
 (0)