This is an automated email from the ASF dual-hosted git repository. yiguolei pushed a commit to branch branch-4.2 in repository https://gitbox.apache.org/repos/asf/doris.git
commit e48911e417ef0c6148c19006f62f8a6ab07aa427 Author: Gabriel <[email protected]> AuthorDate: Fri Oct 9 14:30:19 2026 +0800 [feature](flight) Support native Variant V2 results for ADBC on branch-4.1 (#68667) ### What problem does this PR solve? Flight SQL exposes VARIANT as UTF8, preventing ADBC clients from receiving native binary Variant V2 values. On branch-4.1, return Variant V2 through the `arrow.parquet.variant` extension by default, with `struct<metadata: binary not null, value: binary not null>` storage. SQL NULL is a null struct; Variant null is an encoded non-null value. Clients without a Variant extension implementation can access the physical struct and field metadata. Support **Variant V2 only**. Legacy Variant returns an unsupported error, including nested legacy fields, constants, all-null columns and empty results. Reject legacy types during BE Flight schema conversion and defensively in the legacy SerDe. Remove the output-format session variable and result-sink flag entirely. Users who need text can explicitly `CAST(... AS STRING)`. Keep execution, GetTables, GetSchema and prepared-statement schemas consistent, including nested Variant fields. Native Variant output does not depend on a heartbeat capability or a cluster-wide BE support check. Retain the Variant V2 type requirement and reject non-native execution schemas. Reuse the V2 binary representation and compact dictionaries per selected row, preserving physical scalar types and decimal scales without repeatedly copying unrelated dictionary keys. Keep the value helper private to `data_type_variant_v2_serde.cpp`. Native encoding accepts up to 128 nested levels. Iceberg output keeps its existing behavior. Document the V2 requirement and client decoding limitations in the Python Flight sample README; parsing JSON does not bypass the configured Variant representation. ### Release note Arrow Flight SQL / ADBC returns Variant V2 as native Arrow Variant binary by default, without a session switch. Legacy Variant is unsupported. Use an explicit SQL cast to STRING for text output. ### Validation - The native-output update passed all 46 selected Arrow/Flight tests using a focused ASAN binary; the binary conversion code is unchanged by the capability removal. - For the capability removal, regenerated the heartbeat Thrift bindings, compiled the affected FE sources and BE heartbeat translation unit, and passed all 6 focused FE tests covering native schemas, legacy/UTF8 rejection and the existing Paimon heartbeat capability. - FE Checkstyle, clang-format 16 and diff checks passed. - The updated Groovy regression script and Python sample syntax were checked in the preceding update. - A full production build and live cluster regression/ADBC execution were not rerun for this update; CI validation is requested. ### Check List (For Author) - Test - [x] Regression test updated - [x] Unit Test - Behavior changed: - [x] Yes. Variant V2 uses native binary by default; legacy Variant requires an explicit text cast. - Does this need documentation? - [x] Yes, usage, V2-only scope and client decoding limitations are documented in the Python Flight sample README. --- .../data_type_serde/data_type_variant_serde.cpp | 6 + .../data_type_serde/data_type_variant_v2_serde.cpp | 78 +++- be/src/format/arrow/arrow_block_convertor.cpp | 22 + be/src/format/arrow/arrow_block_convertor.h | 7 + be/src/format/arrow/arrow_row_batch.cpp | 57 ++- be/src/format/arrow/arrow_row_batch.h | 14 +- .../arrow_flight/arrow_flight_batch_reader.cpp | 1 + be/test/format/arrow/arrow_flight_variant_test.cpp | 511 +++++++++++++++++++++ be/test/format/arrow/arrow_row_batch_test.cpp | 11 +- .../java/org/apache/doris/planner/ResultSink.java | 2 +- .../service/arrowflight/FlightSqlQuerySchema.java | 16 +- .../service/arrowflight/FlightSqlSchemaHelper.java | 39 +- .../FlightSqlConnectProcessorSchemaTest.java | 8 +- .../arrowflight/FlightSqlNativeVariantTest.java | 124 +++++ fe/pom.xml | 1 - .../test_flight_cancel_cleanup.groovy | 4 +- .../test_flight_native_variant.groovy | 215 +++++++++ .../flink_connector_p0/flink_connector_type.groovy | 29 +- samples/arrow-flight-sql/python/README.md | 39 +- samples/arrow-flight-sql/python/test_variant.py | 96 ++++ 20 files changed, 1238 insertions(+), 42 deletions(-) diff --git a/be/src/core/data_type_serde/data_type_variant_serde.cpp b/be/src/core/data_type_serde/data_type_variant_serde.cpp index 3bd9be5797f..e0ddd20b7d6 100644 --- a/be/src/core/data_type_serde/data_type_variant_serde.cpp +++ b/be/src/core/data_type_serde/data_type_variant_serde.cpp @@ -157,6 +157,12 @@ Status DataTypeVariantSerDe::write_column_to_arrow(const IColumn& column, const int64_t start, int64_t end, const cctz::time_zone& ctz) const { const auto* var = check_and_get_column<ColumnVariant>(column); + if (array_builder->type()->id() == arrow::Type::STRUCT) { + // Native Flight output must not reinterpret legacy storage as Variant V2. + return Status::NotSupported( + "Native Arrow Flight output only supports Variant V2, not legacy Variant; " + "cast the result to STRING for text output"); + } if (array_builder->type()->id() == arrow::Type::LARGE_STRING) { auto& builder = assert_cast<arrow::LargeStringBuilder&>(*array_builder); return write_variant_column_to_arrow_impl(column, *var, null_map, builder, start, end, ctz); diff --git a/be/src/core/data_type_serde/data_type_variant_v2_serde.cpp b/be/src/core/data_type_serde/data_type_variant_v2_serde.cpp index 0e0d7a58d2f..4e9ae6f872d 100644 --- a/be/src/core/data_type_serde/data_type_variant_v2_serde.cpp +++ b/be/src/core/data_type_serde/data_type_variant_v2_serde.cpp @@ -24,6 +24,7 @@ #include <algorithm> #include <cstring> #include <limits> +#include <optional> #include <orc/Vector.hh> #include <span> #include <utility> @@ -53,6 +54,48 @@ namespace doris { namespace { +// ColumnVariantV2 already validates its encoded dictionaries and value structure. Only walk +// the selected row here: validating every unused dictionary key per row is quadratic for +// shared dictionaries, including when a nested ARRAY invokes this writer on one-row slices. +Status append_flight_variant_value(VariantRef value, VariantBatchBuilder::Row& output, + size_t depth = 0) { + const auto basic_type = value.basic_type(); + if (depth > VARIANT_MAX_NESTING_DEPTH) { + return Status::NotSupported( + "Native Arrow Variant nesting exceeds {}; " + "cast the result to STRING for text output", + VARIANT_MAX_NESTING_DEPTH); + } + if (value.value_size() != value.value.size) { + throw Exception(ErrorCode::CORRUPTION, + "Native Arrow Variant contains trailing value bytes"); + } + if (basic_type == VariantBasicType::OBJECT) { + auto object = output.start_object(); + auto fields = value.object_view(); + for (uint32_t i = 0; i < fields.size(); ++i) { + uint32_t field_id; + auto child = fields.value_at(i, &field_id); + object.add_key(value.metadata.key_at(field_id)); + RETURN_IF_ERROR(append_flight_variant_value(child, output, depth + 1)); + } + object.finish(); + } else if (basic_type == VariantBasicType::ARRAY) { + auto array = output.start_array(); + for (uint32_t i = 0; i < value.num_elements(); ++i) { + RETURN_IF_ERROR(append_flight_variant_value(value.array_at(i), output, depth + 1)); + } + array.finish(); + } else { + // Primitives never reference dictionary keys. Reuse physical import to retain widths, + // decimal scales and non-JSON types; canonical equality encoding normalizes those away. + static constexpr char empty_metadata[] = {0x11, 0, 0}; + value.metadata = {empty_metadata, sizeof(empty_metadata)}; + output.add_value(value); + } + return Status::OK(); +} + using MetaIdsColumn = ColumnVector<TYPE_UINT32>; std::span<const NullMap::value_type> forced_nulls(const NullMap* null_map) { @@ -625,13 +668,14 @@ Status write_paimon_variant(const IColumn& column, const NullMap* null_map, return status; } -Status write_iceberg_variant(const IColumn& column, const NullMap* null_map, - arrow::ArrayBuilder* array_builder, int64_t start, int64_t end) { +Status write_parquet_variant_arrow(const IColumn& column, const NullMap* null_map, + arrow::ArrayBuilder* array_builder, int64_t start, int64_t end, + bool compact_metadata) { if (start < 0 || end < start) { - return Status::InvalidArgument("Invalid Iceberg Variant row range [{}, {})", start, end); + return Status::InvalidArgument("Invalid Variant Arrow row range [{}, {})", start, end); } if (array_builder->type()->id() != arrow::Type::STRUCT) { - return Status::InvalidArgument("Iceberg Variant writer requires a struct builder, got {}", + return Status::InvalidArgument("Variant Arrow writer requires a struct builder, got {}", array_builder->type()->ToString()); } auto& builder = assert_cast<arrow::StructBuilder&>(*array_builder); @@ -640,7 +684,7 @@ Status write_iceberg_variant(const IColumn& column, const NullMap* null_map, type->field(1)->name() != "value" || type->field(0)->type()->id() != arrow::Type::BINARY || type->field(1)->type()->id() != arrow::Type::BINARY) { return Status::InvalidArgument( - "Iceberg Variant writer requires struct<metadata: binary, value: binary>, got {}", + "Variant Arrow writer requires struct<metadata: binary, value: binary>, got {}", type->ToString()); } auto& metadata_builder = assert_cast<arrow::BinaryBuilder&>(*builder.field_builder(0)); @@ -657,10 +701,27 @@ Status write_iceberg_variant(const IColumn& column, const NullMap* null_map, if (!status.ok()) { return; } + std::optional<VariantBatchBuilder> compacted; + if (compact_metadata) { + const auto keys = value.metadata.dict_size(); + // An empty dictionary or an object using every key already has row-local metadata. + if (keys != 0 && (value.basic_type() != VariantBasicType::OBJECT || + value.num_elements() != keys)) { + VariantBatchBuilder encoder; + auto row = encoder.begin_row(); + status = append_flight_variant_value(value, row); + if (!status.ok()) { + return; + } + row.finish(); + compacted.emplace(encoder.finish_batch()); + value = compacted->value_at(0); + } + } if (value.metadata.size > std::numeric_limits<int32_t>::max() || value.value.size > std::numeric_limits<int32_t>::max()) { status = Status::InvalidArgument( - "Iceberg Variant metadata/value exceeds Arrow binary size limit"); + "Variant Arrow metadata/value exceeds Arrow binary size limit"); return; } status = checkArrowStatus(builder.Append(), column, builder); @@ -733,6 +794,9 @@ Status DataTypeVariantV2SerDe::write_column_to_arrow(const IColumn& column, cons options.timezone = &ctz; const size_t first = checked_row(start); const size_t last = checked_row(end); + if (array_builder->type()->id() == arrow::Type::STRUCT) { + return write_parquet_variant_arrow(column, null_map, array_builder, start, end, true); + } if (array_builder->type()->id() == arrow::Type::STRING) { return write_arrow(column, null_map, assert_cast<arrow::StringBuilder&>(*array_builder), first, last, options); @@ -759,7 +823,7 @@ Status DataTypeVariantV2SerDe::write_column_to_iceberg_arrow( const std::shared_ptr<const IDataType>&, const IColumn& column, const NullMap* null_map, const std::shared_ptr<arrow::Field>&, arrow::ArrayBuilder* array_builder, int64_t start, int64_t end, const cctz::time_zone&) const { - return write_iceberg_variant(column, null_map, array_builder, start, end); + return write_parquet_variant_arrow(column, null_map, array_builder, start, end, false); } Status DataTypeVariantV2SerDe::write_column_to_orc(const std::string&, const IColumn& column, diff --git a/be/src/format/arrow/arrow_block_convertor.cpp b/be/src/format/arrow/arrow_block_convertor.cpp index 2148d0934af..685637fe737 100644 --- a/be/src/format/arrow/arrow_block_convertor.cpp +++ b/be/src/format/arrow/arrow_block_convertor.cpp @@ -473,6 +473,28 @@ Status ArrowBlockConvertor::init() { return Status::OK(); } +Status ArrowFlightArrowBlockConvertor::write_column(const std::shared_ptr<const IDataType>& type, + const DataTypeSerDe& serde, + const IColumn& column, const NullMap* null_map, + const std::shared_ptr<arrow::Field>& field, + arrow::ArrayBuilder* array_builder, + int64_t start, int64_t end, + const cctz::time_zone& ctz) const { + if (contains_extension_type(field->type())) { + std::shared_ptr<arrow::DataType> native_type; + RETURN_IF_ERROR( + ArrowFlightSchemaConvertor(ctz.name()).convert_to_arrow_type(type, &native_type)); + // Check the extension identity and its complete nested shape before allowing the + // Variant SerDe to write binary storage. An arbitrary STRUCT is not a Variant binding. + // Timestamp labels may differ for equivalent fixed offsets, including inside containers. + if (is_declared_plain_arrow_binding(type, native_type, field->type())) { + return serde.write_column_to_arrow(column, null_map, array_builder, start, end, ctz); + } + } + return DorisArrowBlockConvertor::write_column(type, serde, column, null_map, field, + array_builder, start, end, ctz); +} + Status ArrowFlightArrowBlockConvertor::convert_to_arrow(const Block& block, arrow::MemoryPool* pool, std::shared_ptr<arrow::RecordBatch>* result, size_t start_row, size_t end_row) const { diff --git a/be/src/format/arrow/arrow_block_convertor.h b/be/src/format/arrow/arrow_block_convertor.h index e0b58d76838..c18c28df72a 100644 --- a/be/src/format/arrow/arrow_block_convertor.h +++ b/be/src/format/arrow/arrow_block_convertor.h @@ -116,6 +116,13 @@ public: Status convert_to_arrow(const Block& block, arrow::MemoryPool* pool, std::shared_ptr<arrow::RecordBatch>* result, size_t start_row = 0, size_t end_row = 0) const override; + +protected: + Status write_column(const std::shared_ptr<const IDataType>& type, const DataTypeSerDe& serde, + const IColumn& column, const NullMap* null_map, + const std::shared_ptr<arrow::Field>& field, + arrow::ArrayBuilder* array_builder, int64_t start, int64_t end, + const cctz::time_zone& ctz) const override; }; class PythonArrowBlockConvertor final : public DorisArrowBlockConvertor { diff --git a/be/src/format/arrow/arrow_row_batch.cpp b/be/src/format/arrow/arrow_row_batch.cpp index a958843e519..9589f6e59ac 100644 --- a/be/src/format/arrow/arrow_row_batch.cpp +++ b/be/src/format/arrow/arrow_row_batch.cpp @@ -18,6 +18,7 @@ #include "format/arrow/arrow_row_batch.h" #include <arrow/buffer.h> +#include <arrow/extension/parquet_variant.h> #include <arrow/io/memory.h> #include <arrow/ipc/writer.h> #include <arrow/record_batch.h> @@ -40,14 +41,30 @@ #include "core/data_type/data_type_array.h" #include "core/data_type/data_type_map.h" #include "core/data_type/data_type_struct.h" +#include "core/data_type/data_type_variant_v2.h" #include "core/data_type/define_primitive_type.h" #include "exprs/vexpr.h" #include "exprs/vexpr_context.h" #include "format/arrow/arrow_block_convertor.h" +#include "format/arrow/arrow_utils.h" #include "runtime/descriptors.h" namespace doris { +Status register_arrow_variant_extension() { + // Remote Flight readers must restore the extension before decoding the result schema. + static const auto status = [] { + if (arrow::GetExtensionType("arrow.parquet.variant") != nullptr) { + return arrow::Status::OK(); + } + return arrow::RegisterExtensionType( + std::static_pointer_cast<arrow::ExtensionType>(arrow::extension::variant( + arrow::struct_({arrow::field("metadata", arrow::binary(), false), + arrow::field("value", arrow::binary(), false)})))); + }(); + return status.ok() ? Status::OK() : Status::InternalError(status.ToString()); +} + Status DorisArrowSchemaConvertor::convert_to_arrow_type( const DataTypePtr& origin_type, std::shared_ptr<arrow::DataType>* result) const { auto type = get_serialized_type(origin_type); @@ -157,10 +174,9 @@ Status DorisArrowSchemaConvertor::convert_to_arrow_type( *result = std::make_shared<arrow::StructType>(fields); break; } - case TYPE_VARIANT: { + case TYPE_VARIANT: *result = arrow::utf8(); break; - } case TYPE_QUANTILE_STATE: case TYPE_BITMAP: case TYPE_HLL: { @@ -229,6 +245,25 @@ std::string DorisArrowSchemaConvertor::timestamp_timezone(PrimitiveType) const { return _timezone == "Z" ? "UTC" : _timezone; } +Status ArrowFlightSchemaConvertor::convert_to_arrow_type( + const DataTypePtr& type, std::shared_ptr<arrow::DataType>* result) const { + // Flight always uses native Variant, including recursively converted children. + if (type->get_primitive_type() == TYPE_VARIANT) { + // Reject by type before reading rows, including empty results and nested legacy leaves. + if (dynamic_cast<const DataTypeVariantV2*>(remove_nullable(type).get()) == nullptr) { + return Status::NotSupported( + "Native Arrow Flight output only supports Variant V2, not legacy Variant; " + "cast the result to STRING for text output"); + } + RETURN_IF_ERROR(register_arrow_variant_extension()); + *result = arrow::extension::variant( + arrow::struct_({arrow::field("metadata", arrow::binary(), false), + arrow::field("value", arrow::binary(), false)})); + return Status::OK(); + } + return DorisArrowSchemaConvertor::convert_to_arrow_type(type, result); +} + std::string ArrowFlightSchemaConvertor::timestamp_timezone(PrimitiveType type) const { // DATETIMEV2 is wall-clock time; TIMESTAMPTZ must still describe an instant. return type == TYPE_DATETIMEV2 ? "" : DorisArrowSchemaConvertor::timestamp_timezone(type); @@ -323,13 +358,17 @@ Status serialize_record_batch(const arrow::RecordBatch& record_batch, std::strin } Status serialize_arrow_schema(std::shared_ptr<arrow::Schema>* schema, std::string* result) { - auto make_empty_result = arrow::RecordBatch::MakeEmpty(*schema); - if (!make_empty_result.ok()) { - return Status::InternalError("serialize_arrow_schema failed, reason: {}", - make_empty_result.status().ToString()); - } - auto batch = make_empty_result.ValueOrDie(); - return serialize_record_batch(*batch, result); + // Schema RPC readers only consume the IPC schema. Building an empty batch would require + // nested extension builders, which Arrow does not provide for ARRAY/MAP/STRUCT<VARIANT>. + std::shared_ptr<arrow::io::BufferOutputStream> sink; + RETURN_DORIS_STATUS_IF_RESULT_ERROR(sink, arrow::io::BufferOutputStream::Create()); + std::shared_ptr<arrow::ipc::RecordBatchWriter> writer; + RETURN_DORIS_STATUS_IF_RESULT_ERROR(writer, arrow::ipc::MakeStreamWriter(sink.get(), *schema)); + RETURN_DORIS_STATUS_IF_ERROR(writer->Close()); + std::shared_ptr<arrow::Buffer> buffer; + RETURN_DORIS_STATUS_IF_RESULT_ERROR(buffer, sink->Finish()); + *result = buffer->ToString(); + return Status::OK(); } } // namespace doris diff --git a/be/src/format/arrow/arrow_row_batch.h b/be/src/format/arrow/arrow_row_batch.h index 60b192716fe..debe570fe10 100644 --- a/be/src/format/arrow/arrow_row_batch.h +++ b/be/src/format/arrow/arrow_row_batch.h @@ -63,8 +63,8 @@ public: std::shared_ptr<arrow::Schema>* result) const; Status get_arrow_schema_from_expr_ctxs(const VExprContextSPtrs& output_vexpr_ctxs, std::shared_ptr<arrow::Schema>* result) const; - Status convert_to_arrow_type(const DataTypePtr& type, - std::shared_ptr<arrow::DataType>* result) const; + virtual Status convert_to_arrow_type(const DataTypePtr& type, + std::shared_ptr<arrow::DataType>* result) const; protected: virtual std::string timestamp_timezone(PrimitiveType type) const; @@ -82,7 +82,13 @@ private: class ArrowFlightSchemaConvertor : public DorisArrowSchemaConvertor { public: - using DorisArrowSchemaConvertor::DorisArrowSchemaConvertor; + explicit ArrowFlightSchemaConvertor(std::string timezone) + : DorisArrowSchemaConvertor(std::move(timezone)) {} + ArrowFlightSchemaConvertor(const Block& header, std::string timezone) + : DorisArrowSchemaConvertor(header, std::move(timezone)) {} + + Status convert_to_arrow_type(const DataTypePtr& type, + std::shared_ptr<arrow::DataType>* result) const override; protected: std::string timestamp_timezone(PrimitiveType type) const override; @@ -103,6 +109,8 @@ protected: PrimitiveType primitive) const override; }; +Status register_arrow_variant_extension(); + std::shared_ptr<arrow::Field> create_arrow_field_with_metadata( const std::string& field_name, const std::shared_ptr<arrow::DataType>& arrow_type, bool is_nullable, PrimitiveType primitive_type); diff --git a/be/src/service/arrow_flight/arrow_flight_batch_reader.cpp b/be/src/service/arrow_flight/arrow_flight_batch_reader.cpp index 7e11fe9a423..32c6ca53d32 100644 --- a/be/src/service/arrow_flight/arrow_flight_batch_reader.cpp +++ b/be/src/service/arrow_flight/arrow_flight_batch_reader.cpp @@ -298,6 +298,7 @@ arrow::Status ArrowFlightBatchRemoteReader::_fetch_schema() { st = Status::create(callback->response_->status()); ARROW_RETURN_NOT_OK(to_arrow_status(st)); + ARROW_RETURN_NOT_OK(to_arrow_status(register_arrow_variant_extension())); if (callback->response_->has_schema() && !callback->response_->schema().empty()) { auto input = arrow::io::BufferReader::FromString(std::string(callback->response_->schema())); diff --git a/be/test/format/arrow/arrow_flight_variant_test.cpp b/be/test/format/arrow/arrow_flight_variant_test.cpp new file mode 100644 index 00000000000..792a3d2b8fe --- /dev/null +++ b/be/test/format/arrow/arrow_flight_variant_test.cpp @@ -0,0 +1,511 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include <arrow/api.h> +#include <arrow/extension/parquet_variant.h> +#include <arrow/io/api.h> +#include <arrow/ipc/api.h> +#include <gtest/gtest.h> + +#include <cmath> +#include <limits> + +#include "core/column/column_array.h" +#include "core/column/column_const.h" +#include "core/column/column_nullable.h" +#include "core/column/column_struct.h" +#include "core/column/column_variant.h" +#include "core/column/variant_v2/column_variant_v2.h" +#include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_factory.hpp" +#include "core/data_type/data_type_map.h" +#include "core/data_type/data_type_nullable.h" +#include "core/data_type/data_type_number.h" +#include "core/data_type/data_type_string.h" +#include "core/data_type/data_type_struct.h" +#include "core/data_type/data_type_variant.h" +#include "core/data_type/data_type_variant_v2.h" +#include "core/data_type_serde/data_type_serde.h" +#include "exprs/function/parse/variant_string_parse.h" +#include "format/arrow/arrow_block_convertor.h" +#include "format/arrow/arrow_row_batch.h" +#include "util/timezone_utils.h" + +namespace doris { +namespace { + +std::shared_ptr<arrow::DataType> native_variant() { + return arrow::extension::variant( + arrow::struct_({arrow::field("metadata", arrow::binary(), false), + arrow::field("value", arrow::binary(), false)})); +} + +MutableColumnPtr documents(const DataTypePtr& type) { + auto column = type->create_column(); + auto serde = type->get_serde(); + DataTypeSerDe::FormatOptions options; + for (std::string json : {R"({"a":[1,null,"x"]})", "null", "42", R"("text")"}) { + Slice slice(json.data(), json.size()); + EXPECT_TRUE(serde->deserialize_one_cell_from_json(*column, slice, options).ok()); + } + if (auto* legacy = check_and_get_column<ColumnVariant>(*column)) { + legacy->finalize(); + } + return column; +} + +VariantRef value_at(const arrow::Array& array, int row) { + const auto& storage = static_cast<const arrow::StructArray&>( + *static_cast<const arrow::ExtensionArray&>(array).storage()); + auto metadata = static_cast<const arrow::BinaryArray&>(*storage.field(0)).GetView(row); + auto value = static_cast<const arrow::BinaryArray&>(*storage.field(1)).GetView(row); + return {{metadata.data(), metadata.size()}, {value.data(), value.size()}}; +} + +TEST(ArrowFlightVariantTest, NativeSchemaRejectsLegacyIncludingNestedAndEmptyResults) { + auto legacy = std::make_shared<DataTypeVariant>(); + auto nullable = make_nullable(legacy); + DataTypes types {legacy, nullable, std::make_shared<DataTypeArray>(nullable), + std::make_shared<DataTypeMap>(std::make_shared<DataTypeString>(), nullable), + std::make_shared<DataTypeStruct>( + DataTypes {std::make_shared<DataTypeVariantV2>(), nullable}, + Strings {"v2", "legacy"})}; + for (const auto& type : types) { + SCOPED_TRACE(type->get_name()); + Block block {{type->create_column(), type, "v"}}; + std::shared_ptr<arrow::Schema> schema; + auto status = ArrowFlightSchemaConvertor(block, "UTC").get_arrow_schema(&schema); + EXPECT_TRUE(status.is<ErrorCode::NOT_IMPLEMENTED_ERROR>()) << status; + EXPECT_NE(status.to_string().find("only supports Variant V2"), std::string::npos); + EXPECT_NE(status.to_string().find("cast the result to STRING"), std::string::npos); + EXPECT_TRUE(DorisArrowSchemaConvertor(block, "UTC").get_arrow_schema(&schema).ok()); + } +} + +TEST(ArrowFlightVariantTest, LegacyNativeOutputRejectsValuesConstantsAndNulls) { + auto type = std::make_shared<DataTypeVariant>(); + auto nulls = ColumnUInt8::create(); + nulls->get_data().resize_fill(4, 1); + auto constant = type->create_column(); + DataTypeSerDe::FormatOptions options; + std::string json = "42"; + Slice slice(json.data(), json.size()); + ASSERT_TRUE(type->get_serde()->deserialize_one_cell_from_json(*constant, slice, options).ok()); + assert_cast<ColumnVariant&>(*constant).finalize(); + ColumnsWithTypeAndName inputs { + {documents(type), type, "v"}, + {ColumnConst::create(std::move(constant), 3), type, "v"}, + {ColumnNullable::create(documents(type), std::move(nulls)), make_nullable(type), "v"}, + {type->create_column(), type, "v"}}; + for (const auto& input : inputs) { + Block block {input}; + ArrowFlightArrowBlockConvertor native(arrow::schema({arrow::field("v", native_variant())}), + cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = native.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_TRUE(status.is<ErrorCode::NOT_IMPLEMENTED_ERROR>()) << status; + EXPECT_NE(status.to_string().find("only supports Variant V2"), std::string::npos); + + DorisArrowBlockConvertor utf8(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(utf8.init().ok()); + ASSERT_TRUE(utf8.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + EXPECT_EQ(batch->num_rows(), block.rows()); + EXPECT_EQ(batch->column(0)->type_id(), arrow::Type::STRING); + if (input.type->is_nullable()) { + EXPECT_EQ(batch->column(0)->null_count(), 4); + } else if (block.rows() == 4) { + const auto& values = static_cast<const arrow::StringArray&>(*batch->column(0)); + EXPECT_EQ(values.GetString(0), R"({"a":[1, null, "x"]})"); + EXPECT_EQ(values.GetString(1), "{}"); + EXPECT_EQ(values.GetString(2), "42"); + EXPECT_EQ(values.GetString(3), R"("text")"); + } + } +} + +TEST(ArrowFlightVariantTest, LegacySerdeRejectsNativeStorageBeforeInspectingRows) { + auto type = std::make_shared<DataTypeVariant>(); + auto column = documents(type); + auto storage = std::static_pointer_cast<arrow::ExtensionType>(native_variant())->storage_type(); + auto builder = arrow::MakeBuilder(storage, arrow::default_memory_pool()).ValueOrDie(); + NullMap nulls(4, 1); + for (int end : {0, 4}) { + auto status = type->get_serde()->write_column_to_arrow(*column, &nulls, builder.get(), 0, + end, cctz::utc_time_zone()); + EXPECT_TRUE(status.is<ErrorCode::NOT_IMPLEMENTED_ERROR>()) << status; + EXPECT_NE(status.to_string().find("only supports Variant V2"), std::string::npos); + EXPECT_EQ(builder->length(), 0); + } +} + +TEST(ArrowFlightVariantTest, NativeResultPreservesValuesAndSqlNulls) { + TimezoneUtils::load_timezones_to_cache(); + ASSERT_TRUE(register_arrow_variant_extension().ok()); + auto type = std::make_shared<DataTypeVariantV2>(); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({0, 0, 0, 1}); + Block block; + block.insert( + {ColumnNullable::create(documents(type), std::move(nulls)), make_nullable(type), "v"}); + auto schema = arrow::schema({arrow::field("v", native_variant())}); + ArrowFlightArrowBlockConvertor converter(schema, cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + ASSERT_TRUE(batch->ValidateFull().ok()); + EXPECT_EQ(batch->num_rows(), 4); + EXPECT_FALSE(batch->column(0)->IsNull(1)); + EXPECT_TRUE(value_at(*batch->column(0), 1).is_null()); + EXPECT_EQ(value_at(*batch->column(0), 2).get_int(), 42); + EXPECT_TRUE(batch->column(0)->IsNull(3)); + EXPECT_EQ(value_at(*batch->column(0), 0).basic_type(), VariantBasicType::OBJECT); + + // Extension metadata and storage must survive the same IPC boundary used by Flight. + auto output = arrow::io::BufferOutputStream::Create().ValueOrDie(); + auto writer = arrow::ipc::MakeStreamWriter(output, batch->schema()).ValueOrDie(); + ASSERT_TRUE(writer->WriteRecordBatch(*batch).ok()); + ASSERT_TRUE(writer->Close().ok()); + auto input = std::make_shared<arrow::io::BufferReader>(output->Finish().ValueOrDie()); + auto reader = arrow::ipc::RecordBatchStreamReader::Open(input).ValueOrDie(); + auto round_trip = reader->Next().ValueOrDie(); + EXPECT_TRUE(batch->Equals(*round_trip)); + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 3).ok()); + EXPECT_TRUE(value_at(*batch->column(0), 0).is_null()); + EXPECT_EQ(value_at(*batch->column(0), 1).get_int(), 42); +} + +TEST(ArrowFlightVariantTest, SchemaMappingAndConstantScalar) { + auto type = std::make_shared<DataTypeVariantV2>(); + std::shared_ptr<arrow::DataType> mapped; + ASSERT_TRUE(DorisArrowSchemaConvertor("UTC").convert_to_arrow_type(type, &mapped).ok()); + EXPECT_TRUE(mapped->Equals(arrow::utf8())); + ASSERT_TRUE(ArrowFlightSchemaConvertor("UTC").convert_to_arrow_type(type, &mapped).ok()); + EXPECT_TRUE(mapped->Equals(native_variant())); + auto column = type->create_column(); + std::string json = R"("te\"xt\n\u4e2d")"; + Slice slice(json.data(), json.size()); + DataTypeSerDe::FormatOptions options; + ASSERT_TRUE(type->get_serde()->deserialize_one_cell_from_json(*column, slice, options).ok()); + Block block; + block.insert({ColumnConst::create(std::move(column), 3), type, "v"}); + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("v", mapped, false)}), + cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + ASSERT_TRUE(batch->ValidateFull().ok()); + EXPECT_EQ(batch->num_rows(), 3); + EXPECT_EQ(value_at(*batch->column(0), 2).get_string().to_string(), "te\"xt\n中"); +} + +TEST(ArrowFlightVariantTest, NativeSchemaPreservesNestedLogicalMetadata) { + auto variant = std::make_shared<DataTypeVariantV2>(); + const auto nullable = make_nullable(variant); + const auto string = std::make_shared<DataTypeString>(); + DataTypes types {nullable, std::make_shared<DataTypeArray>(nullable), + std::make_shared<DataTypeMap>(string, nullable), + std::make_shared<DataTypeStruct>(DataTypes {nullable}, Strings {"v"})}; + Block block; + for (size_t i = 0; i < types.size(); ++i) { + block.insert({types[i]->create_column(), types[i], std::to_string(i)}); + } + { + std::shared_ptr<arrow::Schema> schema; + ASSERT_TRUE(ArrowFlightSchemaConvertor(block, "UTC").get_arrow_schema(&schema).ok()); + const auto& map = static_cast<const arrow::MapType&>(*schema->field(2)->type()); + for (const auto& field : {schema->field(0), schema->field(1)->type()->field(0), + map.item_field(), schema->field(3)->type()->field(0)}) { + EXPECT_TRUE(field->type()->Equals(native_variant())); + EXPECT_TRUE(field->nullable()); + ASSERT_NE(nullptr, field->metadata()); + EXPECT_EQ("VARIANT", field->metadata()->Get("doris_type").ValueOrDie()); + } + EXPECT_FALSE(map.key_field()->nullable()); + } + std::shared_ptr<arrow::Schema> schema; + ASSERT_TRUE(DorisArrowSchemaConvertor(block, "UTC").get_arrow_schema(&schema).ok()); + EXPECT_TRUE(schema->field(0)->type()->Equals(arrow::utf8())); +} + +TEST(ArrowFlightVariantTest, TypedV2AndNestedStructPreserveNonJsonNumbers) { + auto numbers = ColumnFloat64::create(); + numbers->insert_value(std::numeric_limits<double>::quiet_NaN()); + numbers->insert_value(std::numeric_limits<double>::infinity()); + auto values = ColumnVariantV2::create_typed(make_nullable(std::move(numbers)), + std::make_shared<DataTypeFloat64>()); + auto variant = std::make_shared<DataTypeVariantV2>(); + auto type = std::make_shared<DataTypeStruct>(DataTypes {variant}, Strings {"v"}); + Block block; + block.insert({ColumnStruct::create(Columns {std::move(values)}), type, "s"}); + std::shared_ptr<arrow::DataType> mapped; + ASSERT_TRUE(ArrowFlightSchemaConvertor("UTC").convert_to_arrow_type(type, &mapped).ok()); + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("s", mapped, false)}), + cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + ASSERT_TRUE(batch->ValidateFull().ok()); + const auto& child = *static_cast<const arrow::StructArray&>(*batch->column(0)).field(0); + EXPECT_TRUE(std::isnan(value_at(child, 0).get_double())); + EXPECT_TRUE(std::isinf(value_at(child, 1).get_double())); +} + +TEST(ArrowFlightVariantTest, NestedTimezoneAliasesMatchPublishedSchema) { + TimezoneUtils::load_timezones_to_cache(); + auto variant = std::make_shared<DataTypeVariantV2>(); + auto timestamp = DataTypeFactory::instance().create_data_type(TYPE_TIMESTAMPTZ, false, 0, 6); + auto type = + std::make_shared<DataTypeStruct>(DataTypes {variant, timestamp}, Strings {"v", "t"}); + auto times = timestamp->create_column(); + for (int i = 0; i < 4; ++i) { + times->insert_default(); + } + Block block {{ColumnStruct::create(Columns {documents(variant), std::move(times)}), type, "s"}}; + for (const std::string zone : {"+08:00", "+05:45", "-03:30"}) { + cctz::time_zone timezone; + ASSERT_TRUE(TimezoneUtils::find_cctz_time_zone(zone, timezone)); + std::shared_ptr<arrow::DataType> mapped; + ASSERT_TRUE(ArrowFlightSchemaConvertor(zone).convert_to_arrow_type(type, &mapped).ok()); + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("s", mapped, false)}), + timezone); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_TRUE(batch->schema()->field(0)->type()->Equals(mapped)); + } +} + +TEST(ArrowFlightVariantTest, NativeMetadataContainsOnlySelectedRowKeys) { + for (int rows : {32, 64}) { + SCOPED_TRACE(rows); + auto type = std::make_shared<DataTypeVariantV2>(); + auto column = type->create_column(); + JsonStringToVariantEncoder encoder; + for (int i = 0; i < rows; ++i) { + std::string key = std::string(200, 'k') + std::to_string(i); + std::string json = "{\"" + key + "\":" + std::to_string(i) + "}"; + encoder.add_json({json.data(), json.size()}); + } + auto encoded = encoder.finish_batch(); + ASSERT_EQ(encoded.metadata_ref().dict_size(), rows); + assert_cast<ColumnVariantV2&>(*column).insert_encoded_batch(encoded); + auto nulls = ColumnUInt8::create(); + nulls->get_data().resize_fill(rows, 0); + nulls->get_data()[rows / 2] = 1; + Block block {{ColumnNullable::create(std::move(column), std::move(nulls)), + make_nullable(type), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant())}), cctz::utc_time_zone()); + for (int start : {0, 3}) { + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, + start, rows); + ASSERT_TRUE(status.ok()) << status; + size_t metadata_bytes = 0; + for (int i = start; i < rows; ++i) { + if (i == rows / 2) { + EXPECT_TRUE(batch->column(0)->IsNull(i - start)); + continue; + } + VariantRef value = value_at(*batch->column(0), i - start); + // Internal dictionary sharing must not multiply every other row's keys on the wire. + EXPECT_EQ(value.metadata.dict_size(), 1); + metadata_bytes += value.metadata.size; + VariantRef child; + std::string key = std::string(200, 'k') + std::to_string(i); + ASSERT_TRUE(value.object_find({key.data(), key.size()}, &child)); + EXPECT_EQ(child.get_int(), i); + } + EXPECT_LT(metadata_bytes, static_cast<size_t>(rows - start) * 220); + } + } +} + +TEST(ArrowFlightVariantTest, NestedNativeCompactionPreservesPhysicalScalars) { + VariantBatchBuilder builder; + for (std::string key : {"first", "second"}) { + auto row = builder.begin_row(); + auto object = row.start_object(); + object.add_key({key.data(), key.size()}); + row.add_int(1); + object.finish(); + row.finish(); + } + for (int kind = 0; kind < 5; ++kind) { + auto row = builder.begin_row(); + switch (kind) { + case 0: + row.add_decimal(4200, 2, 4); + break; + case 1: + row.add_float(42.0F); + break; + case 2: + row.add_date(1); + break; + case 3: + row.add_binary({"\0\xff", 2}); + break; + case 4: + row.add_string({"42", 2}); + break; + } + row.finish(); + } + { + auto row = builder.begin_row(); + auto array = row.start_array(); + row.add_decimal(4200, 2, 4); + row.add_float(42.0F); + row.add_string({"42", 2}); + array.finish(); + row.finish(); + } + auto encoded = builder.finish_batch(); + ASSERT_EQ(encoded.metadata_ref().dict_size(), 2); + auto values = ColumnVariantV2::create(); + values->insert_encoded_batch(encoded); + auto offsets = ColumnArray::ColumnOffsets::create(); + for (size_t i = 0; i < encoded.num_rows(); ++i) { + offsets->get_data().push_back(i + 1); + } + auto type = std::make_shared<DataTypeArray>(std::make_shared<DataTypeVariantV2>()); + Block block { + {ColumnArray::create(make_nullable(std::move(values)), std::move(offsets)), type, "a"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("a", arrow::list(native_variant()), false)}), + cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + const auto& elements = *static_cast<const arrow::ListArray&>(*batch->column(0)).values(); + for (int i = 0; i < elements.length(); ++i) { + auto actual = value_at(elements, i); + EXPECT_EQ(actual.metadata.dict_size(), i < 2 ? 1 : 0); + if (i >= 2) { + // Integral decimals and floats must retain physical type/scale rather than normalize to integers. + auto expected = encoded.value_at(i); + EXPECT_EQ(actual.basic_type(), expected.basic_type()); + if (actual.basic_type() == VariantBasicType::PRIMITIVE) { + EXPECT_EQ(actual.primitive_id(), expected.primitive_id()); + } + EXPECT_EQ(actual.value, expected.value); + } + } +} + +TEST(ArrowFlightVariantTest, NativeCompactionKeepsTerminalEmptyContainersAtDepthLimit) { + for (std::string terminal : {"[]", "{}"}) { + std::string json = terminal; + for (size_t i = 0; i < VARIANT_MAX_NESTING_DEPTH; ++i) { + json = "[" + json + "]"; + } + JsonStringToVariantEncoder encoder; + encoder.add_json({json.data(), json.size()}); + encoder.add_json({R"({"unused":0})", 12}); + auto encoded = encoder.finish_batch(); + auto values = ColumnVariantV2::create(); + values->insert_encoded_batch(encoded); + Block block {{std::move(values), std::make_shared<DataTypeVariantV2>(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + auto value = value_at(*batch->column(0), 0); + EXPECT_EQ(value.metadata.dict_size(), 0); + for (size_t i = 0; i < VARIANT_MAX_NESTING_DEPTH; ++i) { + value = value.array_at(0); + } + EXPECT_EQ(value.num_elements(), 0); + EXPECT_EQ(value.basic_type(), + terminal == "[]" ? VariantBasicType::ARRAY : VariantBasicType::OBJECT); + } +} + +TEST(ArrowFlightVariantTest, SchemaRpcPreservesNestedVariantExtensions) { + ASSERT_TRUE(register_arrow_variant_extension().ok()); + for (const auto& type : std::vector<std::shared_ptr<arrow::DataType>> { + arrow::utf8(), native_variant(), arrow::list(native_variant()), + arrow::map(arrow::utf8(), native_variant()), + arrow::struct_({arrow::field("v", native_variant())}), + arrow::list(arrow::struct_({arrow::field("v", native_variant())}))}) { + SCOPED_TRACE(type->ToString()); + auto schema = arrow::schema({arrow::field("result", type)}); + std::string serialized; + // Result schema discovery happens before batch conversion and must also support nested extensions. + auto status = serialize_arrow_schema(&schema, &serialized); + ASSERT_TRUE(status.ok()) << status; + auto input = arrow::io::BufferReader::FromString(serialized); + auto opened = arrow::ipc::RecordBatchStreamReader::Open(input.get()); + ASSERT_TRUE(opened.ok()) << opened.status(); + auto reader = opened.ValueOrDie(); + EXPECT_TRUE(reader->schema()->Equals(*schema, true)); + auto next = reader->Next(); + ASSERT_TRUE(next.ok()) << next.status(); + if (next.ValueOrDie() != nullptr) { + EXPECT_EQ(next.ValueOrDie()->num_rows(), 0); + } + } +} + +TEST(ArrowFlightVariantTest, EmptyResultHasNativeSchema) { + auto type = std::make_shared<DataTypeVariantV2>(); + Block block; + block.insert({type->create_column(), type, "v"}); + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + ASSERT_TRUE(batch->ValidateFull().ok()); + EXPECT_EQ(batch->num_rows(), 0); + EXPECT_TRUE(batch->column(0)->type()->Equals(native_variant())); +} + +TEST(ArrowFlightVariantTest, NestedArrayAndDefaultJsonMode) { + auto type = std::make_shared<DataTypeVariantV2>(); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().assign({2, 4}); + auto array_type = std::make_shared<DataTypeArray>(type); + Block block; + block.insert({ColumnArray::create(make_nullable(documents(type)), std::move(offsets)), + array_type, "a"}); + auto schema = arrow::schema( + {arrow::field("a", arrow::list(arrow::field("item", native_variant(), true)), false)}); + ArrowFlightArrowBlockConvertor converter(schema, cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + ASSERT_TRUE(batch->ValidateFull().ok()); + const auto& values = *static_cast<const arrow::ListArray&>(*batch->column(0)).values(); + EXPECT_EQ(value_at(values, 2).get_int(), 42); + EXPECT_EQ(value_at(values, 3).get_string().to_string(), "text"); + + DorisArrowBlockConvertor json(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(json.init().ok()); + ASSERT_TRUE(json.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + EXPECT_EQ(static_cast<const arrow::ListArray&>(*batch->column(0)).values()->type_id(), + arrow::Type::STRING); + // Native Variant bindings belong to Flight, not the ordinary Arrow export path. + EXPECT_FALSE(DorisArrowBlockConvertor(schema, cctz::utc_time_zone()) + .convert_to_arrow(block, arrow::default_memory_pool(), &batch) + .ok()); +} + +} // namespace +} // namespace doris diff --git a/be/test/format/arrow/arrow_row_batch_test.cpp b/be/test/format/arrow/arrow_row_batch_test.cpp index 8b720c675d8..a77b72f7009 100644 --- a/be/test/format/arrow/arrow_row_batch_test.cpp +++ b/be/test/format/arrow/arrow_row_batch_test.cpp @@ -29,6 +29,7 @@ #include "core/data_type/data_type_number.h" #include "core/data_type/data_type_string.h" #include "core/data_type/data_type_struct.h" +#include "core/data_type/data_type_variant_v2.h" #include "exprs/vexpr_context.h" #include "exprs/vslot_ref.h" #include "format/arrow/arrow_block_convertor.h" @@ -104,7 +105,10 @@ class ArrowLogicalTypeMetadataTest TEST_P(ArrowLogicalTypeMetadataTest, PreservesTopLevelAndNestedFields) { const auto [primitive, name] = GetParam(); - auto type = make_nullable(DataTypeFactory::instance().create_data_type(primitive, false)); + auto type = + make_nullable(primitive == TYPE_VARIANT + ? DataTypePtr(std::make_shared<DataTypeVariantV2>()) + : DataTypeFactory::instance().create_data_type(primitive, false)); auto string_type = std::make_shared<DataTypeString>(); DataTypes types {type, std::make_shared<DataTypeArray>(type), std::make_shared<DataTypeStruct>(DataTypes {type, string_type}, @@ -129,7 +133,10 @@ TEST_P(ArrowLogicalTypeMetadataTest, PreservesTopLevelAndNestedFields) { TEST_P(ArrowLogicalTypeMetadataTest, OldFeReceivesLegacySchema) { const auto [primitive, name] = GetParam(); - auto type = make_nullable(DataTypeFactory::instance().create_data_type(primitive, false)); + auto type = + make_nullable(primitive == TYPE_VARIANT + ? DataTypePtr(std::make_shared<DataTypeVariantV2>()) + : DataTypeFactory::instance().create_data_type(primitive, false)); auto string_type = std::make_shared<DataTypeString>(); DataTypes types { type, std::make_shared<DataTypeArray>(type), diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/ResultSink.java b/fe/fe-core/src/main/java/org/apache/doris/planner/ResultSink.java index b7d7dd4d876..30fa7705b0c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/ResultSink.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/ResultSink.java @@ -37,7 +37,7 @@ public class ResultSink extends DataSink { private TResultSinkType resultSinkType = TResultSinkType.MYSQL_PROTOCOL; public ResultSink(PlanNodeId exchNodeId) { - this.exchNodeId = exchNodeId; + this(exchNodeId, TResultSinkType.MYSQL_PROTOCOL); } public ResultSink(PlanNodeId exchNodeId, TResultSinkType resultSinkType) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java index e0f89217690..678cab055a8 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java +++ b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java @@ -378,6 +378,10 @@ final class FlightSqlQuerySchema { type = Type.STRING; } PrimitiveType primitive = type.getPrimitiveType(); + if (primitive == PrimitiveType.VARIANT) { + return FlightSqlSchemaHelper.nativeVariantField(name, nullable, + Collections.singletonMap("doris_type", primitive.toString())); + } int precision = type instanceof ScalarType ? ((ScalarType) type).getScalarPrecision() : 0; int scale = type instanceof ScalarType ? ((ScalarType) type).getScalarScale() : 0; ArrowType arrowType = FlightSqlSchemaHelper.getArrowType(primitive, precision, scale); @@ -400,13 +404,17 @@ final class FlightSqlQuerySchema { field("value", map.getValueType(), true, false, timezone)))); } else if (type instanceof StructType) { for (StructField child : ((StructType) type).getFields()) { - children.add(field(child.getName(), child.getType(), child.getContainsNull(), false, timezone)); + children.add(field(child.getName(), child.getType(), child.getContainsNull(), + false, timezone)); } } Map<String, String> metadata = null; - if (topLevel && (primitive == PrimitiveType.LARGEINT || primitive == PrimitiveType.IPV4 - || primitive == PrimitiveType.IPV6)) { - metadata = Collections.singletonMap("doris_type", primitive.toString()); + // Execution preserves logical type markers at every depth; Prepare must match them exactly. + if (primitive == PrimitiveType.LARGEINT || primitive == PrimitiveType.IPV4 + || primitive == PrimitiveType.IPV6 || primitive == PrimitiveType.VARIANT + || primitive == PrimitiveType.JSONB) { + metadata = Collections.singletonMap("doris_type", + primitive == PrimitiveType.JSONB ? "JSON" : primitive.toString()); } return new Field(name, new FieldType(nullable, arrowType, null, metadata), children); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlSchemaHelper.java b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlSchemaHelper.java index 6af37014132..0ce25bfb81c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlSchemaHelper.java +++ b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlSchemaHelper.java @@ -23,6 +23,7 @@ import org.apache.doris.catalog.MapType; import org.apache.doris.catalog.PrimitiveType; import org.apache.doris.catalog.StructType; import org.apache.doris.catalog.Type; +import org.apache.doris.common.Config; import org.apache.doris.datasource.CatalogIf; import org.apache.doris.qe.ConnectContext; import org.apache.doris.service.ExecuteEnv; @@ -35,8 +36,10 @@ import org.apache.doris.thrift.TGetDbsParams; import org.apache.doris.thrift.TGetDbsResult; import org.apache.doris.thrift.TGetTablesParams; import org.apache.doris.thrift.TListTableStatusResult; +import org.apache.doris.thrift.TPrimitiveType; import org.apache.doris.thrift.TTableStatus; +import org.apache.arrow.flight.CallStatus; import org.apache.arrow.flight.sql.FlightSqlColumnMetadata; import org.apache.arrow.flight.sql.impl.FlightSql.CommandGetDbSchemas; import org.apache.arrow.flight.sql.impl.FlightSql.CommandGetTables; @@ -170,6 +173,15 @@ public class FlightSqlSchemaHelper { } static Field withDorisTypeMetadata(Field field, Type type) { + if (type.isVariantType()) { + requireVariantV2(); + if (field.getMetadata() == null + || !"arrow.parquet.variant".equals(field.getMetadata().get("ARROW:extension:name"))) { + throw CallStatus.UNIMPLEMENTED.withDescription( + "Backend returned a non-native Variant schema; use Variant V2 " + + "or cast the result to STRING").toRuntimeException(); + } + } List<Field> children = new ArrayList<>(field.getChildren()); if (type.isArrayType()) { children.set(0, withDorisTypeMetadata(children.get(0), ((ArrayType) type).getItemType())); @@ -341,6 +353,10 @@ public class FlightSqlSchemaHelper { /** One column, with its nested types described down to the leaves. */ private static Field buildField(String dbName, String tableName, TColumnDesc desc) { + if (desc.getColumnType() == TPrimitiveType.VARIANT) { + return nativeVariantField(desc.getColumnName(), desc.isIsAllowNull(), + createFlightSqlColumnMetadata(dbName, tableName, desc)); + } ArrowType arrowType = columnDescToArrowType(desc); return new Field(desc.getColumnName(), new FieldType(desc.isIsAllowNull(), arrowType, null, @@ -348,6 +364,26 @@ public class FlightSqlSchemaHelper { arrowChildren(dbName, tableName, desc, arrowType)); } + private static void requireVariantV2() { + // Schema discovery must reject legacy Variant before publishing a native binary layout. + if (!Config.enable_variant_v2) { + throw CallStatus.UNIMPLEMENTED.withDescription( + "Native Arrow Flight output only supports Variant V2, not legacy Variant; " + + "cast the result to STRING for text output").toRuntimeException(); + } + } + + static Field nativeVariantField(String name, boolean nullable, Map<String, String> columnMetadata) { + requireVariantV2(); + Map<String, String> metadata = new HashMap<>(columnMetadata); + // Discovery and execution must share the extension metadata as well as its storage type. + metadata.put("ARROW:extension:name", "arrow.parquet.variant"); + metadata.put("ARROW:extension:metadata", ""); + return new Field(name, new FieldType(nullable, new ArrowType.Struct(), null, metadata), + Arrays.asList(Field.notNullable("metadata", new ArrowType.Binary()), + Field.notNullable("value", new ArrowType.Binary()))); + } + /** * The Arrow children of a complex column, built from the descriptor's own children. * @@ -389,7 +425,8 @@ public class FlightSqlSchemaHelper { Field entries = new Field(MapVector.DATA_VECTOR_NAME, new FieldType(false, new ArrowType.Struct(), null), Arrays.asList(new Field(key.getName(), - new FieldType(false, key.getType(), null), key.getChildren()), + new FieldType(false, key.getType(), null, key.getMetadata()), + key.getChildren()), value)); return Collections.singletonList(entries); case Struct: diff --git a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessorSchemaTest.java b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessorSchemaTest.java index b7f6725478b..144d615d5e4 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessorSchemaTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessorSchemaTest.java @@ -75,18 +75,16 @@ class FlightSqlConnectProcessorSchemaTest { field("ip4", new ArrowType.Int(32, true), true, annotated ? "IPV4" : null), field("ip6", new ArrowType.Utf8(), true, annotated ? "IPV6" : null), field("json", new ArrowType.Utf8(), true, annotated ? "JSON" : null), - field("variant", new ArrowType.Utf8(), true, annotated ? "VARIANT" : null), field("text", new ArrowType.Utf8(), true, null)))), - field("json", new ArrowType.Utf8(), true, annotated ? "JSON" : null), - field("variant", new ArrowType.Utf8(), true, annotated ? "VARIANT" : null)); + field("json", new ArrowType.Utf8(), true, annotated ? "JSON" : null)); } private static List<Type> resultTypes() { return Arrays.asList(new ArrayType(Type.LARGEINT), new MapType(Type.LARGEINT, new StructType( new StructField("ip4", Type.IPV4), new StructField("ip6", Type.IPV6), - new StructField("json", Type.JSONB), new StructField("variant", Type.VARIANT), - new StructField("text", Type.STRING))), Type.JSONB, Type.VARIANT); + new StructField("json", Type.JSONB), + new StructField("text", Type.STRING))), Type.JSONB); } private static Schema fetch(List<Type> types, Schema... schemas) throws Exception { diff --git a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariantTest.java b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariantTest.java new file mode 100644 index 00000000000..ebf73875a12 --- /dev/null +++ b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariantTest.java @@ -0,0 +1,124 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.doris.service.arrowflight; + +import org.apache.doris.catalog.ArrayType; +import org.apache.doris.catalog.MapType; +import org.apache.doris.catalog.StructField; +import org.apache.doris.catalog.StructType; +import org.apache.doris.catalog.Type; +import org.apache.doris.common.Config; +import org.apache.doris.common.jmockit.Deencapsulation; +import org.apache.doris.thrift.TColumnDesc; +import org.apache.doris.thrift.TPrimitiveType; + +import org.apache.arrow.vector.ipc.ReadChannel; +import org.apache.arrow.vector.ipc.message.MessageSerializer; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.types.pojo.Schema; +import org.junit.After; +import org.junit.Assert; +import org.junit.Before; +import org.junit.Test; + +import java.io.ByteArrayInputStream; +import java.nio.channels.Channels; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; + +public class FlightSqlNativeVariantTest { + private boolean originalVariantV2; + + @Before + public void enableVariantV2() { + originalVariantV2 = Config.enable_variant_v2; + Config.enable_variant_v2 = true; + } + + @After + public void restoreVariantV2() { + Config.enable_variant_v2 = originalVariantV2; + } + + @Test + public void legacyVariantIsRejectedWithoutAFormatSwitch() { + Config.enable_variant_v2 = false; + Assert.assertThrows(org.apache.arrow.flight.FlightRuntimeException.class, + () -> FlightSqlSchemaHelper.nativeVariantField("v", true, Collections.emptyMap())); + } + + @Test + public void executionCannotPublishUtf8ForVariant() { + Assert.assertThrows(org.apache.arrow.flight.FlightRuntimeException.class, + () -> FlightSqlSchemaHelper.withDorisTypeMetadata(Field.nullable("v", new ArrowType.Utf8()), + Type.VARIANT)); + } + + @Test + public void schemaKeepsExtensionAcrossIpc() throws Exception { + TColumnDesc variant = new TColumnDesc("item", TPrimitiveType.VARIANT); + variant.setIsAllowNull(true); + TColumnDesc array = new TColumnDesc("a", TPrimitiveType.ARRAY); + array.setChildren(Collections.singletonList(variant)); + Field field = Deencapsulation.invoke(FlightSqlSchemaHelper.class, "buildField", + "test_db", "test_table", array); + Field child = field.getChildren().get(0); + Assert.assertEquals(new ArrowType.Struct(), child.getType()); + Assert.assertEquals("arrow.parquet.variant", child.getMetadata().get("ARROW:extension:name")); + Assert.assertEquals("metadata", child.getChildren().get(0).getName()); + Assert.assertEquals("value", child.getChildren().get(1).getName()); + for (Field storage : child.getChildren()) { + Assert.assertFalse(storage.isNullable()); + Assert.assertEquals(new ArrowType.Binary(), storage.getType()); + } + Schema schema = new Schema(Collections.singletonList(field)); + try (ReadChannel channel = new ReadChannel(Channels.newChannel( + new ByteArrayInputStream(schema.serializeAsMessage())))) { + Assert.assertEquals(schema, MessageSerializer.deserializeSchema(channel)); + } + } + + @Test + public void querySchemaPreservesNativeVariantInNestedFields() { + Type nested = new StructType(new ArrayList<>(Arrays.asList( + new StructField("scalar", Type.VARIANT), + new StructField("array", new ArrayType(Type.VARIANT, true)), + new StructField("map", new MapType(Type.STRING, Type.VARIANT))))); + Field result = Deencapsulation.invoke(FlightSqlQuerySchema.class, "field", + "s", nested, true, true, "UTC"); + Field scalar = result.getChildren().get(0); + Field item = result.getChildren().get(1).getChildren().get(0); + Field value = result.getChildren().get(2).getChildren().get(0).getChildren().get(1); + for (Field leaf : Arrays.asList(scalar, item, value)) { + Assert.assertEquals(new ArrowType.Struct(), leaf.getType()); + Assert.assertEquals("arrow.parquet.variant", leaf.getMetadata().get("ARROW:extension:name")); + Assert.assertEquals("", leaf.getMetadata().get("ARROW:extension:metadata")); + Assert.assertEquals(Arrays.asList(Field.notNullable("metadata", new ArrowType.Binary()), + Field.notNullable("value", new ArrowType.Binary())), leaf.getChildren()); + Assert.assertEquals("VARIANT", leaf.getMetadata().get("doris_type")); + } + // Model execution's metadata enrichment to catch Prepare/DoGet schema mismatches. + Field execution = FlightSqlSchemaHelper.withDorisTypeMetadata(result, nested); + Assert.assertTrue(FlightSqlQuerySchema.matchesExecutionSchema( + new Schema(Collections.singletonList(result)), + new Schema(Collections.singletonList(execution)), Collections.singletonList("s"))); + } + +} diff --git a/fe/pom.xml b/fe/pom.xml index 5f1223aa793..556c47639a6 100644 --- a/fe/pom.xml +++ b/fe/pom.xml @@ -341,7 +341,6 @@ under the License. <hudi-spark.version>hudi-spark3.4.x</hudi-spark.version> <hive.version>3.1.3</hive.version> <hive.common.version>2.3.9</hive.common.version> - <nimbusds.version>9.35</nimbusds.version> <mapreduce.client.version>2.10.1</mapreduce.client.version> <calcite.version>1.33.0</calcite.version> <avatica.version>1.22.0</avatica.version> diff --git a/regression-test/suites/arrow_flight_sql_p0/test_flight_cancel_cleanup.groovy b/regression-test/suites/arrow_flight_sql_p0/test_flight_cancel_cleanup.groovy index aff57aeac4d..9b3a21dd30b 100644 --- a/regression-test/suites/arrow_flight_sql_p0/test_flight_cancel_cleanup.groovy +++ b/regression-test/suites/arrow_flight_sql_p0/test_flight_cancel_cleanup.groovy @@ -73,8 +73,10 @@ suite("test_flight_cancel_cleanup", "arrow_flight_sql") { assertTrue(info.endpoints.size() > 1, "Expected distributed Flight result endpoints") } def ids = info.endpoints.collect { endpoint -> - Any.parseFrom(endpoint.ticket.bytes).unpack(FlightSql.TicketStatementQuery.class) + def queryId = Any.parseFrom(endpoint.ticket.bytes).unpack(FlightSql.TicketStatementQuery.class) .statementHandle.toStringUtf8().split("&")[0] + // Flight tickets omit leading zeroes; the BE diagnostic API requires two 16-digit halves. + queryId.split("-").collect { it.padLeft(16, "0") }.join("-") }.unique() // Abort only one endpoint; the cancellation must reach every participating BE. [info.endpoints[0]].each { endpoint -> diff --git a/regression-test/suites/arrow_flight_sql_p0/test_flight_native_variant.groovy b/regression-test/suites/arrow_flight_sql_p0/test_flight_native_variant.groovy new file mode 100644 index 00000000000..ad490bf332f --- /dev/null +++ b/regression-test/suites/arrow_flight_sql_p0/test_flight_native_variant.groovy @@ -0,0 +1,215 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import org.apache.arrow.driver.jdbc.shaded.com.google.protobuf.Any +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.CallOptions +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.FlightClient +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.Location +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.sql.FlightSqlClient +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.sql.impl.FlightSql +import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.memory.RootAllocator + +import java.util.concurrent.TimeUnit + +suite("test_flight_native_variant", "arrow_flight_sql") { + def frontend = jdbc_sql_return_maparray("SHOW FRONTENDS").find { + it.IsMaster.toString().equalsIgnoreCase("true") && it.Alive.toString().equalsIgnoreCase("true") + } + assertNotNull(frontend) + assertTrue(frontend.ArrowFlightSqlPort.toString().toInteger() > 0) + def database = jdbc_sql("SELECT DATABASE()")[0][0] + def variantV2 = getFeConfig("enable_variant_v2").toBoolean() + def variantV2Function = variantV2 ? "parse_to_variant" : "" + def table = "${database}.flight_native_variant_input" + def allocator = new RootAllocator(Long.MAX_VALUE) + def feClient = FlightClient.builder(allocator, + Location.forGrpcInsecure(frontend.Host.toString(), frontend.ArrowFlightSqlPort.toString().toInteger())).build() + def client = new FlightSqlClient(feClient) + def auth + def read = { String query, Closure inspect, boolean parallel = false, int resultBackendCount = 1, def prepared = null -> + int count = 0 + def info = prepared == null ? client.execute(query, auth) : prepared.execute(auth) + assertFalse(info.endpoints.isEmpty()) + // Multiple buckets alone do not prove coverage of native output from different result BEs. + if (parallel) { + def resultAddresses = info.endpoints.collect { endpoint -> + def fields = Any.parseFrom(endpoint.ticket.bytes).unpack(FlightSql.TicketStatementQuery.class) + .statementHandle.toStringUtf8().split("&") + "${fields[1]}:${fields[2]}".toString() + } + assertEquals(info.endpoints.size(), resultAddresses.toSet().size(), "Duplicate result backends") + if (resultBackendCount > 1) { + assertTrue(info.endpoints.size() > 1, "Expected multiple native Variant result backends") + } + } + info.endpoints.each { endpoint -> + FlightClient.builder(allocator, endpoint.locations[0]).build().withCloseable { beClient -> + beClient.getStream(endpoint.ticket, auth, CallOptions.timeout(30, TimeUnit.SECONDS)).withCloseable { stream -> + assertEquals(info.schema, stream.schema, "Result endpoints must share the published schema") + while (stream.next()) { + inspect(stream.root) + count += stream.root.rowCount + } + } + } + } + count + } + def executeSetting = { String query -> read(query, { root -> }) } + def expectUnsupported = { Closure operation -> + Exception failure = null + try { + operation() + } catch (Exception e) { + failure = e + } + assertNotNull(failure, "Native legacy Variant output must fail") + assertTrue(failure.toString().contains("only supports Variant V2"), failure.toString()) + } + try { + auth = feClient.authenticateBasicToken(context.config.otherConfigs.get("extArrowFlightSqlUser"), + context.config.otherConfigs.get("extArrowFlightSqlPassword")).get() + executeSetting("SET enable_sql_cache=false") + executeSetting("SET enable_nereids_distribute_planner=true") + executeSetting("SET parallel_pipeline_task_num=8") + jdbc_sql("DROP TABLE IF EXISTS ${table}") + jdbc_sql("""CREATE TABLE ${table} (id INT, v VARIANT) + DUPLICATE KEY(id) DISTRIBUTED BY HASH(id) BUCKETS 60 + PROPERTIES("replication_num"="1")""") + jdbc_sql("""INSERT INTO ${table} + SELECT number + 1, ${variantV2Function}(CASE number % 4 + WHEN 0 THEN '42' WHEN 1 THEN '"text"' + WHEN 2 THEN '{"a":[1,null,"x"]}' ELSE NULL END) + FROM numbers("number"="60")""") + def resultBackendCount = jdbc_sql_return_maparray("SHOW TABLETS FROM ${table}") + .collect { it.BackendId }.unique().size() + [false, true].each { parallel -> + executeSetting("SET enable_parallel_result_sink=${parallel}") + // Text output is an explicit SQL conversion, for either Variant representation. + assertEquals(60, read("SELECT CAST(v AS STRING) AS text_value FROM ${table}", { root -> + assertEquals("Utf8", root.getVector(0).field.type.toString()) + })) + if (!variantV2) { + // Reject by type even for constants, SQL NULL, and empty result sets. + ["SELECT id, v FROM ${table}", + "SELECT CAST(42 AS VARIANT) AS v", + "SELECT CAST(NULL AS VARIANT) AS v", + "SELECT v FROM ${table} WHERE id < 0"].each { query -> + expectUnsupported { read(query.toString(), { root -> }) } + expectUnsupported { client.getExecuteSchema(query.toString(), auth) } + expectUnsupported { + def prepared = client.prepare(query.toString(), auth) + try { + prepared.resultSetSchema + } finally { + prepared.close(auth) + } + } + } + return + } + def seen = [] + assertEquals(60, read("SELECT id, v FROM ${table}", { root -> + def vector = root.getVector(1) + def field = vector.field + assertEquals("arrow.parquet.variant", field.metadata.get("ARROW:extension:name")) + assertEquals("Struct", field.type.toString()) + assertEquals(["metadata", "value"], field.children.collect { it.name }) + field.children.each { child -> + assertFalse(child.nullable) + assertEquals("Binary", child.type.toString()) + } + for (int i = 0; i < root.rowCount; i++) { + int id = root.getVector(0).get(i) + seen.add(id) + assertEquals(id % 4 == 0, vector.isNull(i)) + if (id % 4 != 0) { + assertTrue(vector.getChild("metadata").get(i).length > 0) + assertTrue(vector.getChild("value").get(i).length > 0) + if (id % 4 == 1) { + assertEquals([12, 42], vector.getChild("value").get(i).collect { it & 0xff }) + } + } + } + }, parallel, resultBackendCount)) + assertEquals((1..60).toList(), seen.sort()) + def scannedColumns = "v, ARRAY(v) AS a, MAP('key', v) AS m, STRUCT(v) AS s" + // Prepare and GetSchema must advertise the same Variant leaves as execution. + ["SELECT CAST(42 AS VARIANT) AS v", + "SELECT ${scannedColumns} FROM ${table} WHERE id = 1", + "SELECT ${scannedColumns} FROM ${table} WHERE id < 0"].eachWithIndex { query, index -> + def prepared = client.prepare(query.toString(), auth) + try { + def schema = prepared.resultSetSchema + assertEquals(schema, client.getExecuteSchema(query.toString(), auth).schema) + assertEquals(schema, prepared.fetchSchema(auth).schema) + def leaves = [schema.fields[0]] + if (index > 0) { + leaves.add(schema.fields[1].children[0]) + leaves.add(schema.fields[2].children[0].children[1]) + leaves.add(schema.fields[3].children[0]) + } + leaves.each { field -> + assertEquals("Struct", field.type.toString()) + assertEquals("arrow.parquet.variant", field.metadata.get("ARROW:extension:name")) + } + assertEquals(index == 2 ? 0 : 1, read(query.toString(), { root -> }, false, 1, prepared)) + } finally { + prepared.close(auth) + } + } + } + if (!variantV2) { + return + } + // Each row owns its wire dictionary; unrelated keys must not multiply Arrow metadata. + def keyExpression = "CONCAT(REPEAT('k', 244), LPAD(CAST(number AS STRING), 6, '0'))" + def variantExpression = """parse_to_variant(CONCAT('{"', ${keyExpression}, '":', CAST(number AS STRING), '}'))""" + def metadataRows = [] + assertEquals(128, read("SELECT number AS id, ${variantExpression} AS v FROM numbers(\"number\"=\"128\")", { root -> + def variant = root.getVector(1) + assertEquals("arrow.parquet.variant", variant.field.metadata.get("ARROW:extension:name")) + for (int i = 0; i < root.rowCount; ++i) { + metadataRows.add(root.getVector(0).get(i).intValue()) + assertTrue(variant.getChild("metadata").get(i).length < 270, + "A row must not repeat the batch dictionary") + } + })) + assertEquals((0..<128).toList(), metadataRows.sort()) + // A folded constant must use the same wire representation as a scanned Variant column. + assertEquals(1, read("SELECT parse_to_variant('42') AS v", { root -> + assertEquals("arrow.parquet.variant", root.getVector(0).field.metadata.get("ARROW:extension:name")) + })) + assertEquals(0, read("SELECT parse_to_variant(CAST(id AS STRING)) AS v FROM ${table} WHERE id < 0", { root -> })) + } finally { + try { + if (auth != null) { + feClient.closeSession(new org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.CloseSessionRequest(), auth) + } + } finally { + try { + client.close() + } finally { + try { + allocator.close() + } finally { + jdbc_sql("DROP TABLE IF EXISTS ${table}") + } + } + } + } +} diff --git a/regression-test/suites/flink_connector_p0/flink_connector_type.groovy b/regression-test/suites/flink_connector_p0/flink_connector_type.groovy index 26d4b33b1b8..6e6a8d25744 100644 --- a/regression-test/suites/flink_connector_p0/flink_connector_type.groovy +++ b/regression-test/suites/flink_connector_p0/flink_connector_type.groovy @@ -24,13 +24,15 @@ import org.awaitility.Awaitility suite("flink_connector_type") { + def inputTable = "test_types_input" def tableName1 = "test_types_source" def tableName2 = "test_types_sink" + sql """DROP TABLE IF EXISTS ${inputTable}""" sql """DROP TABLE IF EXISTS ${tableName1}""" sql """DROP TABLE IF EXISTS ${tableName2}""" sql """ - CREATE TABLE `test_types_source` ( + CREATE TABLE `${inputTable}` ( `id` int, `c1` boolean, `c2` tinyint, @@ -59,9 +61,9 @@ PROPERTIES ( ); """; - sql """CREATE TABLE `test_types_sink` like `test_types_source` """ + sql """CREATE TABLE `${tableName2}` like `${inputTable}` """ - sql """ INSERT INTO `test_types_source` + sql """ INSERT INTO `${inputTable}` VALUES ( 1, @@ -106,6 +108,17 @@ VALUES '{"B":"variant_value1"}' );""" + // The connector declares c18 as STRING and cannot decode native Arrow Variant. + // Materialize an explicit text projection while retaining Variant ingestion in the sink. + sql """ + CREATE TABLE `${tableName1}` + DISTRIBUTED BY HASH(`id`) BUCKETS 1 + PROPERTIES ("replication_num" = "1") + AS SELECT id, c1, c2, c3, c4, c5, c6, c7, c8, c9, + c10, c11, c12, c13, c14, c15, c16, c17, CAST(c18 AS STRING) AS c18 + FROM `${inputTable}` + """ + def thisDb = sql """select database()"""; thisDb = thisDb[0][0]; logger.info("current database is ${thisDb}"); @@ -143,19 +156,21 @@ VALUES run_cmd.addAll(addOpens.tokenize()) run_cmd.addAll(["-cp", jarName, "org.apache.doris.FlinkConnectorTypeCase", "--doris-fe-address", context.config.feHttpAddress, - "--doris-database", "regression_test_flink_connector_p0", + "--doris-database", thisDb, "--doris-user", context.config.feHttpUser, "--doris-password", context.config.feHttpPassword]) run_cmd.addAll(getDorisConnectorTlsArgs()) logger.info("run_cmd : ${run_cmd.join(' ')}") - def run_flink_jar = run_cmd.execute().getText() - logger.info("result: $run_flink_jar") + // Drain both streams and check the exit code so a failed Flink job cannot look successful. + def flinkProcess = new ProcessBuilder(run_cmd.collect { it.toString() }).redirectErrorStream(true).start() + logger.info("result: ${flinkProcess.text}") + assertEquals(0, flinkProcess.waitFor()) // The publish in the commit phase is asynchronous Awaitility.await().atMost(30, SECONDS).pollInterval(1, SECONDS).await().until( { def resultTbl = sql """ select count(1) from test_types_sink""" logger.info("retry test_types_sink count: $resultTbl") - resultTbl.size() >= 1 + resultTbl[0][0] == 2 }) logger.info("flink job execute finished."); diff --git a/samples/arrow-flight-sql/python/README.md b/samples/arrow-flight-sql/python/README.md index 7a7586940a1..42762c29cd7 100644 --- a/samples/arrow-flight-sql/python/README.md +++ b/samples/arrow-flight-sql/python/README.md @@ -60,4 +60,41 @@ The tests execute only read-only queries. # Notes For more details, refer to [Python Usage] in the document https://doris.apache.org/zh-CN/docs/dev/db-connect/arrow-flight-sql-connect - \ No newline at end of file + + +# Native Variant V2 results + +On branch-4.1 builds with this feature and Variant V2 enabled on FE and BE, Arrow +Flight SQL / ADBC returns Variant V2 as native binary values automatically: + +```sql +SELECT parse_to_variant('{"key":42}') AS v; +-- Request text explicitly when the client needs JSON strings. +SELECT CAST(variant_column AS STRING) FROM example_table; +``` + +Variant V2 fields, including nested fields, use the `arrow.parquet.variant` extension +with `struct<metadata: binary not null, value: binary not null>` storage. SQL NULL is +a null struct. Variant null is a non-null struct containing the encoded null value. +V2 values retain their physical scalar types and decimal scales. Each Arrow row carries +only the dictionary keys it uses, rather than copying keys from unrelated rows. + +Legacy Variant is unsupported, including nested legacy fields, SQL NULL and empty +query results. Use an explicit SQL cast to STRING for text output. `parse_to_variant` +follows the configured Variant representation; it does not convert legacy storage to +V2. Native encoding accepts up to 128 nested levels. + +ADBC can transport this schema and its binary values. A client without a registered +Variant extension exposes the struct with `ARROW:extension:name` field metadata. +Receiving native VARIANT does not automatically decode it to Python dictionaries or +pandas objects; use a Parquet Variant decoder, or explicitly cast the result to STRING. + +To check ADBC query and partition reads against a running Variant V2 cluster: + +```bash +pip install adbc_driver_flightsql pyarrow +export DORIS_FLIGHT_URI='grpc://localhost:8815' +export DORIS_USER='root' +# Set DORIS_PASSWORD in the environment if authentication requires it. +python test_variant.py +``` diff --git a/samples/arrow-flight-sql/python/test_variant.py b/samples/arrow-flight-sql/python/test_variant.py new file mode 100644 index 00000000000..35ab784f5b9 --- /dev/null +++ b/samples/arrow-flight-sql/python/test_variant.py @@ -0,0 +1,96 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +""" +Licensed to the Apache Software Foundation (ASF) under one +or more contributor license agreements. See the NOTICE file +distributed with this work for additional information +regarding copyright ownership. The ASF licenses this file +to you under the Apache License, Version 2.0 (the +"License"); you may not use this file except in compliance +with the License. You may obtain a copy of the License at + +http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, +software distributed under the License is distributed on an +"AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +KIND, either express or implied. See the License for the +specific language governing permissions and limitations +under the License. +""" + +"""Run against a Variant V2 cluster with DORIS_FLIGHT_URI, DORIS_USER and DORIS_PASSWORD.""" + +import os +import unittest + +import adbc_driver_flightsql +import adbc_driver_manager +import pyarrow as pa + + +class NativeVariantTest(unittest.TestCase): + def test_query_and_partitions(self): + uri = os.environ["DORIS_FLIGHT_URI"] + options = { + adbc_driver_manager.DatabaseOptions.USERNAME.value: os.environ.get("DORIS_USER", "root"), + adbc_driver_manager.DatabaseOptions.PASSWORD.value: os.environ.get("DORIS_PASSWORD", ""), + } + query = """SELECT 1 AS id, parse_to_variant('42') AS v + UNION ALL SELECT 2, parse_to_variant(CAST(NULL AS STRING))""" + with adbc_driver_flightsql.connect(uri, db_kwargs=options) as database: + with adbc_driver_manager.AdbcConnection(database) as connection: + def execute(sql): + with adbc_driver_manager.AdbcStatement(connection) as statement: + statement.set_sql_query(sql) + stream, _ = statement.execute_query() + return pa.RecordBatchReader._import_from_c(stream.address).read_all() + + execute("SET enable_sql_cache=false") + for parallel in (False, True): + execute(f"SET enable_parallel_result_sink={str(parallel).lower()}") + table = execute(query) + self.check_result(table) + # ExecuteSchema and Prepare must agree with the subsequently fetched batches. + with adbc_driver_manager.AdbcStatement(connection) as statement: + statement.set_sql_query(query) + schema_handle = statement.execute_schema() + schema = pa.Schema._import_from_c(schema_handle.address) + self.assertEqual(schema, table.schema) + statement.prepare() + stream, _ = statement.execute_query() + prepared = pa.RecordBatchReader._import_from_c(stream.address).read_all() + self.assertEqual(prepared.schema, schema) + self.check_result(prepared) + + with adbc_driver_manager.AdbcStatement(connection) as statement: + statement.set_sql_query(query) + partitions, _, _ = statement.execute_partitions() + tables = [] + for partition in partitions: + stream = connection.read_partition(partition) + tables.append(pa.RecordBatchReader._import_from_c(stream.address).read_all()) + self.check_result(pa.concat_tables(tables)) + + def check_result(self, table): + table = table.sort_by("id") + self.assertEqual(table.num_rows, 2) + field = table.schema.field("v") + values = table.column("v").combine_chunks() + # Clients without a registered extension expose its storage plus field metadata. + if isinstance(values, pa.ExtensionArray): + self.assertEqual(values.type.extension_name, "arrow.parquet.variant") + values = values.storage + else: + self.assertEqual(field.metadata[b"ARROW:extension:name"], b"arrow.parquet.variant") + self.assertTrue(pa.types.is_struct(values.type)) + self.assertEqual([child.name for child in values.type], ["metadata", "value"]) + self.assertIsNone(values[1].as_py()) + row = values[0].as_py() + self.assertTrue(row["metadata"]) + # The Parquet Variant primitive INT8 tag is 3 << 2, followed by its signed byte. + self.assertEqual(row["value"], bytes([12, 42])) + + +if __name__ == "__main__": + unittest.main() --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
