diff --git a/be/src/agent/heartbeat_server.cpp b/be/src/agent/heartbeat_server.cpp index 5cf8c7df01da7b..9e1712b9e020b5 100644 --- a/be/src/agent/heartbeat_server.cpp +++ b/be/src/agent/heartbeat_server.cpp @@ -85,6 +85,7 @@ void HeartbeatServer::heartbeat(THeartbeatResult& heartbeat_result, heartbeat_result.backend_info.__set_be_rpc_port(-1); heartbeat_result.backend_info.__set_brpc_port(config::brpc_port); heartbeat_result.backend_info.__set_arrow_flight_sql_port(config::arrow_flight_sql_port); + heartbeat_result.backend_info.__set_arrow_flight_native_variant_supported(true); heartbeat_result.backend_info.__set_version(get_short_version()); heartbeat_result.backend_info.__set_be_start_time(_be_epoch); heartbeat_result.backend_info.__set_be_node_role(config::be_node_role); diff --git a/be/src/core/column/column_variant.h b/be/src/core/column/column_variant.h index 28a6df280cf3aa..b60b3689f31dae 100644 --- a/be/src/core/column/column_variant.h +++ b/be/src/core/column/column_variant.h @@ -366,6 +366,9 @@ class ColumnVariant final : public COWHelper { // Only single scalar root column bool is_scalar_variant() const; + // Output adapters must use the same root/document precedence as legacy JSON serialization. + bool is_visible_root_value(size_t nrow) const; + ColumnPtr get_root() const { return subcolumns.get_root()->data.get_finalized_column_ptr(); } bool has_subcolumn(const PathInData& key) const; @@ -680,8 +683,6 @@ class ColumnVariant final : public COWHelper { size_t start, size_t length); bool try_add_new_subcolumn(const PathInData& path); - - bool is_visible_root_value(size_t nrow) const; }; } // namespace doris 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 3bd9be5797fb5a..0eed9ac289edcc 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 @@ -19,8 +19,12 @@ #include +#include +#include #include #include +#include +#include #include "common/cast_set.h" #include "common/config.h" @@ -28,19 +32,262 @@ #include "common/status.h" #include "core/assert_cast.h" #include "core/column/column.h" +#include "core/column/column_array.h" +#include "core/column/column_map.h" +#include "core/column/column_struct.h" #include "core/column/column_variant.h" +#include "core/column/variant_v2/column_variant_v2.h" +#include "core/column/variant_v2/column_variant_v2_typed_column.h" +#include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_map.h" +#include "core/data_type/data_type_nullable.h" +#include "core/data_type/data_type_struct.h" #include "core/data_type_serde/data_type_serde.h" +#include "core/data_type_serde/data_type_variant_v2_serde.h" +#include "core/data_type_serde/variant_arrow_utils.h" #include "core/field.h" #include "core/string_ref.h" #include "core/types.h" #include "core/value/jsonb_value.h" #include "exec/common/variant_util.h" +#include "exprs/function/parse/variant_jsonb_parse.h" +#include "exprs/function/parse/variant_string_parse.h" #include "util/json/json_parser.h" #include "util/jsonb_writer.h" namespace doris { namespace { +Status append_legacy_arrow_document(const ColumnVariant& column, size_t index, + VariantBatchBuilder::Row& output, + const DataTypeSerDe::FormatOptions& options, size_t depth); + +// Legacy CAST accepts more root families than V2 CAST. Encode their structure here so +// Flight output does not reject valid roots or lose typed leaves through JSON reparsing. +Status append_legacy_arrow_value(const IColumn& column, const DataTypePtr& type, size_t index, + VariantBatchBuilder::Row& output, + const DataTypeSerDe::FormatOptions& options, size_t depth = 0) { + if (depth > VARIANT_MAX_NESTING_DEPTH) { + return Status::NotSupported( + "Native Arrow Variant nesting exceeds {}; " + "use enable_arrow_flight_sql_native_variant=false for UTF8 output", + VARIANT_MAX_NESTING_DEPTH); + } + if (const auto* constant = check_and_get_column(column)) { + return append_legacy_arrow_value(constant->get_data_column(), type, 0, output, options, + depth); + } + if (const auto* nullable = check_and_get_column(column)) { + if (nullable->is_null_at(index)) { + output.add_null(); + return Status::OK(); + } + return append_legacy_arrow_value(nullable->get_nested_column(), remove_nullable(type), + index, output, options, depth); + } + const auto primitive = type->get_primitive_type(); + if (is_supported_variant_typed_identity(primitive)) { + dispatch_variant_typed_column( + column, primitive, [&](const auto& scalar) { + with_variant_typed_scalar( + scalar, index, cast_set(type->get_scale()), + [&](const VariantScalarRef& value) { output.add_scalar(value); }); + }); + } else if (primitive == TYPE_TIMEV2) { + // TIMEV2 already stores microseconds; treating its physical double as a number loses its type. + const double micros = assert_cast(column).get_data()[index]; + // Parquet TIME is a time of day, whereas Doris TIME also represents signed durations. + // Reject unrepresentable durations instead of wrapping them or emitting invalid TIME values. + constexpr int64_t micros_per_day = 86400000000; + if (!std::isfinite(micros) || micros < 0 || micros >= micros_per_day || + std::llround(micros) >= micros_per_day) { + return Status::NotSupported( + "Native Arrow Variant TIMEV2 requires a time in [00:00:00, 24:00:00); " + "use enable_arrow_flight_sql_native_variant=false for UTF8 output"); + } + output.add_time_ntz_micros(std::llround(micros)); + } else if (primitive == TYPE_VARBINARY) { + // Binary leaves must retain arbitrary bytes, including NUL and non-UTF8 data. + output.add_binary(column.get_data_at(index)); + } else if (primitive == TYPE_JSONB) { + // JSONB leaves retain their own depth limit, but also consume the enclosing Variant depth. + try { + jsonb_to_variant(column.get_data_at(index), output, cast_set(depth)); + } catch (const Exception& e) { + if (e.code() != ErrorCode::INVALID_ARGUMENT) { + return e.to_status(); + } + return Status::NotSupported( + "Native Arrow Variant cannot encode JSONB leaf: {}; " + "use enable_arrow_flight_sql_native_variant=false for UTF8 output", + e.what()); + } + } else if (primitive == TYPE_ARRAY) { + const auto& array = assert_cast(column); + const auto& array_type = assert_cast(*type); + auto scope = output.start_array(); + for (size_t element = array.offset_at(index); element < array.get_offsets()[index]; + ++element) { + RETURN_IF_ERROR(append_legacy_arrow_value(array.get_data(), + array_type.get_nested_type(), element, output, + options, depth + 1)); + } + scope.finish(); + } else if (primitive == TYPE_MAP) { + const auto& map = assert_cast(column); + const auto& map_type = assert_cast(*type); + auto scope = output.start_object(); + for (size_t element = map.get_offsets()[static_cast(index) - 1]; + element < map.get_offsets()[index]; ++element) { + // Variant object keys cannot distinguish SQL NULL from the literal string "null". + if (map.get_keys().is_null_at(element)) { + return Status::NotSupported( + "Native Arrow Variant cannot represent MAP with NULL keys; " + "use enable_arrow_flight_sql_native_variant=false for UTF8 output"); + } + auto key = map_type.get_key_type()->to_string(map.get_keys(), element, options); + scope.add_key({key.data(), key.size()}); + RETURN_IF_ERROR(append_legacy_arrow_value(map.get_values(), map_type.get_value_type(), + element, output, options, depth + 1)); + } + scope.finish(); + } else if (primitive == TYPE_STRUCT) { + const auto& structure = assert_cast(column); + const auto& struct_type = assert_cast(*type); + auto scope = output.start_object(); + for (size_t field = 0; field < struct_type.get_elements().size(); ++field) { + const auto& name = struct_type.get_element_names()[field]; + scope.add_key({name.data(), name.size()}); + RETURN_IF_ERROR(append_legacy_arrow_value(structure.get_column(field), + struct_type.get_element(field), index, output, + options, depth + 1)); + } + scope.finish(); + } else if (primitive == TYPE_VARIANT) { + if (const auto* legacy = check_and_get_column(column)) { + const bool visible = legacy->is_scalar_variant() + ? !legacy->get_root()->is_null_at(index) + : legacy->is_visible_root_value(index); + if (visible) { + return append_legacy_arrow_value(*legacy->get_root(), legacy->get_root_type(), + index, output, options, depth); + } + RETURN_IF_ERROR(append_legacy_arrow_document(*legacy, index, output, options, depth)); + } else { + // Reuse the selected-value import: each leaf may share a large dictionary, and + // its enclosing legacy containers must count toward the native depth limit. + Status status = Status::OK(); + visit_variant_v2_values( + column, index, index + 1, {}, [&](size_t) { output.add_null(); }, + [&](size_t, VariantRef value) { + status = append_flight_variant_value(value, output, depth); + }); + RETURN_IF_ERROR(status); + } + } else { + return Status::NotSupported("Native Arrow Variant does not support {} roots", + type->get_name()); + } + return Status::OK(); +} + +// Assemble the same flattened document paths as legacy JSON output, but encode each +// leaf with its stored type. JSON reparsing loses decimal precision and date identity. +Status append_legacy_arrow_document(const ColumnVariant& column, size_t index, + VariantBatchBuilder::Row& output, + const DataTypeSerDe::FormatOptions& options, size_t depth) { + struct Field { + std::string path; + ColumnPtr values; + DataTypePtr type; + size_t row; + }; + std::vector fields; + auto append_serialized = [&](const ColumnMap& map) { + const auto& keys = assert_cast(map.get_keys()); + const auto& values = assert_cast(map.get_values()); + for (size_t i = map.get_offsets()[static_cast(index) - 1]; + i < map.get_offsets()[index]; ++i) { + ColumnVariant::Subcolumn subcolumn(0, true); + subcolumn.deserialize_from_binary_column(&values, i); + subcolumn.finalize(); + fields.push_back({keys.get_data_at(i).to_string(), subcolumn.get_finalized_column_ptr(), + subcolumn.get_least_common_type(), 0}); + } + }; + const auto& snapshot = assert_cast(*column.get_doc_value_column()); + if (snapshot.get_offsets()[index] != 0) { + // Document snapshots are authoritative, including empty rows after a populated snapshot. + append_serialized(snapshot); + } else { + for (const auto& subcolumn : column.get_subcolumns()) { + if (subcolumn->path.empty() || subcolumn->data.is_null_at(index) || + subcolumn->data.is_empty_nested(index)) { + continue; + } + if (subcolumn->data.is_finalized()) { + fields.push_back({subcolumn->path.get_path(), + subcolumn->data.get_finalized_column_ptr(), + subcolumn->data.get_least_common_type(), index}); + } else { + // Scans may leave lazy defaults or multiple parts. Materialize only this row, + // without mutating shared input or repeatedly copying an entire batch. + auto value = subcolumn->data.cut(index, 1); + value.finalize(); + fields.push_back({subcolumn->path.get_path(), value.get_finalized_column_ptr(), + value.get_least_common_type(), 0}); + } + } + append_serialized(assert_cast(*column.get_sparse_column())); + } + std::sort(fields.begin(), fields.end(), + [](const auto& a, const auto& b) { return a.path < b.path; }); + std::vector prefix; + std::vector objects; + objects.push_back(output.start_object()); + for (const auto& field : fields) { + std::vector parts; + std::string_view path(field.path); + while (true) { + const auto dot = path.find('.'); + parts.push_back(path.substr(0, dot)); + if (depth + parts.size() > VARIANT_MAX_NESTING_DEPTH) { + return Status::NotSupported( + "Native Arrow Variant nesting exceeds {}; " + "use enable_arrow_flight_sql_native_variant=false for UTF8 output", + VARIANT_MAX_NESTING_DEPTH); + } + if (dot == std::string_view::npos) { + break; + } + path.remove_prefix(dot + 1); + } + size_t common = 0; + while (common < prefix.size() && common + 1 < parts.size() && + prefix[common] == parts[common]) { + ++common; + } + while (prefix.size() > common) { + objects.back().finish(); + objects.pop_back(); + prefix.pop_back(); + } + for (size_t i = common; i + 1 < parts.size(); ++i) { + objects.back().add_key({parts[i].data(), parts[i].size()}); + objects.push_back(output.start_object()); + prefix.push_back(parts[i]); + } + objects.back().add_key({parts.back().data(), parts.back().size()}); + RETURN_IF_ERROR(append_legacy_arrow_value(*field.values, field.type, field.row, output, + options, depth + parts.size())); + } + while (!objects.empty()) { + objects.back().finish(); + objects.pop_back(); + } + return Status::OK(); +} + template Status write_variant_column_to_arrow_impl(const IColumn& column, const ColumnVariant& var, const NullMap* null_map, BuilderType& builder, @@ -157,6 +404,97 @@ 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(column); + if (array_builder->type()->id() == arrow::Type::STRUCT) { + // Keep legacy scalar and document leaves in their original types. + // The outer null map must remain SQL NULL on the wire. + if (start < 0 || end < start || end > column.size() || + (null_map != nullptr && end > null_map->size())) { + return Status::InvalidArgument("Invalid Variant Arrow row range [{}, {})", start, end); + } + // A legacy null root renders as {}, not Variant null, even in a scalar-only batch. + if (var->is_scalar_variant() && !var->get_root()->has_null(start, end)) { + auto scalar_type = remove_nullable(var->get_root_type()); + // Unsupported roots use the row path below, after applying the outer SQL null mask. + if (is_supported_variant_typed_identity(scalar_type->get_primitive_type())) { + // Avoid a JSON round trip that would turn exact decimal roots into doubles. + auto typed = + ColumnVariantV2::create_typed(make_nullable(var->get_root()), scalar_type); + return DataTypeVariantV2SerDe().write_column_to_arrow( + *typed, null_map, array_builder, start, end, ctz); + } + } + const size_t rows = end - start; + NullMap selected_nulls(rows, 0); + NullMap root_mask(rows, 1); + bool has_roots = false; + bool has_documents = false; + for (size_t row = 0; row < rows; ++row) { + selected_nulls[row] = null_map != nullptr && (*null_map)[start + row]; + if (selected_nulls[row]) { + continue; + } + const bool root_visible = var->is_scalar_variant() + ? !var->get_root()->is_null_at(start + row) + : var->is_visible_root_value(start + row); + root_mask[row] = !root_visible; + has_roots |= root_visible; + has_documents |= !root_visible; + } + ColumnPtr roots; + if (has_roots) { + VariantBatchBuilder builder(VariantBatchBuilder::ReserveHint {.rows = rows}); + FormatOptions options; + options.timezone = &ctz; + for (size_t index = 0; index < rows; ++index) { + auto row = builder.begin_row(); + if (root_mask[index]) { + row.add_null(); + } else { + RETURN_IF_ERROR(append_legacy_arrow_value( + *var->get_root(), var->get_root_type(), start + index, row, options)); + } + row.finish(); + } + auto values = builder.finish_batch(); + auto encoded = ColumnVariantV2::create(); + encoded->insert_encoded_batch(values); + roots = std::move(encoded); + } + ColumnPtr documents; + if (has_documents || !has_roots) { + VariantBatchBuilder builder(VariantBatchBuilder::ReserveHint {.rows = rows}); + FormatOptions options; + options.timezone = &ctz; + for (size_t index = 0; index < rows; ++index) { + auto row = builder.begin_row(); + if (root_mask[index] && !selected_nulls[index]) { + RETURN_IF_ERROR( + append_legacy_arrow_document(*var, start + index, row, options, 0)); + } else { + row.add_null(); + } + row.finish(); + } + auto values = builder.finish_batch(); + auto encoded = ColumnVariantV2::create(); + encoded->insert_encoded_batch(values); + documents = std::move(encoded); + } + // Write contiguous root/document runs without building another copy of the encoded batch. + for (size_t first = 0; first < rows;) { + const bool use_root = static_cast(roots) && (!root_mask[first] || !documents); + size_t last = first + 1; + while (last < rows && + use_root == (static_cast(roots) && (!root_mask[last] || !documents))) { + ++last; + } + RETURN_IF_ERROR(DataTypeVariantV2SerDe().write_column_to_arrow( + *(use_root ? roots : documents), &selected_nulls, array_builder, first, last, + ctz)); + first = last; + } + return Status::OK(); + } if (array_builder->type()->id() == arrow::Type::LARGE_STRING) { auto& builder = assert_cast(*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 0e0d7a58d2f091..ebcca4b1ab59c9 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 #include #include +#include #include #include #include @@ -42,6 +43,7 @@ #include "core/data_type/data_type_nullable.h" #include "core/data_type/data_type_string.h" #include "core/data_type_serde/data_type_string_serde.h" +#include "core/data_type_serde/variant_arrow_utils.h" #include "core/types.h" #include "core/value/jsonb_value.h" #include "core/value/variant/variant_batch_builder.h" @@ -51,6 +53,49 @@ #include "util/mysql_row_buffer.h" namespace doris { + +// 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) { + const auto basic_type = value.basic_type(); + if (depth > VARIANT_MAX_NESTING_DEPTH) { + return Status::NotSupported( + "Native Arrow Variant nesting exceeds {}; " + "use enable_arrow_flight_sql_native_variant=false for UTF8 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(); +} + namespace { using MetaIdsColumn = ColumnVector; @@ -625,13 +670,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(*array_builder); @@ -640,7 +686,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, got {}", + "Variant Arrow writer requires struct, got {}", type->ToString()); } auto& metadata_builder = assert_cast(*builder.field_builder(0)); @@ -657,10 +703,27 @@ Status write_iceberg_variant(const IColumn& column, const NullMap* null_map, if (!status.ok()) { return; } + std::optional 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::max() || value.value.size > std::numeric_limits::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 +796,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(*array_builder), first, last, options); @@ -759,7 +825,7 @@ Status DataTypeVariantV2SerDe::write_column_to_iceberg_arrow( const std::shared_ptr&, const IColumn& column, const NullMap* null_map, const std::shared_ptr&, 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/core/data_type_serde/variant_arrow_utils.h b/be/src/core/data_type_serde/variant_arrow_utils.h new file mode 100644 index 00000000000000..0f0e00c3536a1b --- /dev/null +++ b/be/src/core/data_type_serde/variant_arrow_utils.h @@ -0,0 +1,32 @@ +// 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. + +#pragma once + +#include + +#include "common/status.h" +#include "core/value/variant/variant_batch_builder.h" + +namespace doris { + +// Import only values obtained from validated ColumnVariantV2 storage. Depth includes any +// enclosing legacy containers so all native Flight paths enforce the same nesting limit. +Status append_flight_variant_value(VariantRef value, VariantBatchBuilder::Row& output, + size_t depth = 0); + +} // namespace doris diff --git a/be/src/exec/operator/result_sink_operator.cpp b/be/src/exec/operator/result_sink_operator.cpp index 24c9f6ad18eed1..05d9f3e9b2ff51 100644 --- a/be/src/exec/operator/result_sink_operator.cpp +++ b/be/src/exec/operator/result_sink_operator.cpp @@ -57,9 +57,9 @@ Status ResultSinkLocalState::init(RuntimeState* state, LocalSinkStateInfo& info) } else { std::shared_ptr arrow_schema; if (p._sink_type == TResultSinkType::ARROW_FLIGHT_PROTOCOL) { - RETURN_IF_ERROR(get_arrow_schema_from_expr_ctxs(_output_vexpr_ctxs, &arrow_schema, - state->timezone(), - /*datetime_naive=*/true)); + RETURN_IF_ERROR(get_arrow_schema_from_expr_ctxs( + _output_vexpr_ctxs, &arrow_schema, state->timezone(), + /*datetime_naive=*/true, p._native_variant)); } VLOG_DEBUG << "create sender in INIT with instance id " << fragment_instance_id; RETURN_IF_ERROR(state->exec_env()->result_mgr()->create_sender( @@ -103,6 +103,7 @@ ResultSinkOperatorX::ResultSinkOperatorX(int operator_id, int node_id, _sink_type(!sink.__isset.type || sink.type == TResultSinkType::MYSQL_PROTOCOL ? TResultSinkType::MYSQL_PROTOCOL : sink.type), + _native_variant(sink.native_variant), _result_sink_buffer_size_rows(_sink_type == TResultSinkType::ARROW_FLIGHT_PROTOCOL ? config::arrow_flight_result_sink_buffer_size_rows : RESULT_SINK_BUFFER_SIZE), @@ -123,9 +124,9 @@ Status ResultSinkOperatorX::prepare(RuntimeState* state) { if (state->query_options().enable_parallel_result_sink) { std::shared_ptr arrow_schema; if (_sink_type == TResultSinkType::ARROW_FLIGHT_PROTOCOL) { - RETURN_IF_ERROR(get_arrow_schema_from_expr_ctxs(_output_vexpr_ctxs, &arrow_schema, - state->timezone(), - /*datetime_naive=*/true)); + RETURN_IF_ERROR(get_arrow_schema_from_expr_ctxs( + _output_vexpr_ctxs, &arrow_schema, state->timezone(), + /*datetime_naive=*/true, _native_variant)); } VLOG_DEBUG << "create sender in prepare with query id " << state->query_id(); RETURN_IF_ERROR(state->exec_env()->result_mgr()->create_sender( diff --git a/be/src/exec/operator/result_sink_operator.h b/be/src/exec/operator/result_sink_operator.h index 84c86f1127e319..e193d382b4f2db 100644 --- a/be/src/exec/operator/result_sink_operator.h +++ b/be/src/exec/operator/result_sink_operator.h @@ -167,6 +167,7 @@ class ResultSinkOperatorX final : public DataSinkOperatorX Status _second_phase_fetch_data(RuntimeState* state, Block* final_block); const TResultSinkType::type _sink_type; + const bool _native_variant; const int _result_sink_buffer_size_rows; // set file options when sink type is FILE std::unique_ptr _file_opts = nullptr; diff --git a/be/src/format/arrow/arrow_block_convertor.cpp b/be/src/format/arrow/arrow_block_convertor.cpp index 6b2ed5a71e8a83..d5c43f1b1778ce 100644 --- a/be/src/format/arrow/arrow_block_convertor.cpp +++ b/be/src/format/arrow/arrow_block_convertor.cpp @@ -472,6 +472,27 @@ Status ArrowBlockConvertor::init() { return Status::OK(); } +Status ArrowFlightArrowBlockConvertor::write_column(const std::shared_ptr& type, + const DataTypeSerDe& serde, + const IColumn& column, const NullMap* null_map, + const std::shared_ptr& 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 native_type; + RETURN_IF_ERROR(convert_to_arrow_type(type, &native_type, ctz.name(), true, true)); + // 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* 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 c923507786d3f9..93d4551a0d7bff 100644 --- a/be/src/format/arrow/arrow_block_convertor.h +++ b/be/src/format/arrow/arrow_block_convertor.h @@ -119,6 +119,13 @@ class ArrowFlightArrowBlockConvertor final : public DorisArrowBlockConvertor { Status convert_to_arrow(const Block& block, arrow::MemoryPool* pool, std::shared_ptr* result, size_t start_row = 0, size_t end_row = 0) const override; + +protected: + Status write_column(const std::shared_ptr& type, const DataTypeSerDe& serde, + const IColumn& column, const NullMap* null_map, + const std::shared_ptr& 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 f291cf33ac38db..c78adf586ebc96 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 +#include #include #include #include @@ -44,13 +45,34 @@ #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::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 convert_to_arrow_type(const DataTypePtr& origin_type, std::shared_ptr* result, const std::string& timezone, bool datetime_naive) { + return convert_to_arrow_type(origin_type, result, timezone, datetime_naive, false); +} + +Status convert_to_arrow_type(const DataTypePtr& origin_type, + std::shared_ptr* result, const std::string& timezone, + bool datetime_naive, bool native_variant) { auto type = get_serialized_type(origin_type); switch (type->get_primitive_type()) { case TYPE_NULL: @@ -137,7 +159,7 @@ Status convert_to_arrow_type(const DataTypePtr& origin_type, const auto* type_arr = assert_cast(remove_nullable(type).get()); std::shared_ptr item_type; RETURN_IF_ERROR(convert_to_arrow_type(type_arr->get_nested_type(), &item_type, timezone, - datetime_naive)); + datetime_naive, native_variant)); *result = std::make_shared(item_type); break; } @@ -146,9 +168,9 @@ Status convert_to_arrow_type(const DataTypePtr& origin_type, std::shared_ptr key_type; std::shared_ptr val_type; RETURN_IF_ERROR(convert_to_arrow_type(type_map->get_key_type(), &key_type, timezone, - datetime_naive)); + datetime_naive, native_variant)); RETURN_IF_ERROR(convert_to_arrow_type(type_map->get_value_type(), &val_type, timezone, - datetime_naive)); + datetime_naive, native_variant)); *result = std::make_shared(key_type, val_type); break; } @@ -158,7 +180,7 @@ Status convert_to_arrow_type(const DataTypePtr& origin_type, for (size_t i = 0; i < type_struct->get_elements().size(); i++) { std::shared_ptr field_type; RETURN_IF_ERROR(convert_to_arrow_type(type_struct->get_element(i), &field_type, - timezone, datetime_naive)); + timezone, datetime_naive, native_variant)); fields.push_back( std::make_shared(type_struct->get_element_name(i), field_type, type_struct->get_element(i)->is_nullable())); @@ -167,7 +189,13 @@ Status convert_to_arrow_type(const DataTypePtr& origin_type, break; } case TYPE_VARIANT: { - *result = arrow::utf8(); + if (native_variant) { + RETURN_IF_ERROR(register_arrow_variant_extension()); + } + *result = native_variant ? arrow::extension::variant(arrow::struct_( + {arrow::field("metadata", arrow::binary(), false), + arrow::field("value", arrow::binary(), false)})) + : arrow::utf8(); break; } case TYPE_QUANTILE_STATE: @@ -224,12 +252,20 @@ Status get_arrow_schema_from_block(const Block& block, std::shared_ptr* result, const std::string& timezone, bool datetime_naive) { + return get_arrow_schema_from_expr_ctxs(output_vexpr_ctxs, result, timezone, datetime_naive, + false); +} + +Status get_arrow_schema_from_expr_ctxs(const VExprContextSPtrs& output_vexpr_ctxs, + std::shared_ptr* result, + const std::string& timezone, bool datetime_naive, + bool native_variant) { std::vector> fields; for (int i = 0; i < output_vexpr_ctxs.size(); i++) { std::shared_ptr arrow_type; auto root_expr = output_vexpr_ctxs.at(i)->root(); RETURN_IF_ERROR(convert_to_arrow_type(root_expr->data_type(), &arrow_type, timezone, - datetime_naive)); + datetime_naive, native_variant)); auto field_name = root_expr->is_slot_ref() && !root_expr->expr_label().empty() ? root_expr->expr_label() : fmt::format("{}_{}", root_expr->data_type()->get_name(), i); @@ -286,13 +322,17 @@ Status serialize_record_batch(const arrow::RecordBatch& record_batch, std::strin } Status serialize_arrow_schema(std::shared_ptr* 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. + std::shared_ptr sink; + RETURN_DORIS_STATUS_IF_RESULT_ERROR(sink, arrow::io::BufferOutputStream::Create()); + std::shared_ptr writer; + RETURN_DORIS_STATUS_IF_RESULT_ERROR(writer, arrow::ipc::MakeStreamWriter(sink.get(), *schema)); + RETURN_DORIS_STATUS_IF_ERROR(writer->Close()); + std::shared_ptr 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 d5ba5cb0ed023c..9f47c25bd09e76 100644 --- a/be/src/format/arrow/arrow_row_batch.h +++ b/be/src/format/arrow/arrow_row_batch.h @@ -49,6 +49,12 @@ class RowDescriptor; Status convert_to_arrow_type(const DataTypePtr& type, std::shared_ptr* result, const std::string& timezone, bool datetime_naive = false); +Status register_arrow_variant_extension(); + +// Native Variant is opt-in for Flight; ordinary Arrow exports retain their existing types. +Status convert_to_arrow_type(const DataTypePtr& type, std::shared_ptr* result, + const std::string& timezone, bool datetime_naive, bool native_variant); + std::shared_ptr create_arrow_field_with_metadata( const std::string& field_name, const std::shared_ptr& arrow_type, bool is_nullable, PrimitiveType primitive_type); @@ -60,6 +66,11 @@ Status get_arrow_schema_from_expr_ctxs(const VExprContextSPtrs& output_vexpr_ctx std::shared_ptr* result, const std::string& timezone, bool datetime_naive = false); +Status get_arrow_schema_from_expr_ctxs(const VExprContextSPtrs& output_vexpr_ctxs, + std::shared_ptr* result, + const std::string& timezone, bool datetime_naive, + bool native_variant); + Status serialize_record_batch(const arrow::RecordBatch& record_batch, std::string* result); Status serialize_arrow_schema(std::shared_ptr* schema, std::string* result); 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 7e11fe9a4230fb..32c6ca53d327cd 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 00000000000000..6ec10b46406b30 --- /dev/null +++ b/be/test/format/arrow/arrow_flight_variant_test.cpp @@ -0,0 +1,1096 @@ +// 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 +#include +#include +#include +#include + +#include +#include + +#include "core/column/column_array.h" +#include "core/column/column_const.h" +#include "core/column/column_map.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_date_or_datetime_v2.h" +#include "core/data_type/data_type_decimal.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_time.h" +#include "core/data_type/data_type_varbinary.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 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(*column)) { + legacy->finalize(); + } + return column; +} + +VariantRef value_at(const arrow::Array& array, int row) { + const auto& storage = static_cast( + *static_cast(array).storage()); + auto metadata = static_cast(*storage.field(0)).GetView(row); + auto value = static_cast(*storage.field(1)).GetView(row); + return {{metadata.data(), metadata.size()}, {value.data(), value.size()}}; +} + +Status convert_legacy_root(MutableColumnPtr column, DataTypePtr type, bool nested, + std::shared_ptr* batch) { + if (nested) { + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(column->size()); + column = ColumnArray::create(make_nullable(std::move(column)), std::move(offsets)); + type = std::make_shared(type); + } + if (type->get_primitive_type() != TYPE_VARIANT) { + auto legacy = ColumnVariant::create(0); + legacy->create_root(type, std::move(column)); + column = std::move(legacy); + } + assert_cast(*column).finalize(); + Block block {{std::move(column), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("v", native_variant())}), + cctz::utc_time_zone()); + return converter.convert_to_arrow(block, arrow::default_memory_pool(), batch); +} + +TEST(ArrowFlightVariantTest, LegacyTimeRejectsDurationsOutsideDay) { + for (double micros : {-3600000000.0, -1.0, 0.0, 86399999999.0, 86400000000.0, 90000000000.0}) { + for (bool nested : {false, true}) { + auto times = ColumnTimeV2::create(); + times->insert_value(micros); + std::shared_ptr batch; + auto status = convert_legacy_root(std::move(times), std::make_shared(6), + nested, &batch); + if (micros < 0 || micros >= 86400000000.0) { + EXPECT_FALSE(status.ok()); + EXPECT_NE(status.to_string().find("enable_arrow_flight_sql_native_variant=false"), + std::string::npos); + } else { + ASSERT_TRUE(status.ok()) << status; + auto value = value_at(*batch->column(0), 0); + if (nested) { + value = value.array_at(0); + } + EXPECT_EQ(value.get_time_ntz_micros(), static_cast(micros)); + } + } + } +} + +TEST(ArrowFlightVariantTest, LegacyMapRejectsNullKeysWithoutCollidingWithStrings) { + for (bool null_key : {false, true}) { + for (bool nested : {false, true}) { + auto keys = ColumnString::create(); + keys->insert_data("key", 3); + keys->insert_data("null", 4); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({static_cast(null_key), 0}); + auto values = ColumnInt32::create(); + values->get_data().assign({1, 2}); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(2); + auto map = ColumnMap::create(ColumnNullable::create(std::move(keys), std::move(nulls)), + make_nullable(std::move(values)), std::move(offsets)); + auto type = + std::make_shared(make_nullable(std::make_shared()), + make_nullable(std::make_shared())); + std::shared_ptr batch; + auto status = convert_legacy_root(std::move(map), type, nested, &batch); + if (null_key) { + EXPECT_FALSE(status.ok()); + EXPECT_NE(status.to_string().find("MAP with NULL keys"), std::string::npos); + EXPECT_NE(status.to_string().find("enable_arrow_flight_sql_native_variant=false"), + std::string::npos); + } else { + ASSERT_TRUE(status.ok()) << status; + auto value = value_at(*batch->column(0), 0); + if (nested) { + value = value.array_at(0); + } + VariantRef child; + ASSERT_TRUE(value.object_find({"null", 4}, &child)); + EXPECT_EQ(child.get_int(), 2); + } + } + } +} + +TEST(ArrowFlightVariantTest, LegacyBinaryPreservesBytesInRootsAndArrays) { + for (bool nested : {false, true}) { + auto type = std::make_shared(); + auto values = type->create_column(); + const std::string bytes( + "\0\xff\x80" + "42", + 5); + values->insert_data(bytes.data(), bytes.size()); + values->insert_data("", 0); + std::shared_ptr batch; + auto status = convert_legacy_root(std::move(values), type, nested, &batch); + ASSERT_TRUE(status.ok()) << status; + for (int i = 0; i < 2; ++i) { + auto value = nested ? value_at(*batch->column(0), 0).array_at(i) + : value_at(*batch->column(0), i); + EXPECT_EQ(value.get_binary().to_string(), i == 0 ? bytes : ""); + } + } +} + +TEST(ArrowFlightVariantTest, LegacyDocumentsPreserveTypedPaths) { + // Dense paths, sparse paths and document snapshots must all retain typed values. + for (int storage = 0; storage < 3; ++storage) { + for (bool nested : {false, true}) { + SCOPED_TRACE(::testing::Message() << "storage=" << storage << " nested=" << nested); + auto legacy = ColumnVariant::create(0); + auto root_type = make_nullable(std::make_shared()); + auto roots = root_type->create_column(); + roots->insert_default(); + legacy->create_root(root_type, std::move(roots)); + const Int128 exact = 900719925474099301LL; + auto decimal_type = make_nullable(std::make_shared(20, 2)); + auto decimals = ColumnDecimal128V3::create(0, 2); + decimals->insert_value(Decimal128V3(exact)); + auto decimal_column = make_nullable(std::move(decimals)); + auto date_type = make_nullable(std::make_shared()); + auto dates = date_type->create_column(); + DateV2Value date; + date.unchecked_set_time(2020, 1, 2, 0, 0, 0); + dates->insert(Field::create_field(date)); + if (storage == 0) { + ASSERT_TRUE(legacy->add_sub_column(PathInData("nested.amount"), + decimal_column->assert_mutable(), decimal_type)); + ASSERT_TRUE(legacy->add_sub_column(PathInData("nested.date"), std::move(dates), + date_type)); + } else { + auto& map = assert_cast( + storage == 1 ? legacy->get_sparse_column_mutable() + : legacy->get_doc_value_column_mutable()); + auto& keys = assert_cast(map.get_keys()); + auto& values = assert_cast(map.get_values()); + ColumnVariant::Subcolumn decimal(decimal_column->assert_mutable(), decimal_type, + true); + ColumnVariant::Subcolumn date_column(std::move(dates), date_type, true); + decimal.serialize_to_binary_column(&keys, "nested.amount", &values, 0); + date_column.serialize_to_binary_column(&keys, "nested.date", &values, 0); + map.get_offsets()[0] = 2; + } + legacy->finalize(); + std::shared_ptr batch; + auto status = convert_legacy_root(std::move(legacy), + std::make_shared(), nested, &batch); + ASSERT_TRUE(status.ok()) << status; + auto value = value_at(*batch->column(0), 0); + if (nested) { + value = value.array_at(0); + } + VariantRef object; + ASSERT_TRUE(value.object_find({"nested", 6}, &object)); + VariantRef amount; + ASSERT_TRUE(object.object_find({"amount", 6}, &amount)); + ASSERT_EQ(amount.primitive_id(), VariantPrimitiveId::DECIMAL16); + EXPECT_EQ(amount.get_decimal().unscaled, exact); + EXPECT_EQ(amount.get_decimal().scale, 2); + VariantRef day; + ASSERT_TRUE(object.object_find({"date", 4}, &day)); + EXPECT_EQ(day.primitive_id(), VariantPrimitiveId::DATE); + EXPECT_EQ(day.get_date(), 18263); + } + } +} + +TEST(ArrowFlightVariantTest, LegacyPendingDocumentDefaultsPreserveValuesAndSlices) { + auto legacy = ColumnVariant::create(0); + auto root_type = make_nullable(std::make_shared()); + auto roots = root_type->create_column(); + roots->insert_many_defaults(3); + legacy->create_root(root_type, std::move(roots)); + auto decimal_type = std::make_shared(20, 2); + auto decimals = ColumnDecimal128V3::create(0, 2); + const Int128 exact = 900719925474099301LL; + decimals->insert_value(Decimal128V3(exact)); + ASSERT_TRUE(legacy->add_sub_column(PathInData("amount"), 3)); + auto* amount = legacy->get_subcolumn(PathInData("amount")); + *amount = ColumnVariant::Subcolumn(1, true); + amount->insert(decimal_type->get_field_with_data_type(*decimals, 0)); + amount->insert_default(); + // Scans can return lazy prefix/suffix defaults; output must not finalize the shared input. + ASSERT_FALSE(amount->is_finalized()); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({0, 0, 1}); + Block block {{ColumnNullable::create(std::move(legacy), std::move(nulls)), + make_nullable(std::make_shared()), "v"}}; + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("v", native_variant())}), + cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(value_at(*batch->column(0), 0).num_elements(), 0); + VariantRef value; + ASSERT_TRUE(value_at(*batch->column(0), 1).object_find({"amount", 6}, &value)); + EXPECT_EQ(value.get_decimal().unscaled, exact); + EXPECT_EQ(value.get_decimal().scale, 2); + EXPECT_TRUE(batch->column(0)->IsNull(2)); + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 3).ok()); + ASSERT_TRUE(value_at(*batch->column(0), 0).object_find({"amount", 6}, &value)); + EXPECT_EQ(value.get_decimal().unscaled, exact); + EXPECT_TRUE(batch->column(0)->IsNull(1)); + EXPECT_FALSE(amount->is_finalized()); +} + +TEST(ArrowFlightVariantTest, LegacyScanDocumentWithPendingArrayDefaults) { + auto type = std::make_shared(9); + auto column = type->create_column(); + DataTypeSerDe::FormatOptions options; + for (std::string json : {"42", R"("text")", R"({"a":[1,null,"x"]})", "null"}) { + Slice slice(json.data(), json.size()); + ASSERT_TRUE( + type->get_serde()->deserialize_one_cell_from_json(*column, slice, options).ok()); + } + auto& legacy = assert_cast(*column); + legacy.get_subcolumns().get_mutable_root()->data.finalize(); + auto* array = legacy.get_subcolumn(PathInData("a")); + ASSERT_NE(array, nullptr); + ASSERT_FALSE(array->is_finalized()); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({0, 0, 0, 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()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(value_at(*batch->column(0), 0).get_int(), 42); + EXPECT_EQ(value_at(*batch->column(0), 1).get_string().to_string(), "text"); + VariantRef value; + ASSERT_TRUE(value_at(*batch->column(0), 2).object_find({"a", 1}, &value)); + ASSERT_EQ(value.num_elements(), 3); + EXPECT_EQ(value.array_at(0).get_int(), 1); + EXPECT_TRUE(value.array_at(1).is_null()); + EXPECT_EQ(value.array_at(2).get_string().to_string(), "x"); + EXPECT_TRUE(batch->column(0)->IsNull(3)); + EXPECT_FALSE(array->is_finalized()); +} + +TEST(ArrowFlightVariantTest, NativeResultPreservesValuesAndSqlNulls) { + TimezoneUtils::load_timezones_to_cache(); + ASSERT_TRUE(register_arrow_variant_extension().ok()); + for (DataTypePtr type : {DataTypePtr(std::make_shared()), + DataTypePtr(std::make_shared())}) { + 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 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)); + // Legacy Variant already renders a null root as {}; only V2 retains Variant null. + EXPECT_EQ(value_at(*batch->column(0), 1).is_null(), + dynamic_cast(type.get()) != nullptr); + 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(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_EQ(value_at(*batch->column(0), 0).is_null(), + dynamic_cast(type.get()) != nullptr); + EXPECT_EQ(value_at(*batch->column(0), 1).get_int(), 42); + } +} + +TEST(ArrowFlightVariantTest, SchemaMappingAndConstantScalar) { + for (DataTypePtr type : {DataTypePtr(std::make_shared()), + DataTypePtr(std::make_shared())}) { + std::shared_ptr mapped; + ASSERT_TRUE(convert_to_arrow_type(type, &mapped, "UTC", true).ok()); + EXPECT_TRUE(mapped->Equals(arrow::utf8())); + ASSERT_TRUE(convert_to_arrow_type(type, &mapped, "UTC", true, true).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()); + if (auto* legacy = check_and_get_column(*column)) { + legacy->finalize(); + } + 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 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, TypedV2AndNestedStructPreserveNonJsonNumbers) { + auto numbers = ColumnFloat64::create(); + numbers->insert_value(std::numeric_limits::quiet_NaN()); + numbers->insert_value(std::numeric_limits::infinity()); + auto values = ColumnVariantV2::create_typed(make_nullable(std::move(numbers)), + std::make_shared()); + auto variant = std::make_shared(); + auto type = std::make_shared(DataTypes {variant}, Strings {"v"}); + Block block; + block.insert({ColumnStruct::create(Columns {std::move(values)}), type, "s"}); + std::shared_ptr mapped; + ASSERT_TRUE(convert_to_arrow_type(type, &mapped, "UTC", true, true).ok()); + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("s", mapped, false)}), + cctz::utc_time_zone()); + std::shared_ptr 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(*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, LegacyTypedDecimalDoesNotRoundThroughDouble) { + auto decimal_type = std::make_shared(20, 2); + auto decimals = ColumnDecimal128V3::create(0, 2); + const __int128 unscaled = 900719925474099301LL; + decimals->insert_value(Decimal128V3(unscaled)); + auto values = ColumnVariant::create(0); + values->create_root(decimal_type, std::move(decimals)); + values->finalize(); + auto type = std::make_shared(); + Block block; + block.insert({std::move(values), type, "v"}); + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + auto decimal = value_at(*batch->column(0), 0).get_decimal(); + EXPECT_EQ(decimal.unscaled, unscaled); + EXPECT_EQ(decimal.scale, 2); +} + +TEST(ArrowFlightVariantTest, LegacyArrayPreservesExactDecimal) { + auto decimal_type = std::make_shared(20, 2); + auto decimals = ColumnDecimal128V3::create(0, 2); + const __int128 unscaled = 900719925474099301LL; + decimals->insert_value(Decimal128V3(unscaled)); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(1); + auto values = ColumnVariant::create(0); + values->create_root( + std::make_shared(decimal_type), + ColumnArray::create(make_nullable(std::move(decimals)), std::move(offsets))); + values->finalize(); + Block block {{std::move(values), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + auto element = value_at(*batch->column(0), 0).array_at(0); + ASSERT_EQ(element.primitive_id(), VariantPrimitiveId::DECIMAL16); + EXPECT_EQ(element.get_decimal().unscaled, unscaled); + EXPECT_EQ(element.get_decimal().scale, 2); +} + +TEST(ArrowFlightVariantTest, LegacyArrayPreservesNonFiniteNumbers) { + auto numbers = ColumnFloat64::create(); + numbers->insert_value(std::numeric_limits::quiet_NaN()); + numbers->insert_value(std::numeric_limits::infinity()); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(2); + auto values = ColumnVariant::create(0); + values->create_root(std::make_shared(std::make_shared()), + ColumnArray::create(make_nullable(std::move(numbers)), std::move(offsets))); + values->finalize(); + Block block {{std::move(values), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + auto array = value_at(*batch->column(0), 0); + EXPECT_TRUE(std::isnan(array.array_at(0).get_double())); + EXPECT_TRUE(std::isinf(array.array_at(1).get_double())); +} + +TEST(ArrowFlightVariantTest, MixedDateRootRetainsDateIdentityAndSlices) { + auto date_type = make_nullable(std::make_shared()); + auto dates = date_type->create_column(); + DateV2Value date; + date.unchecked_set_time(2020, 1, 2, 0, 0, 0); + dates->insert(Field::create_field(date)); + dates->insert_default(); + auto values = ColumnVariant::create(0); + values->create_root(date_type, std::move(dates)); + auto object_type = make_nullable(std::make_shared()); + auto object = object_type->create_column(); + object->insert_default(); + object->insert(Field::create_field(7)); + ASSERT_TRUE(values->add_sub_column(PathInData("a"), std::move(object), object_type)); + values->finalize(); + ASSERT_FALSE(values->is_scalar_variant()); + Block block {{std::move(values), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(value_at(*batch->column(0), 0).get_date(), 18263); + EXPECT_EQ(value_at(*batch->column(0), 1).basic_type(), VariantBasicType::OBJECT); + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 2).ok()); + EXPECT_EQ(value_at(*batch->column(0), 0).basic_type(), VariantBasicType::OBJECT); +} + +TEST(ArrowFlightVariantTest, MixedStringRootsKeepStringIdentity) { + for (auto primitive : {TYPE_CHAR, TYPE_VARCHAR, TYPE_STRING}) { + auto type = std::make_shared(-1, primitive); + auto root = type->create_column(); + for (const auto& text : {"hello", "true", "42", ""}) { + root->insert_data(text, strlen(text)); + } + auto root_nulls = ColumnUInt8::create(); + root_nulls->get_data().assign({0, 0, 0, 1}); + auto values = ColumnVariant::create(0); + values->create_root(make_nullable(type), + ColumnNullable::create(std::move(root), std::move(root_nulls))); + auto object_type = make_nullable(std::make_shared()); + auto object = object_type->create_column(); + object->insert_default(); + object->insert_default(); + object->insert_default(); + object->insert(Field::create_field(7)); + ASSERT_TRUE(values->add_sub_column(PathInData("a"), std::move(object), object_type)); + values->finalize(); + ASSERT_FALSE(values->is_scalar_variant()); + Block block {{std::move(values), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(value_at(*batch->column(0), 0).get_string().to_string(), "hello"); + EXPECT_EQ(value_at(*batch->column(0), 1).get_string().to_string(), "true"); + EXPECT_EQ(value_at(*batch->column(0), 2).get_string().to_string(), "42"); + EXPECT_EQ(value_at(*batch->column(0), 3).basic_type(), VariantBasicType::OBJECT); + } +} + +TEST(ArrowFlightVariantTest, NestedTimezoneAliasesMatchPublishedSchema) { + TimezoneUtils::load_timezones_to_cache(); + auto variant = std::make_shared(); + auto timestamp = DataTypeFactory::instance().create_data_type(TYPE_TIMESTAMPTZ, false, 0, 6); + auto type = + std::make_shared(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 mapped; + ASSERT_TRUE(convert_to_arrow_type(type, &mapped, zone, true, true).ok()); + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("s", mapped, false)}), + timezone); + std::shared_ptr 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 (bool legacy : {false, true}) { + for (int rows : {32, 64}) { + SCOPED_TRACE(::testing::Message() << "legacy=" << legacy << " rows=" << rows); + DataTypePtr type = legacy ? DataTypePtr(std::make_shared()) + : DataTypePtr(std::make_shared()); + auto column = type->create_column(); + JsonStringToVariantEncoder encoder; + DataTypeSerDe::FormatOptions options; + 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) + "}"; + if (legacy) { + Slice slice(json.data(), json.size()); + ASSERT_TRUE(type->get_serde() + ->deserialize_one_cell_from_json(*column, slice, options) + .ok()); + } else { + encoder.add_json({json.data(), json.size()}); + } + } + if (legacy) { + assert_cast(*column).finalize(); + } else { + auto encoded = encoder.finish_batch(); + ASSERT_EQ(encoded.metadata_ref().dict_size(), rows); + assert_cast(*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 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(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(std::make_shared()); + 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 batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + const auto& elements = *static_cast(*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(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + std::shared_ptr 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, LegacyV2LeafIncludesEnclosingDepth) { + for (const std::string terminal : {"1", "[]", "{}"}) { + for (size_t depth : {VARIANT_MAX_NESTING_DEPTH - 1, VARIANT_MAX_NESTING_DEPTH}) { + std::string json = terminal; + for (size_t i = 0; i < depth; ++i) { + json = "[" + json + "]"; + } + JsonStringToVariantEncoder encoder; + encoder.add_json({json.data(), json.size()}); + auto encoded = encoder.finish_batch(); + auto values = ColumnVariantV2::create(); + values->insert_encoded_batch(encoded); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(1); + auto legacy = ColumnVariant::create(0); + legacy->create_root( + std::make_shared(std::make_shared()), + ColumnArray::create(make_nullable(std::move(values)), std::move(offsets))); + legacy->finalize(); + Block block {{std::move(legacy), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor utf8(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(utf8.init().ok()); + std::shared_ptr batch; + ASSERT_TRUE(utf8.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + EXPECT_EQ(static_cast(*batch->column(0)).GetString(0), + "[" + json + "]"); + ArrowFlightArrowBlockConvertor native( + arrow::schema({arrow::field("v", native_variant(), false)}), + cctz::utc_time_zone()); + auto status = native.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + // A valid V2 leaf may exceed the native limit once its legacy container is included. + if (depth == VARIANT_MAX_NESTING_DEPTH) { + EXPECT_EQ(status.code(), ErrorCode::NOT_IMPLEMENTED_ERROR) << status; + EXPECT_NE(status.to_string().find("enable_arrow_flight_sql_native_variant=false"), + std::string::npos); + } else { + ASSERT_TRUE(status.ok()) << status; + auto actual = value_at(*batch->column(0), 0); + for (size_t i = 0; i <= depth; ++i) { + actual = actual.array_at(0); + } + EXPECT_EQ(actual.basic_type(), terminal == "1" ? VariantBasicType::PRIMITIVE + : terminal == "[]" ? VariantBasicType::ARRAY + : VariantBasicType::OBJECT); + } + } + } +} + +TEST(ArrowFlightVariantTest, LegacyArrayV2LeavesShareLargeDictionary) { + constexpr int rows = 4096; + VariantBatchBuilder encoder; + for (int i = 0; i < rows; ++i) { + auto row = encoder.begin_row(); + auto object = row.start_object(); + auto key = "key_" + std::to_string(i); + object.add_key({key.data(), key.size()}); + row.add_int(i); + object.finish(); + row.finish(); + } + auto encoded = encoder.finish_batch(); + ASSERT_EQ(encoded.metadata_ref().dict_size(), rows); + auto values = ColumnVariantV2::create(); + values->insert_encoded_batch(encoded); + std::shared_ptr batch; + ASSERT_TRUE(convert_legacy_root(std::move(values), std::make_shared(), true, + &batch) + .ok()); + auto actual = value_at(*batch->column(0), 0); + ASSERT_EQ(actual.num_elements(), rows); + EXPECT_EQ(actual.metadata.dict_size(), rows); + for (int i = 0; i < rows; ++i) { + auto object = actual.array_at(i).object_view(); + ASSERT_EQ(object.size(), 1); + uint32_t field_id; + auto child = object.value_at(0, &field_id); + EXPECT_EQ(actual.metadata.key_at(field_id).to_string(), "key_" + std::to_string(i)); + EXPECT_EQ(child.get_int(), i); + } +} + +TEST(ArrowFlightVariantTest, LegacyDepthLimitExplainsNativeModeRestriction) { + std::string json = "1"; + for (int i = 0; i < 129; ++i) { + json = R"({"a":)" + json + "}"; + } + auto type = std::make_shared(); + auto values = type->create_column(); + Slice slice(json.data(), json.size()); + DataTypeSerDe::FormatOptions options; + ASSERT_TRUE(type->get_serde()->deserialize_one_cell_from_json(*values, slice, options).ok()); + assert_cast(*values).finalize(); + Block block {{std::move(values), type, "v"}}; + ArrowFlightArrowBlockConvertor utf8(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(utf8.init().ok()); + std::shared_ptr batch; + ASSERT_TRUE(utf8.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + ArrowFlightArrowBlockConvertor native( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + auto status = native.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_FALSE(status.ok()); + EXPECT_NE(status.to_string().find("enable_arrow_flight_sql_native_variant=false"), + std::string::npos); +} + +TEST(ArrowFlightVariantTest, LegacyJsonbLeafIncludesEnclosingDepth) { + for (int outer_depth : {28, 29}) { + std::string json = "1"; + for (int i = 0; i < 99; ++i) { + json = R"({"a":)" + json + "}"; + } + json = "[" + json + "]"; + for (int i = 0; i < outer_depth; ++i) { + json = R"({"a":)" + json + "}"; + } + auto type = std::make_shared(); + auto values = type->create_column(); + Slice slice(json.data(), json.size()); + DataTypeSerDe::FormatOptions options; + ASSERT_TRUE( + type->get_serde()->deserialize_one_cell_from_json(*values, slice, options).ok()); + assert_cast(*values).finalize(); + Block block {{std::move(values), type, "v"}}; + ArrowFlightArrowBlockConvertor utf8(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(utf8.init().ok()); + std::shared_ptr batch; + ASSERT_TRUE(utf8.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + ArrowFlightArrowBlockConvertor native( + arrow::schema({arrow::field("v", native_variant(), false)}), cctz::utc_time_zone()); + auto status = native.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + // The leaf fits JSONB's limit, but its enclosing paths count toward Variant's limit. + if (outer_depth == 28) { + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(value_at(*batch->column(0), 0).basic_type(), VariantBasicType::OBJECT); + } else { + EXPECT_FALSE(status.ok()); + EXPECT_NE(status.to_string().find("enable_arrow_flight_sql_native_variant=false"), + std::string::npos); + } + } +} + +TEST(ArrowFlightVariantTest, LegacyDecimal256RespectsOuterSqlNulls) { + auto decimal_type = std::make_shared(76, 2); + auto decimals = decimal_type->create_column(); + for (int i = 0; i < 3; ++i) { + decimals->insert_default(); + } + auto values = ColumnVariant::create(0); + values->create_root(decimal_type, std::move(decimals)); + values->finalize(); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({1, 1, 1}); + Block block {{ColumnNullable::create(std::move(values), std::move(nulls)), + make_nullable(std::make_shared()), "v"}}; + ArrowFlightArrowBlockConvertor converter(arrow::schema({arrow::field("v", native_variant())}), + cctz::utc_time_zone()); + std::shared_ptr batch; + // A masked physical decimal is not a value to encode, even for scalar-only batches. + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(batch->column(0)->null_count(), 3); + auto& nullable = + assert_cast(*block.get_by_position(0).column->assert_mutable()); + nullable.get_null_map_data()[0] = 0; + status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 3); + ASSERT_TRUE(status.ok()) << status; + EXPECT_EQ(batch->column(0)->null_count(), 2); + status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 0, 1); + EXPECT_FALSE(status.ok()); + EXPECT_NE(status.to_string().find(decimal_type->get_name()), std::string::npos); +} + +TEST(ArrowFlightVariantTest, LegacyScalarNullKeepsEmptyObjectAndSqlNull) { + auto type = std::make_shared(); + auto column = type->create_column(); + DataTypeSerDe::FormatOptions options; + for (std::string json : {"42", "null", "null"}) { + Slice slice(json.data(), json.size()); + ASSERT_TRUE( + type->get_serde()->deserialize_one_cell_from_json(*column, slice, options).ok()); + } + auto& legacy = assert_cast(*column); + legacy.finalize(); + ASSERT_TRUE(legacy.is_scalar_variant()); + std::string text; + legacy.serialize_one_row_to_string(1, &text, options); + ASSERT_EQ(text, "{}"); + auto nulls = ColumnUInt8::create(); + nulls->get_data().assign({0, 0, 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()); + std::shared_ptr batch; + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + EXPECT_EQ(value_at(*batch->column(0), 0).get_int(), 42); + EXPECT_EQ(value_at(*batch->column(0), 1).basic_type(), VariantBasicType::OBJECT); + EXPECT_TRUE(batch->column(0)->IsNull(2)); + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 3).ok()); + EXPECT_EQ(value_at(*batch->column(0), 0).basic_type(), VariantBasicType::OBJECT); + EXPECT_TRUE(batch->column(0)->IsNull(1)); +} + +TEST(ArrowFlightVariantTest, LegacyCompositeRootsAndNestedLeaves) { + const __int128 exact = 900719925474099301LL; + for (int family = 0; family < 5; ++family) { + for (int depth = family == 3 ? 1 : 0; depth < 3; ++depth) { + SCOPED_TRACE(::testing::Message() << "family=" << family << " depth=" << depth); + DataTypePtr type; + MutableColumnPtr column; + if (family <= 1 || family == 4) { + auto decimal_type = std::make_shared(20, 2); + auto decimals = ColumnDecimal128V3::create(0, 2); + decimals->insert_value(Decimal128V3(exact)); + if (family != 1) { + DataTypePtr key_type = family == 0 + ? DataTypePtr(std::make_shared()) + : DataTypePtr(std::make_shared()); + auto keys = key_type->create_column(); + if (family == 0) { + keys->insert_data("d", 1); + } else { + keys->insert(Field::create_field(-7)); + } + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(1); + type = std::make_shared(make_nullable(key_type), + make_nullable(decimal_type)); + column = ColumnMap::create(make_nullable(std::move(keys)), + make_nullable(std::move(decimals)), + std::move(offsets)); + } else { + type = std::make_shared(DataTypes {decimal_type}, + Strings {"d"}); + column = ColumnStruct::create(Columns {std::move(decimals)}); + } + } else if (family == 2) { + type = std::make_shared(6); + auto times = ColumnTimeV2::create(); + times->insert_value(3'723'123'456.0); + column = std::move(times); + } else { + type = std::make_shared(); + column = documents(type)->cut(0, 1)->assert_mutable(); + } + for (int level = 0; level < depth; ++level) { + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().push_back(1); + column = ColumnArray::create(make_nullable(std::move(column)), std::move(offsets)); + type = std::make_shared(type); + } + auto variant = ColumnVariant::create(0); + variant->create_root(type, std::move(column)); + variant->finalize(); + Block block {{std::move(variant), std::make_shared(), "v"}}; + ArrowFlightArrowBlockConvertor converter( + arrow::schema({arrow::field("v", native_variant(), false)}), + cctz::utc_time_zone()); + std::shared_ptr batch; + auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_TRUE(status.ok()) << status; + if (!status.ok()) { + continue; + } + auto value = value_at(*batch->column(0), 0); + for (int level = 0; level < depth; ++level) { + value = value.array_at(0); + } + if (family <= 1 || family == 4) { + VariantRef decimal; + const std::string key = family == 4 ? "-7" : "d"; + ASSERT_TRUE(value.object_find({key.data(), key.size()}, &decimal)); + EXPECT_EQ(decimal.get_decimal().unscaled, exact); + EXPECT_EQ(decimal.get_decimal().scale, 2); + } else if (family == 2) { + EXPECT_EQ(value.get_time_ntz_micros(), 3'723'123'456LL); + } else { + VariantRef array; + ASSERT_TRUE(value.object_find({"a", 1}, &array)); + EXPECT_EQ(array.array_at(0).get_int(), 1); + EXPECT_TRUE(array.array_at(1).is_null()); + EXPECT_EQ(array.array_at(2).get_string().to_string(), "x"); + } + } + } +} + +TEST(ArrowFlightVariantTest, SchemaRpcPreservesNestedVariantExtensions) { + ASSERT_TRUE(register_arrow_variant_extension().ok()); + for (const auto& type : std::vector> { + 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) { + for (DataTypePtr type : {DataTypePtr(std::make_shared()), + DataTypePtr(std::make_shared())}) { + 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 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(); + auto offsets = ColumnArray::ColumnOffsets::create(); + offsets->get_data().assign({2, 4}); + auto array_type = std::make_shared(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 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(*batch->column(0)).values(); + EXPECT_EQ(value_at(values, 2).get_int(), 42); + EXPECT_EQ(value_at(values, 3).get_string().to_string(), "text"); + + ArrowFlightArrowBlockConvertor 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(*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/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 25e72ed7598621..fc50079f9d7042 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 @@ -17,6 +17,8 @@ package org.apache.doris.planner; +import org.apache.doris.qe.ConnectContext; +import org.apache.doris.service.arrowflight.FlightSqlNativeVariant; import org.apache.doris.thrift.TDataSink; import org.apache.doris.thrift.TDataSinkType; import org.apache.doris.thrift.TExplainLevel; @@ -34,15 +36,20 @@ public class ResultSink extends DataSink { // Two phase fetch option private TFetchOption fetchOption; + private final boolean nativeVariant; private TResultSinkType resultSinkType = TResultSinkType.MYSQL_PROTOCOL; public ResultSink(PlanNodeId exchNodeId) { - this.exchNodeId = exchNodeId; + this(exchNodeId, TResultSinkType.MYSQL_PROTOCOL); } public ResultSink(PlanNodeId exchNodeId, TResultSinkType resultSinkType) { this.exchNodeId = exchNodeId; this.resultSinkType = resultSinkType; + ConnectContext context = ConnectContext.get(); + // The session may change before deferred result fetching; pin the format during planning. + nativeVariant = resultSinkType == TResultSinkType.ARROW_FLIGHT_PROTOCOL + && FlightSqlNativeVariant.isEnabled(context); } @Override @@ -73,6 +80,7 @@ protected TDataSink toThrift() { tResultSink.setFetchOption(fetchOption); } tResultSink.setType(resultSinkType); + tResultSink.setNativeVariant(nativeVariant); result.setResultSink(tResultSink); return result; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java index f6eaf55e5f13c5..baaeb4276de845 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java @@ -275,6 +275,7 @@ public class SessionVariable implements Serializable, Writable { public static final String RUNTIME_FILTER_TREE_PUBLISH_MAX_SEND_BYTES = "runtime_filter_tree_publish_max_send_bytes"; + public static final String ENABLE_ARROW_FLIGHT_SQL_NATIVE_VARIANT = "enable_arrow_flight_sql_native_variant"; public static final String ENABLE_PARALLEL_RESULT_SINK = "enable_parallel_result_sink"; public static final String HIVE_TEXT_COMPRESSION = "hive_text_compression"; @@ -1838,6 +1839,9 @@ public enum IgnoreSplitType { @VariableMgr.VarAttr(name = "runtime_filter_max_build_row_count", needForward = true, fuzzy = false) public long runtimeFilterMaxBuildRowCount = 64L * 1024L * 1024L; + @VariableMgr.VarAttr(name = ENABLE_ARROW_FLIGHT_SQL_NATIVE_VARIANT, needForward = true) + private boolean enableArrowFlightSqlNativeVariant = false; + @VariableMgr.VarAttr(name = ENABLE_PARALLEL_RESULT_SINK, needForward = true, fuzzy = true) private boolean enableParallelResultSink = true; @@ -6291,6 +6295,14 @@ public boolean getEnableAggregateFunctionNullV2() { return enableAggregateFunctionNullV2; } + public boolean isEnableArrowFlightSqlNativeVariant() { + return enableArrowFlightSqlNativeVariant; + } + + public void setEnableArrowFlightSqlNativeVariant(boolean enabled) { + enableArrowFlightSqlNativeVariant = enabled; + } + public boolean enableParallelResultSink() { return enableParallelResultSink; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariant.java b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariant.java new file mode 100644 index 00000000000000..14fd06d70ae835 --- /dev/null +++ b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariant.java @@ -0,0 +1,51 @@ +// 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.Env; +import org.apache.doris.common.AnalysisException; +import org.apache.doris.qe.ConnectContext; +import org.apache.doris.system.Backend; + +import java.util.Collection; + +public final class FlightSqlNativeVariant { + private FlightSqlNativeVariant() { + } + + public static boolean isEnabled(ConnectContext context) { + if (context == null) { + return false; + } + // Schema analysis temporarily installs SET_VAR state under this monitor. Metadata + // requests must wait for that scope to end instead of observing another query's hints. + synchronized (context) { + if (!context.getSessionVariable().isEnableArrowFlightSqlNativeVariant()) { + return false; + } + } + try { + Collection backends = Env.getCurrentSystemInfo().getAllBackendsByAllCluster().values(); + // A Flight ticket may be proxied through a BE outside the query's result sinks. + // Require every registered BE, including unknown heartbeat capabilities, to support the format. + return !backends.isEmpty() && backends.stream().allMatch(Backend::isArrowFlightNativeVariantSupported); + } catch (AnalysisException e) { + return false; + } + } +} 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 e0f892176904ea..40ff5340f90fc1 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 @@ -194,9 +194,11 @@ static Schema analyze(ConnectContext context, String query) throws Exception { Plan analyzed = cascades.getRewritePlan(); // PrepareCommandPlanner stops before the rewrite phase that normally checks privileges. new CheckPrivileges().rewriteRoot(analyzed, cascades.getCurrentJobContext()); + // Prepare/GetSchema must use the sink capability gate, including scoped SET_VAR hints. + boolean nativeVariant = FlightSqlNativeVariant.isEnabled(context); for (Slot slot : analyzed.getOutput()) { fields.add(field(slot.getName(), slot.getDataType().toCatalogDataType(), slot.nullable(), - true, context.getSessionVariable().getTimeZone())); + true, context.getSessionVariable().getTimeZone(), nativeVariant)); } } return new Schema(fields); @@ -371,13 +373,17 @@ private static void resolveNamespace(ConnectContext context, Plan plan, Map children = new ArrayList<>(); if (type instanceof ArrayType) { // BE constructs ListType and MapType from data types, so item/value fields are nullable. - children.add(field("item", ((ArrayType) type).getItemType(), true, false, timezone)); + children.add(field("item", ((ArrayType) type).getItemType(), true, false, timezone, nativeVariant)); } else if (type instanceof MapType) { MapType map = (MapType) type; children.add(new Field("entries", FieldType.notNullable(new ArrowType.Struct()), Arrays.asList( - field("key", map.getKeyType(), false, false, timezone), - field("value", map.getValueType(), true, false, timezone)))); + field("key", map.getKeyType(), false, false, timezone, nativeVariant), + field("value", map.getValueType(), true, false, timezone, nativeVariant)))); } 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, nativeVariant)); } } Map metadata = null; 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 9d91153ebd6bcd..b2da760c073204 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 @@ -31,6 +31,7 @@ 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.sql.FlightSqlColumnMetadata; @@ -280,6 +281,7 @@ private TDescribeTablesResult describeTables(String dbName, String catalogName, private Map> buildTableToFields(String dbName, TDescribeTablesResult describeTablesResult, List tablesName) { Map> tableToFields = new HashMap<>(); + boolean nativeVariant = FlightSqlNativeVariant.isEnabled(ctx); int columnIndex = 0; for (int tableIndex = 0; tableIndex < describeTablesResult.getTablesOffsetSize(); tableIndex++) { String tableName = tablesName.get(tableIndex); @@ -287,7 +289,7 @@ private Map> buildTableToFields(String dbName, TDescribeTabl Integer tableOffset = describeTablesResult.getTablesOffset().get(tableIndex); for (; columnIndex < tableOffset; columnIndex++) { TColumnDef columnDef = describeTablesResult.getColumns().get(columnIndex); - fields.add(buildField(dbName, tableName, columnDef.getColumnDesc())); + fields.add(buildField(dbName, tableName, columnDef.getColumnDesc(), nativeVariant)); } tableToFields.put(tableName, fields); } @@ -296,11 +298,29 @@ private Map> buildTableToFields(String dbName, TDescribeTabl /** One column, with its nested types described down to the leaves. */ private static Field buildField(String dbName, String tableName, TColumnDesc desc) { + return buildField(dbName, tableName, desc, false); + } + + private static Field buildField(String dbName, String tableName, TColumnDesc desc, boolean nativeVariant) { + if (nativeVariant && 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, createFlightSqlColumnMetadata(dbName, tableName, desc)), - arrowChildren(dbName, tableName, desc, arrowType)); + arrowChildren(dbName, tableName, desc, arrowType, nativeVariant)); + } + + static Field nativeVariantField(String name, boolean nullable, Map columnMetadata) { + Map 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()))); } /** @@ -319,7 +339,7 @@ private static Field buildField(String dbName, String tableName, TColumnDesc des * that cannot describe its nested types is no worse off than before. */ private static List arrowChildren(String dbName, String tableName, TColumnDesc desc, - ArrowType arrowType) { + ArrowType arrowType, boolean nativeVariant) { List children = desc.isSetChildren() ? desc.getChildren() : Collections.emptyList(); switch (arrowType.getTypeID()) { case List: @@ -330,7 +350,7 @@ private static List arrowChildren(String dbName, String tableName, TColum Field.notNullable(BaseRepeatedValueVector.DATA_VECTOR_NAME, ZeroVector.INSTANCE.getField().getType())); } - return Collections.singletonList(buildField(dbName, tableName, children.get(0))); + return Collections.singletonList(buildField(dbName, tableName, children.get(0), nativeVariant)); case Map: // Arrow spells a map as list>, with the entries struct and // the key both non-nullable -- the descriptor's key nullability is not carried over, @@ -339,12 +359,13 @@ private static List arrowChildren(String dbName, String tableName, TColum return Collections.singletonList( Field.notNullable(MapVector.DATA_VECTOR_NAME, new ArrowType.List())); } - Field key = buildField(dbName, tableName, children.get(0)); - Field value = buildField(dbName, tableName, children.get(1)); + Field key = buildField(dbName, tableName, children.get(0), nativeVariant); + Field value = buildField(dbName, tableName, children.get(1), nativeVariant); 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: @@ -353,7 +374,7 @@ private static List arrowChildren(String dbName, String tableName, TColum } List structFields = new ArrayList<>(children.size()); for (TColumnDesc child : children) { - structFields.add(buildField(dbName, tableName, child)); + structFields.add(buildField(dbName, tableName, child, nativeVariant)); } return structFields; default: diff --git a/fe/fe-core/src/main/java/org/apache/doris/system/Backend.java b/fe/fe-core/src/main/java/org/apache/doris/system/Backend.java index adb4a57d448385..8148dc22a463f9 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/system/Backend.java +++ b/fe/fe-core/src/main/java/org/apache/doris/system/Backend.java @@ -73,6 +73,8 @@ public class Backend implements Writable { @SerializedName("host") private volatile String host; private String version; + @SerializedName("arrowFlightNativeVariantSupported") + private volatile boolean arrowFlightNativeVariantSupported; @SerializedName("heartbeatPort") private int heartbeatPort; // heartbeat @@ -255,6 +257,10 @@ public String getHost() { return host; } + public boolean isArrowFlightNativeVariantSupported() { + return arrowFlightNativeVariantSupported; + } + public String getVersion() { return version; } @@ -882,6 +888,11 @@ public boolean handleHbResponse(BackendHbResponse hbResponse, boolean isReplay) isChanged = true; supportsPaimonRustReader = hbResponse.isPaimonRustReaderSupported(); } + // An absent capability bit from an older BE must also clear previously advertised support. + if (arrowFlightNativeVariantSupported != hbResponse.isArrowFlightNativeVariantSupported()) { + arrowFlightNativeVariantSupported = hbResponse.isArrowFlightNativeVariantSupported(); + isChanged = true; + } if (!this.version.equals(hbResponse.getVersion())) { isChanged = true; this.version = hbResponse.getVersion(); @@ -966,6 +977,12 @@ public boolean handleHbResponse(BackendHbResponse hbResponse, boolean isReplay) // Only set backend to dead if the heartbeat failure counter exceed threshold. // And if it is a replay process, must set backend to dead. if (isReplay || ++this.heartbeatFailureCounter >= Config.max_backend_heartbeat_failure_tolerance_count) { + // Every journaled BAD means death on replay, including on older FEs. Retain the + // last successful capability during tolerated misses and clear it with the death event. + if (arrowFlightNativeVariantSupported) { + arrowFlightNativeVariantSupported = false; + isChanged = true; + } if (isAlive.compareAndSet(true, false)) { isChanged = true; LOG.warn("{} is dead,", this.toString()); diff --git a/fe/fe-core/src/main/java/org/apache/doris/system/BackendHbResponse.java b/fe/fe-core/src/main/java/org/apache/doris/system/BackendHbResponse.java index 55fb0c9f456a74..80099c46a9d42f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/system/BackendHbResponse.java +++ b/fe/fe-core/src/main/java/org/apache/doris/system/BackendHbResponse.java @@ -49,6 +49,8 @@ public class BackendHbResponse extends HeartbeatResponse implements Writable { private long lastFragmentUpdateTime; @SerializedName(value = "isShutDown") private boolean isShutDown = false; + @SerializedName(value = "arrowFlightNativeVariantSupported") + private boolean arrowFlightNativeVariantSupported; // The physical memory available for use by BE. private long beMemory = 0; @SerializedName("supportsPaimonRustReader") @@ -106,6 +108,14 @@ public BackendHbResponse(long beId, String host, long lastHbTime, String errMsg) this.msg = errMsg; } + public boolean isArrowFlightNativeVariantSupported() { + return arrowFlightNativeVariantSupported; + } + + public void setArrowFlightNativeVariantSupported(boolean supported) { + arrowFlightNativeVariantSupported = supported; + } + public long getFragmentNum() { return fragmentNum; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/system/HeartbeatMgr.java b/fe/fe-core/src/main/java/org/apache/doris/system/HeartbeatMgr.java index 5e13e9191af31d..c59701a3bc0c0c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/system/HeartbeatMgr.java +++ b/fe/fe-core/src/main/java/org/apache/doris/system/HeartbeatMgr.java @@ -374,6 +374,8 @@ private HeartbeatResponse pingOnce() { fragmentNum, lastFragmentUpdateTime, isShutDown, arrowFlightSqlPort, beMemory); response.setPaimonRustReaderSupported(tBackendInfo.isSetSupportsPaimonRustReader() && tBackendInfo.isSupportsPaimonRustReader()); + response.setArrowFlightNativeVariantSupported( + tBackendInfo.isArrowFlightNativeVariantSupported()); return response; } else { return new BackendHbResponse(backendId, backend.getHost(), backend.getLastUpdateMs(), 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 00000000000000..daa9fae74040c4 --- /dev/null +++ b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlNativeVariantTest.java @@ -0,0 +1,342 @@ +// 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.Env; +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.persist.gson.GsonUtils; +import org.apache.doris.planner.PlanNodeId; +import org.apache.doris.planner.ResultSink; +import org.apache.doris.qe.ConnectContext; +import org.apache.doris.qe.SessionVariable; +import org.apache.doris.system.Backend; +import org.apache.doris.system.BackendHbResponse; +import org.apache.doris.system.SystemInfoService; +import org.apache.doris.thrift.TBackendInfo; +import org.apache.doris.thrift.TColumnDesc; +import org.apache.doris.thrift.TDataSink; +import org.apache.doris.thrift.TPrimitiveType; +import org.apache.doris.thrift.TResultSinkType; + +import com.google.common.collect.ImmutableMap; +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.Assert; +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; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; + +public class FlightSqlNativeVariantTest { + @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, true); + 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)); + } + Field legacy = Deencapsulation.invoke(FlightSqlSchemaHelper.class, "buildField", + "test_db", "test_table", variant, false); + Assert.assertEquals(new ArrowType.Utf8(), legacy.getType()); + } + + @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))))); + for (boolean nativeVariant : new boolean[] {false, true}) { + Field result = Deencapsulation.invoke(FlightSqlQuerySchema.class, "field", + "s", nested, true, true, "UTC", nativeVariant); + 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(nativeVariant ? new ArrowType.Struct() : new ArrowType.Utf8(), leaf.getType()); + if (nativeVariant) { + 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()); + } else { + Assert.assertTrue(leaf.getMetadata().isEmpty()); + } + } + } + } + + @Test + public void unknownBackendKeepsUtf8DuringUpgrade() throws Exception { + ConnectContext previous = ConnectContext.get(); + ConnectContext context = new ConnectContext(); + context.setThreadLocalInfo(); + SystemInfoService system = Env.getCurrentSystemInfo(); + Object original = system.getAllBackendsByAllCluster(); + try { + Backend upgraded = new Backend(12346, "127.0.0.1", 9051); + upgraded.handleHbResponse(heartbeat(upgraded.getId(), true), true); + Deencapsulation.setField(system, "idToBackendRef", ImmutableMap.of(upgraded.getId(), upgraded)); + context.getSessionVariable().setEnableArrowFlightSqlNativeVariant(true); + Assert.assertTrue(FlightSqlNativeVariant.isEnabled(context)); + system.addBackend(new Backend(12345, "127.0.0.1", 9050)); + Assert.assertFalse(FlightSqlNativeVariant.isEnabled(context)); + ResultSink flight = new ResultSink(new PlanNodeId(0), TResultSinkType.ARROW_FLIGHT_PROTOCOL); + TDataSink thrift = Deencapsulation.invoke(flight, "toThrift"); + Assert.assertFalse(thrift.getResultSink().isNativeVariant()); + } finally { + Deencapsulation.setField(system, "idToBackendRef", original); + ConnectContext.remove(); + if (previous != null) { + previous.setThreadLocalInfo(); + } + } + } + + @Test + public void metadataWaitsForScopedQuerySessionToBeRestored() throws Exception { + SystemInfoService system = Env.getCurrentSystemInfo(); + Object original = system.getAllBackendsByAllCluster(); + ExecutorService executor = Executors.newSingleThreadExecutor(); + try { + Backend backend = new Backend(12350, "127.0.0.1", 9050); + backend.handleHbResponse(heartbeat(backend.getId(), true), true); + Deencapsulation.setField(system, "idToBackendRef", ImmutableMap.of(backend.getId(), backend)); + ConnectContext context = new ConnectContext(); + for (boolean nativeVariant : new boolean[] {false, true}) { + SessionVariable permanent = context.getSessionVariable(); + permanent.setEnableArrowFlightSqlNativeVariant(nativeVariant); + Assert.assertEquals(nativeVariant, FlightSqlNativeVariant.isEnabled(context)); + Future metadata; + synchronized (context) { + SessionVariable scoped = new SessionVariable(); + scoped.setEnableArrowFlightSqlNativeVariant(!nativeVariant); + context.setSessionVariable(scoped); + try { + // Schema analysis holds this monitor while a SET_VAR clone is temporarily installed. + CountDownLatch started = new CountDownLatch(1); + metadata = executor.submit(() -> { + started.countDown(); + return FlightSqlNativeVariant.isEnabled(context); + }); + Assert.assertTrue(started.await(5, TimeUnit.SECONDS)); + Assert.assertThrows(TimeoutException.class, () -> metadata.get(200, TimeUnit.MILLISECONDS)); + } finally { + context.setSessionVariable(permanent); + } + } + Assert.assertEquals(nativeVariant, metadata.get(5, TimeUnit.SECONDS)); + } + } finally { + try { + executor.shutdownNow(); + Assert.assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } finally { + Deencapsulation.setField(system, "idToBackendRef", original); + } + } + } + + @Test + public void sinkCapturesOptInWithoutChangingMysql() throws Exception { + Assert.assertFalse(new SessionVariable().isEnableArrowFlightSqlNativeVariant()); + SystemInfoService system = Env.getCurrentSystemInfo(); + Object original = system.getAllBackendsByAllCluster(); + Backend backend = new Backend(12346, "127.0.0.1", 9050); + BackendHbResponse heartbeat = heartbeat(backend.getId(), true); + backend.handleHbResponse(heartbeat, true); + Deencapsulation.setField(system, "idToBackendRef", ImmutableMap.of(backend.getId(), backend)); + ConnectContext previous = ConnectContext.get(); + ConnectContext context = new ConnectContext(); + context.setThreadLocalInfo(); + try { + context.getSessionVariable().setEnableArrowFlightSqlNativeVariant(true); + ResultSink flight = new ResultSink(new PlanNodeId(0), TResultSinkType.ARROW_FLIGHT_PROTOCOL); + ResultSink mysql = new ResultSink(new PlanNodeId(0)); + context.getSessionVariable().setEnableArrowFlightSqlNativeVariant(false); + TDataSink flightSink = Deencapsulation.invoke(flight, "toThrift"); + TDataSink mysqlSink = Deencapsulation.invoke(mysql, "toThrift"); + Assert.assertTrue(flightSink.getResultSink().isNativeVariant()); + Assert.assertFalse(mysqlSink.getResultSink().isNativeVariant()); + } finally { + Deencapsulation.setField(system, "idToBackendRef", original); + ConnectContext.remove(); + if (previous != null) { + previous.setThreadLocalInfo(); + } + } + } + + @Test + public void heartbeatReplayPreservesLivenessAndCapability() { + long originalTolerance = Config.max_backend_heartbeat_failure_tolerance_count; + try { + for (int tolerance : new int[] {0, 1, 3}) { + Config.max_backend_heartbeat_failure_tolerance_count = tolerance; + for (boolean supported : new boolean[] {false, true}) { + Backend leader = new Backend(12348, "127.0.0.1", 9050); + Backend follower = new Backend(12348, "127.0.0.1", 9050); + Assert.assertTrue(applyHeartbeatAndReplay(leader, follower, heartbeat(leader.getId(), supported))); + int deathThreshold = Math.max(1, tolerance); + for (int failure = 1; failure <= deathThreshold + 1; failure++) { + BackendHbResponse failed = new BackendHbResponse(leader.getId(), "127.0.0.1", 1, "timeout"); + // HeartbeatMgr journals only changed responses; every journaled BAD means death on replay. + Assert.assertEquals(failure == deathThreshold, applyHeartbeatAndReplay(leader, follower, failed)); + Assert.assertEquals(failure < deathThreshold, leader.isAlive()); + Assert.assertEquals(leader.isAlive(), follower.isAlive()); + Assert.assertEquals(supported && leader.isAlive(), leader.isArrowFlightNativeVariantSupported()); + Assert.assertEquals(leader.isArrowFlightNativeVariantSupported(), + follower.isArrowFlightNativeVariantSupported()); + } + Assert.assertTrue(applyHeartbeatAndReplay(leader, follower, heartbeat(leader.getId(), supported))); + Assert.assertTrue(leader.isAlive()); + Assert.assertTrue(follower.isAlive()); + Assert.assertEquals(supported, leader.isArrowFlightNativeVariantSupported()); + Assert.assertEquals(supported, follower.isArrowFlightNativeVariantSupported()); + } + } + } finally { + Config.max_backend_heartbeat_failure_tolerance_count = originalTolerance; + } + } + + @Test + public void toleratedFailureAndRecoveryRetainLastSuccessfulCapability() throws Exception { + SystemInfoService system = Env.getCurrentSystemInfo(); + Object original = system.getAllBackendsByAllCluster(); + long tolerance = Config.max_backend_heartbeat_failure_tolerance_count; + try { + Config.max_backend_heartbeat_failure_tolerance_count = 3; + Backend leader = new Backend(12348, "127.0.0.1", 9050); + Backend follower = new Backend(12348, "127.0.0.1", 9050); + Deencapsulation.setField(system, "idToBackendRef", ImmutableMap.of(leader.getId(), leader)); + ConnectContext context = new ConnectContext(); + context.getSessionVariable().setEnableArrowFlightSqlNativeVariant(true); + applyHeartbeatAndReplay(leader, follower, heartbeat(leader.getId(), true)); + BackendHbResponse failed = new BackendHbResponse(leader.getId(), "127.0.0.1", 1, "timeout"); + Assert.assertFalse(applyHeartbeatAndReplay(leader, follower, failed)); + Assert.assertTrue(leader.isAlive()); + Assert.assertTrue(follower.isAlive()); + Assert.assertTrue(FlightSqlNativeVariant.isEnabled(context)); + Assert.assertTrue(follower.isArrowFlightNativeVariantSupported()); + applyHeartbeatAndReplay(leader, follower, heartbeat(leader.getId(), true)); + // A successful heartbeat resets the failure count before a later missed heartbeat. + Assert.assertFalse(applyHeartbeatAndReplay(leader, follower, failed)); + Assert.assertFalse(applyHeartbeatAndReplay(leader, follower, failed)); + Assert.assertTrue(applyHeartbeatAndReplay(leader, follower, failed)); + Assert.assertFalse(FlightSqlNativeVariant.isEnabled(context)); + Assert.assertFalse(leader.isAlive()); + Assert.assertFalse(follower.isAlive()); + applyHeartbeatAndReplay(leader, follower, heartbeat(leader.getId(), true)); + Assert.assertTrue(FlightSqlNativeVariant.isEnabled(context)); + // The next successful process report remains authoritative, including an absent legacy bit. + String legacy = GsonUtils.GSON.toJson(heartbeat(leader.getId(), true)) + .replace(",\"arrowFlightNativeVariantSupported\":true", ""); + applyHeartbeatAndReplay(leader, follower, GsonUtils.GSON.fromJson(legacy, BackendHbResponse.class)); + Assert.assertTrue(leader.isAlive()); + Assert.assertTrue(follower.isAlive()); + Assert.assertFalse(FlightSqlNativeVariant.isEnabled(context)); + Assert.assertFalse(follower.isArrowFlightNativeVariantSupported()); + } finally { + Config.max_backend_heartbeat_failure_tolerance_count = tolerance; + Deencapsulation.setField(system, "idToBackendRef", original); + } + } + + private static boolean applyHeartbeatAndReplay(Backend leader, Backend follower, BackendHbResponse response) { + boolean changed = leader.handleHbResponse(response, false); + if (changed) { + follower.handleHbResponse(GsonUtils.GSON.fromJson( + GsonUtils.GSON.toJson(response), BackendHbResponse.class), true); + } + return changed; + } + + @Test + public void heartbeatCapabilitiesRemainIndependent() { + // Field 11 already belongs to Paimon; sharing its wire ID would enable the wrong reader. + Assert.assertEquals(11, TBackendInfo._Fields.SUPPORTS_PAIMON_RUST_READER.getThriftFieldId()); + Assert.assertEquals(12, TBackendInfo._Fields.ARROW_FLIGHT_NATIVE_VARIANT_SUPPORTED.getThriftFieldId()); + Backend backend = new Backend(12349, "127.0.0.1", 9050); + for (boolean paimon : new boolean[] {false, true}) { + for (boolean variant : new boolean[] {false, true}) { + BackendHbResponse response = heartbeat(backend.getId(), variant); + response.setPaimonRustReaderSupported(paimon); + backend.handleHbResponse(GsonUtils.GSON.fromJson( + GsonUtils.GSON.toJson(response), BackendHbResponse.class), true); + Assert.assertEquals(paimon, backend.isPaimonRustReaderSupported()); + Assert.assertEquals(variant, backend.isArrowFlightNativeVariantSupported()); + } + } + } + + private static BackendHbResponse heartbeat(long id, boolean supported) { + BackendHbResponse response = new BackendHbResponse(id, 9060, 8040, 8060, + 1, 1, "test", "mix", 0, 0, false, 8815); + response.setArrowFlightNativeVariantSupported(supported); + return response; + } + + @Test + public void heartbeatCapabilitySurvivesReplayAndClearsOnDowngrade() { + Backend backend = new Backend(12347, "127.0.0.1", 9050); + Assert.assertFalse(backend.isArrowFlightNativeVariantSupported()); + BackendHbResponse advertised = heartbeat(backend.getId(), true); + String serialized = GsonUtils.GSON.toJson(advertised); + backend.handleHbResponse(GsonUtils.GSON.fromJson(serialized, BackendHbResponse.class), true); + Assert.assertTrue(backend.isArrowFlightNativeVariantSupported()); + // An old BE omits the new heartbeat field after a rollback. + String oldHeartbeat = serialized.replace(",\"arrowFlightNativeVariantSupported\":true", ""); + backend.handleHbResponse(GsonUtils.GSON.fromJson(oldHeartbeat, BackendHbResponse.class), true); + Assert.assertFalse(backend.isArrowFlightNativeVariantSupported()); + } + +} diff --git a/gensrc/thrift/DataSinks.thrift b/gensrc/thrift/DataSinks.thrift index 82f7cbb0a14873..29bd090f79b29f 100644 --- a/gensrc/thrift/DataSinks.thrift +++ b/gensrc/thrift/DataSinks.thrift @@ -231,6 +231,8 @@ struct TResultSink { 1: optional TResultSinkType type; 2: optional TResultFileSinkOptions file_options; // deprecated 3: optional TFetchOption fetch_option; + // Freeze the Flight result representation with the query's sink and schema. + 4: optional bool native_variant = false; } struct TResultFileSink { diff --git a/gensrc/thrift/HeartbeatService.thrift b/gensrc/thrift/HeartbeatService.thrift index 8608853b29f81a..072623329e6b67 100644 --- a/gensrc/thrift/HeartbeatService.thrift +++ b/gensrc/thrift/HeartbeatService.thrift @@ -64,6 +64,7 @@ struct TBackendInfo { 9: optional Types.TPort arrow_flight_sql_port 10: optional i64 be_mem // The physical memory available for use by BE. 11: optional bool supports_paimon_rust_reader + 12: optional bool arrow_flight_native_variant_supported = false // For cloud 1000: optional i64 fragment_executing_count 1001: optional i64 fragment_last_active_time 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 aff57aeac4d3bb..9b3a21dd30b3b1 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 00000000000000..495ef224ee512c --- /dev/null +++ b/regression-test/suites/arrow_flight_sql_p0/test_flight_native_variant.groovy @@ -0,0 +1,196 @@ +// 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] + // Match ingestion to the configured Variant representation; legacy expression roots + // cannot be cast to a table Variant with a different subcolumn limit. + def variantV2Function = getFeConfig("enable_variant_v2").toBoolean() ? "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 -> }) } + 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}") + [false, true, false].each { nativeVariant -> + executeSetting("SET enable_arrow_flight_sql_native_variant=${nativeVariant}") + def seen = [] + assertEquals(60, read("SELECT id, v FROM ${table}", { root -> + def vector = root.getVector(1) + def field = vector.field + if (nativeVariant) { + 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()) + } + } else { + assertEquals("Utf8", field.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 (nativeVariant && 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 }) + } + } else if (!nativeVariant && id % 4 == 1) { + assertEquals("42", vector.getObject(i).toString()) + } + } + }, parallel, resultBackendCount)) + assertEquals((1..60).toList(), seen.sort()) + // Legacy Variant is not a legal ARRAY() argument; nested SQL coverage requires V2. + def scannedColumns = variantV2Function + ? "v, ARRAY(v) AS a, MAP('key', v) AS m, STRUCT(v) AS s" : "v" + // 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 && variantV2Function) { + 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(nativeVariant ? "Struct" : "Utf8", field.type.toString()) + if (nativeVariant) { + 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) + } + } + } + } + executeSetting("SET enable_arrow_flight_sql_native_variant=true") + // 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 = variantV2Function + ? """parse_to_variant(CONCAT('{"', ${keyExpression}, '":', CAST(number AS STRING), '}'))""" + : "CAST(MAP(${keyExpression}, number) AS VARIANT)" + 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 CAST(42 AS VARIANT) AS v", { root -> + assertEquals("arrow.parquet.variant", root.getVector(0).field.metadata.get("ARROW:extension:name")) + })) + assertEquals(0, read("SELECT 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/samples/arrow-flight-sql/python/README.md b/samples/arrow-flight-sql/python/README.md index e729054b81359b..8cdafd1ba34edd 100644 --- a/samples/arrow-flight-sql/python/README.md +++ b/samples/arrow-flight-sql/python/README.md @@ -38,4 +38,61 @@ under the License. # 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 results + +On branch-4.1 builds with native VARIANT support, enable it on the same Flight SQL +connection that executes the query: + +```sql +SET enable_arrow_flight_sql_native_variant = true; +SELECT variant_column FROM example_table; +``` + +The default is `false`, which retains the existing UTF8 representation. When enabled +and every registered BE has advertised native Variant support in its heartbeat, +VARIANT fields (including nested fields) use the `arrow.parquet.variant` extension +with `struct` storage. SQL NULL is +a null struct. V2 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 roots use recursive typed encoding, +including MAP, STRUCT, ARRAY, VARBINARY, TIMEV2 and nested VARIANT values. VARBINARY +retains its original bytes. MAP keys become object field names; NULL keys are rejected +because they cannot be distinguished from a literal `"null"` object key. TIMEV2 values +in `[00:00:00, 24:00:00)` retain their microseconds as a native Variant time value; +negative and longer durations are rejected because Parquet TIME is a time of day. +Use `enable_arrow_flight_sql_native_variant=false` to read these unsupported values. +Legacy document fields also use typed encoding, preserving DECIMAL precision and +DATE identity across dense paths, sparse paths and document snapshots. Existing +null/missing semantics are retained. In particular, +a legacy null root remains an empty object, including in a scalar-only batch; outer +SQL NULL remains a null struct. +Decimal256 scalar roots are rejected because the wire format has no Decimal256 primitive. +Native encoding currently accepts at most 128 nested levels. Deeper legacy documents +remain readable with `enable_arrow_flight_sql_native_variant=false`. +During a rolling upgrade, missing support on any registered BE keeps both query results +and GetTables metadata in UTF8 mode, including when an older BE may proxy a result. +Capability follows the last successful heartbeat during tolerated heartbeat failures. +When heartbeat failures mark a BE dead, its capability is cleared on every FE until a +successful heartbeat advertises support again. A successful heartbeat from an older BE +also clears the capability. These changes affect newly planned queries; outstanding +Flight tickets are not migrated across BE replacement. Heartbeat discovery cannot make +an in-place downgrade atomic with query planning. Drain active queries and stop new +native-mode queries before downgrading a BE. + +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 keep the default string mode. + +To check ADBC query and partition reads against a running 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 00000000000000..5eb1c6b44cf1f9 --- /dev/null +++ b/samples/arrow-flight-sql/python/test_variant.py @@ -0,0 +1,102 @@ +#!/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 with DORIS_FLIGHT_URI, DORIS_USER and DORIS_PASSWORD in the environment.""" + +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, CAST(42 AS VARIANT) AS v + UNION ALL SELECT 2, CAST(NULL AS VARIANT)""" + 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()}") + for native in (False, True, False): + execute(f"SET enable_arrow_flight_sql_native_variant={str(native).lower()}") + table = execute(query) + self.check_result(table, native) + # 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, native) + + 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), native) + + def check_result(self, table, native): + table = table.sort_by("id") + self.assertEqual(table.num_rows, 2) + field = table.schema.field("v") + values = table.column("v").combine_chunks() + if not native: + self.assertEqual(field.type, pa.string()) + self.assertEqual(values.to_pylist(), ["42", None]) + return + # 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()