2929
3030namespace 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-
7032namespace {
7133
7234bool 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>
11275bool IsSupportedTypedValue (const std::shared_ptr<Field>& field);
11376
77+ template <bool strict, bool in_object>
11478bool 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>
140110bool 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
209232Result<std::shared_ptr<DataType>> VariantExtensionType::Make (
0 commit comments