From fb6670861a84de7adf4b94532b08c15d352632b8 Mon Sep 17 00:00:00 2001 From: happenlee Date: Mon, 10 Aug 2026 12:21:58 +0800 Subject: [PATCH 1/4] [fix](be) Support complex array_agg state serialization ### What problem does this PR solve? Issue Number: close #65976 Related PR: None Problem Summary: Multi-phase array_agg_foreach queries on complex element types failed because the generic array_agg state could neither serialize, deserialize, nor merge its buffered column. Persist complex states with the native data type serialization format and merge state columns by value, covering both distributed aggregation and ForEach state growth. ### Release note Fix array_agg and array_agg_foreach failures for complex element types in multi-phase aggregation. ### Check List (For Author) - Test: Unit Test and Regression test - Unit Test: GLIBC_COMPATIBILITY=OFF ./run-be-ut.sh -j 48 --run --filter=AggregateFunctionArrayAggTest.* - Regression test: ./run-regression-test.sh --run -d function_p0 -s test_agg_foreach - Behavior changed: Yes. Complex array_agg states now support serialization and merge. - Does this need documentation: No --- .../aggregate/aggregate_function_array_agg.h | 42 +++-- .../exprs/aggregate/agg_array_agg_test.cpp | 145 +++++++++++++++++- .../function_p0/test_agg_foreach.groovy | 2 +- 3 files changed, 177 insertions(+), 12 deletions(-) diff --git a/be/src/exprs/aggregate/aggregate_function_array_agg.h b/be/src/exprs/aggregate/aggregate_function_array_agg.h index f2c7a1feadb3b9..c7adc0ff460953 100644 --- a/be/src/exprs/aggregate/aggregate_function_array_agg.h +++ b/be/src/exprs/aggregate/aggregate_function_array_agg.h @@ -17,6 +17,7 @@ #pragma once +#include "common/check.h" #include "core/assert_cast.h" #include "core/column/column.h" #include "core/column/column_array.h" @@ -234,11 +235,13 @@ struct AggregateFunctionArrayAggData { using ElementType = StringRef; using Self = AggregateFunctionArrayAggData; MutableColumnPtr column_data; + DataTypePtr column_type; + int be_exec_version; - AggregateFunctionArrayAggData(const DataTypes& argument_types) { - DataTypePtr column_type = argument_types[0]; - column_data = column_type->create_column(); - } + AggregateFunctionArrayAggData(const DataTypes& argument_types, int be_exec_version_ = 0) + : column_data(argument_types[0]->create_column()), + column_type(argument_types[0]), + be_exec_version(be_exec_version_) {} void add(const IColumn& column, size_t row_num) { column_data->insert_from(column, row_num); } @@ -265,15 +268,27 @@ struct AggregateFunctionArrayAggData { } void write(BufferWritable& buf) const { - throw Exception(ErrorCode::NOT_IMPLEMENTED_ERROR, "array_agg not support write"); + const auto serialized_bytes = + column_type->get_uncompressed_serialized_bytes(*column_data, be_exec_version); + std::string serialized_buffer(serialized_bytes, '\0'); + const auto* end = + column_type->serialize(*column_data, serialized_buffer.data(), be_exec_version); + DORIS_CHECK_LE(end, serialized_buffer.data() + serialized_bytes); + serialized_buffer.resize(end - serialized_buffer.data()); + buf.write_binary(serialized_buffer); } void read(BufferReadable& buf) { - throw Exception(ErrorCode::NOT_IMPLEMENTED_ERROR, "array_agg not support read"); + DORIS_CHECK(column_data->empty()); + std::string serialized_buffer; + buf.read_binary(serialized_buffer); + const auto* end = + column_type->deserialize(serialized_buffer.data(), &column_data, be_exec_version); + DORIS_CHECK_EQ(end, serialized_buffer.data() + serialized_buffer.size()); } void merge(const Self& rhs) { - throw Exception(ErrorCode::NOT_IMPLEMENTED_ERROR, "array_agg not support merge"); + column_data->insert_range_from(*rhs.column_data, 0, rhs.column_data->size()); } }; @@ -285,15 +300,24 @@ class AggregateFunctionArrayAgg final UnaryExpression, NotNullableAggregateFunction { public: + using Base = IAggregateFunctionDataHelper, true>; + AggregateFunctionArrayAgg(const DataTypes& argument_types_) - : IAggregateFunctionDataHelper, true>( - {argument_types_}), + : Base({argument_types_}), return_type(std::make_shared(make_nullable(argument_types_[0]))) {} std::string get_name() const override { return "array_agg"; } DataTypePtr get_return_type() const override { return return_type; } + void create(AggregateDataPtr __restrict place) const override { + if constexpr (Data::PType == INVALID_TYPE) { + new (place) Data(this->argument_types, this->version); + } else { + Base::create(place); + } + } + void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num, Arena& arena) const override { this->data(place).add(*columns[0], row_num); diff --git a/be/test/exprs/aggregate/agg_array_agg_test.cpp b/be/test/exprs/aggregate/agg_array_agg_test.cpp index d65565a4c99b18..dac424f3850040 100644 --- a/be/test/exprs/aggregate/agg_array_agg_test.cpp +++ b/be/test/exprs/aggregate/agg_array_agg_test.cpp @@ -17,13 +17,14 @@ #include #include -#include -#include +#include +#include #include #include #include +#include "agent/be_exec_version_manager.h" #include "common/logging.h" #include "core/arena.h" #include "core/column/column_array.h" @@ -34,9 +35,11 @@ #include "core/data_type/data_type_date.h" #include "core/data_type/data_type_date_time.h" #include "core/data_type/data_type_decimal.h" +#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/string_buffer.hpp" #include "core/types.h" #include "exprs/aggregate/agg_function_test.h" @@ -53,6 +56,73 @@ namespace doris { struct AggregateFunctionArrayAggTest : public AggregateFunctiontest {}; +namespace { + +Field array_field(std::initializer_list values) { + return Field::create_field(Array(values)); +} + +Field map_field(std::initializer_list keys, std::initializer_list values) { + return Field::create_field(Map {array_field(keys), array_field(values)}); +} + +void add_column_to_state(const IAggregateFunction& function, AggregateDataPtr state, + const IColumn& column, Arena& arena) { + const IColumn* columns[] = {&column}; + for (size_t row = 0; row < column.size(); ++row) { + function.add(state, columns, row, arena); + } +} + +void check_complex_array_agg_state(const DataTypePtr& data_type, + std::initializer_list source_values, + std::initializer_list rhs_values) { + SCOPED_TRACE(data_type->get_name()); + const auto nullable_type = make_nullable(data_type); + auto function = AggregateFunctionSimpleFactory::instance().get( + "array_agg", {nullable_type}, nullptr, false, + BeExecVersionManager::get_newest_version()); + ASSERT_NE(function, nullptr); + function->set_version(BeExecVersionManager::get_newest_version()); + + auto source_column = nullable_type->create_column(); + for (const auto& value : source_values) { + source_column->insert(value); + } + auto rhs_column = nullable_type->create_column(); + for (const auto& value : rhs_values) { + rhs_column->insert(value); + } + + Arena arena; + AggregateFunctionGuard source(function.get()); + AggregateFunctionGuard restored(function.get()); + AggregateFunctionGuard rhs(function.get()); + add_column_to_state(*function, source.data(), *source_column, arena); + add_column_to_state(*function, rhs.data(), *rhs_column, arena); + + auto serialized_column = ColumnString::create(); + BufferWritable writer(*serialized_column); + function->serialize(source.data(), writer); + writer.commit(); + ASSERT_EQ(serialized_column->size(), 1); + + auto serialized_data = serialized_column->get_data_at(0); + BufferReadable reader(serialized_data); + function->deserialize(restored.data(), reader, arena); + function->merge(restored.data(), rhs.data(), arena); + + auto result_column = function->get_return_type()->create_column(); + function->insert_result_into(restored.data(), *result_column); + auto expected_column = function->get_return_type()->create_column(); + Array expected_values(source_values); + expected_values.insert(expected_values.end(), rhs_values.begin(), rhs_values.end()); + expected_column->insert(Field::create_field(expected_values)); + EXPECT_TRUE(ColumnHelper::column_equal(std::move(result_column), std::move(expected_column))); +} + +} // namespace + TEST_F(AggregateFunctionArrayAggTest, test_array_agg_aint64) { create_agg("array_agg", false, {std::make_shared()}, std::make_shared()); @@ -201,6 +271,77 @@ TEST_F(AggregateFunctionArrayAggTest, test_array_agg_aint64_foreach) { ColumnWithTypeAndName(std::move(array_array_column), array_array_data_type, "column")); } +TEST_F(AggregateFunctionArrayAggTest, complex_type_state_serialize_deserialize_and_merge) { + auto nullable_int = make_nullable(std::make_shared()); + auto nullable_string = make_nullable(std::make_shared()); + + check_complex_array_agg_state( + std::make_shared(nullable_int), + {array_field({Field::create_field(1), Field()}), Field()}, + {array_field({Field::create_field(2), Field::create_field(3)})}); + + check_complex_array_agg_state( + std::make_shared(DataTypes {nullable_int, nullable_string}), + {Field::create_field(Struct {Field::create_field(1), + Field::create_field("one")}), + Field()}, + {Field::create_field( + Struct {Field(), Field::create_field("two")})}); + + check_complex_array_agg_state(std::make_shared(nullable_string, nullable_int), + {map_field({Field::create_field("one"), + Field::create_field("null")}, + {Field::create_field(1), Field()}), + Field()}, + {map_field({Field::create_field("two")}, + {Field::create_field(2)})}); +} + +TEST_F(AggregateFunctionArrayAggTest, foreach_complex_type_state_growth_and_round_trip) { + auto nullable_int = make_nullable(std::make_shared()); + auto nullable_inner_array = make_nullable(std::make_shared(nullable_int)); + auto input_type = std::make_shared(nullable_inner_array); + auto function = AggregateFunctionSimpleFactory::instance().get( + "array_agg_foreach", {input_type}, input_type, false, + BeExecVersionManager::get_newest_version(), {.is_foreach = true, .column_names = {}}); + ASSERT_NE(function, nullptr); + function->set_version(BeExecVersionManager::get_newest_version()); + + auto input_column = input_type->create_column(); + input_column->insert(array_field({array_field({Field::create_field(1)})})); + input_column->insert(array_field( + {array_field({Field::create_field(2)}), + array_field({Field::create_field(3), Field::create_field(4)}), + Field()})); + + Arena arena; + AggregateFunctionGuard source(function.get()); + AggregateFunctionGuard restored(function.get()); + AggregateFunctionGuard merged(function.get()); + add_column_to_state(*function, source.data(), *input_column, arena); + + auto serialized_column = ColumnString::create(); + BufferWritable writer(*serialized_column); + function->serialize(source.data(), writer); + writer.commit(); + + auto serialized_data = serialized_column->get_data_at(0); + BufferReadable reader(serialized_data); + function->deserialize(restored.data(), reader, arena); + function->merge(merged.data(), restored.data(), arena); + + auto result_column = function->get_return_type()->create_column(); + function->insert_result_into(merged.data(), *result_column); + auto expected_column = function->get_return_type()->create_column(); + expected_column->insert( + array_field({array_field({array_field({Field::create_field(1)}), + array_field({Field::create_field(2)})}), + array_field({array_field({Field::create_field(3), + Field::create_field(4)})}), + array_field({Field()})})); + EXPECT_TRUE(ColumnHelper::column_equal(std::move(result_column), std::move(expected_column))); +} + TEST(AggregateFunctionSortDataTest, merge_does_not_share_rhs_block) { auto data_type = std::make_shared(); Block prototype({ColumnWithTypeAndName(data_type->create_column(), data_type, "value"), diff --git a/regression-test/suites/function_p0/test_agg_foreach.groovy b/regression-test/suites/function_p0/test_agg_foreach.groovy index d2280e0c9ce2b1..72fea2f4cfc955 100644 --- a/regression-test/suites/function_p0/test_agg_foreach.groovy +++ b/regression-test/suites/function_p0/test_agg_foreach.groovy @@ -119,7 +119,7 @@ suite("test_agg_foreach") { qt_sql """select array_agg_foreach(s) from foreach_table;""" qt_array_agg_nested """ - select /*+ SET_VAR(parallel_pipeline_task_num=1) */ + select /*+ SET_VAR(agg_phase=2, parallel_pipeline_task_num=1) */ size(array_agg_foreach(b)), size(element_at(array_agg_foreach(b), 1)), size(element_at(array_agg_foreach(b), 2)), From 012fb525044ff1b89420aca3807acd2b27e07a1a Mon Sep 17 00:00:00 2001 From: happenlee Date: Mon, 10 Aug 2026 15:25:05 +0800 Subject: [PATCH 2/4] [fix](be) Preserve padding for complex array_agg states ### What problem does this PR solve? Issue Number: close #65976 Related PR: #66601 Problem Summary: Native DataType deserialization may read STREAMVBYTE_PADDING bytes past the logical compressed payload. Copying the framed complex array_agg state into an exact-size std::string discarded the padded ColumnString backing buffer and also imposed the generic 1 GiB string limit. Deserialize directly from the framed padded buffer, validate the logical payload length, and advance the reader by that length. ### Release note Fix complex array_agg state deserialization for StreamVByte-compressed payloads. ### Check List (For Author) - Test: Unit Test - Unit Test: GLIBC_COMPATIBILITY=OFF ./run-be-ut.sh -j 48 --run --filter=AggregateFunctionArrayAggTest.* - Behavior changed: Yes. StreamVByte-compressed complex array_agg states now retain the required readable padding during deserialization. - Does this need documentation: No --- be/src/exprs/aggregate/aggregate_function_array_agg.h | 11 ++++++----- be/test/exprs/aggregate/agg_array_agg_test.cpp | 8 ++++++++ 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/be/src/exprs/aggregate/aggregate_function_array_agg.h b/be/src/exprs/aggregate/aggregate_function_array_agg.h index c7adc0ff460953..40e7426c0c9112 100644 --- a/be/src/exprs/aggregate/aggregate_function_array_agg.h +++ b/be/src/exprs/aggregate/aggregate_function_array_agg.h @@ -280,11 +280,12 @@ struct AggregateFunctionArrayAggData { void read(BufferReadable& buf) { DORIS_CHECK(column_data->empty()); - std::string serialized_buffer; - buf.read_binary(serialized_buffer); - const auto* end = - column_type->deserialize(serialized_buffer.data(), &column_data, be_exec_version); - DORIS_CHECK_EQ(end, serialized_buffer.data() + serialized_buffer.size()); + UInt64 serialized_bytes = 0; + buf.read_var_uint(serialized_bytes); + const auto* serialized_data = buf.data(); + const auto* end = column_type->deserialize(serialized_data, &column_data, be_exec_version); + DORIS_CHECK_EQ(end, serialized_data + serialized_bytes); + buf.add_offset(serialized_bytes); } void merge(const Self& rhs) { diff --git a/be/test/exprs/aggregate/agg_array_agg_test.cpp b/be/test/exprs/aggregate/agg_array_agg_test.cpp index dac424f3850040..3f2d0c693abae3 100644 --- a/be/test/exprs/aggregate/agg_array_agg_test.cpp +++ b/be/test/exprs/aggregate/agg_array_agg_test.cpp @@ -275,6 +275,14 @@ TEST_F(AggregateFunctionArrayAggTest, complex_type_state_serialize_deserialize_a auto nullable_int = make_nullable(std::make_shared()); auto nullable_string = make_nullable(std::make_shared()); + Array streamvbyte_values; + for (int32_t value = 0; value < 65; ++value) { + streamvbyte_values.emplace_back(Field::create_field(value)); + } + check_complex_array_agg_state(std::make_shared(nullable_int), + {Field::create_field(std::move(streamvbyte_values))}, + {}); + check_complex_array_agg_state( std::make_shared(nullable_int), {array_field({Field::create_field(1), Field()}), Field()}, From 6d36bea59589a4c1e519cc9d9db9b1beca766cc4 Mon Sep 17 00:00:00 2001 From: happenlee Date: Mon, 10 Aug 2026 16:16:23 +0800 Subject: [PATCH 3/4] [fix](be) Make foreach state growth exception-safe ### What problem does this PR solve? Issue Number: close #65976 Related PR: #66601 Problem Summary: Growing a multi-position foreach aggregate state merged and destroyed old positions one at a time before publishing the replacement buffer. If a later nested merge failed, query cleanup could destroy an already destroyed old state while newly constructed states were leaked. Keep all old states alive while constructing and merging the replacement, clean every constructed replacement state on failure, and destroy and publish only after all merges succeed. ### Release note Preserve recoverable query errors during foreach aggregate state growth. ### Check List (For Author) - Test: Unit Test - Unit Test: GLIBC_COMPATIBILITY=OFF ./run-be-ut.sh -j 48 --run --filter=AggregateFunctionExceptionTest.*:AggregateFunctionArrayAggTest.* - Behavior changed: Yes. Failed foreach state relocation now leaves the original state valid for cleanup. - Does this need documentation: No --- .../aggregate/aggregate_function_foreach.h | 19 ++--- .../aggregate_function_exception_test.cpp | 70 ++++++++++++++++++- 2 files changed, 79 insertions(+), 10 deletions(-) diff --git a/be/src/exprs/aggregate/aggregate_function_foreach.h b/be/src/exprs/aggregate/aggregate_function_foreach.h index 28313dee8b6379..d6324cd7694933 100644 --- a/be/src/exprs/aggregate/aggregate_function_foreach.h +++ b/be/src/exprs/aggregate/aggregate_function_foreach.h @@ -97,24 +97,25 @@ class AggregateFunctionForEach : public AggregateFunctionNonFinalBase, char* new_state = arena.aligned_alloc(allocation_size, nested_function->align_of_data()); - size_t i; + size_t num_created = 0; try { - for (i = 0; i < new_size; ++i) { - nested_function->create(&new_state[i * nested_size_of_data]); + for (; num_created < new_size; ++num_created) { + nested_function->create(&new_state[num_created * nested_size_of_data]); } - } catch (...) { - size_t cleanup_size = i; - for (i = 0; i < cleanup_size; ++i) { + for (size_t i = 0; i < old_size; ++i) { + nested_function->merge(&new_state[i * nested_size_of_data], + &old_state[i * nested_size_of_data], arena); + } + } catch (...) { + for (size_t i = 0; i < num_created; ++i) { nested_function->destroy(&new_state[i * nested_size_of_data]); } throw; } - for (i = 0; i < old_size; ++i) { - nested_function->merge(&new_state[i * nested_size_of_data], - &old_state[i * nested_size_of_data], arena); + for (size_t i = 0; i < old_size; ++i) { nested_function->destroy(&old_state[i * nested_size_of_data]); } diff --git a/be/test/exprs/aggregate/aggregate_function_exception_test.cpp b/be/test/exprs/aggregate/aggregate_function_exception_test.cpp index 21ee64dba4aef1..d571e8c032000a 100644 --- a/be/test/exprs/aggregate/aggregate_function_exception_test.cpp +++ b/be/test/exprs/aggregate/aggregate_function_exception_test.cpp @@ -17,10 +17,14 @@ #include +#include #include #include "core/arena.h" +#include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_string.h" #include "exprs/aggregate/aggregate_function.h" +#include "exprs/aggregate/aggregate_function_foreach.h" namespace doris { @@ -31,14 +35,17 @@ struct TrackingAggregateState { static void reset_counters() { construct_count = 0; destroy_count = 0; + merge_count = 0; } static int construct_count; static int destroy_count; + static int merge_count; }; int TrackingAggregateState::construct_count = 0; int TrackingAggregateState::destroy_count = 0; +int TrackingAggregateState::merge_count = 0; class ThrowOnDeserializeAggregateFunction final : public IAggregateFunctionDataHelper { +public: + ThrowOnSecondMergeAggregateFunction() + : IAggregateFunctionDataHelper( + DataTypes {std::make_shared()}) {} + + String get_name() const override { return "throw_on_second_merge"; } + + DataTypePtr get_return_type() const override { return std::make_shared(); } + + void add(AggregateDataPtr, const IColumn**, ssize_t, Arena&) const override {} + + void merge(AggregateDataPtr, ConstAggregateDataPtr, Arena&) const override { + if (++TrackingAggregateState::merge_count == 2) { + throw Exception(ErrorCode::MEM_ALLOC_FAILED, "mock merge allocation failure"); + } + } + + void serialize(ConstAggregateDataPtr, BufferWritable&) const override {} + + void deserialize(AggregateDataPtr, BufferReadable&, Arena&) const override {} + + void insert_result_into(ConstAggregateDataPtr, IColumn&) const override {} +}; + class AggregateFunctionExceptionTest : public testing::Test { protected: void SetUp() override { TrackingAggregateState::reset_counters(); } @@ -159,4 +194,37 @@ TEST_F(AggregateFunctionExceptionTest, EXPECT_EQ(TrackingAggregateState::construct_count, TrackingAggregateState::destroy_count); } -} // namespace doris \ No newline at end of file +TEST_F(AggregateFunctionExceptionTest, ForEachGrowthPreservesOldStatesWhenMergeThrows) { + auto nested_function = std::make_shared(); + auto input_type = std::make_shared(std::make_shared()); + AggregateFunctionForEach foreach_function(nested_function, DataTypes {input_type}); + auto input_column = input_type->create_column(); + input_column->insert( + Field::create_field(Array {Field::create_field(String("a")), + Field::create_field(String("b"))})); + input_column->insert( + Field::create_field(Array {Field::create_field(String("c")), + Field::create_field(String("d")), + Field::create_field(String("e"))})); + const IColumn* columns[] = {input_column.get()}; + + { + AggregateFunctionGuard state(&foreach_function); + foreach_function.add(state.data(), columns, 0, arena); + + try { + foreach_function.add(state.data(), columns, 1, arena); + FAIL() << "Expected doris::Exception"; + } catch (const doris::Exception& e) { + EXPECT_EQ(e.code(), doris::ErrorCode::MEM_ALLOC_FAILED); + } + + EXPECT_EQ(TrackingAggregateState::merge_count, 2); + EXPECT_EQ(TrackingAggregateState::construct_count, 5); + EXPECT_EQ(TrackingAggregateState::destroy_count, 3); + } + + EXPECT_EQ(TrackingAggregateState::construct_count, TrackingAggregateState::destroy_count); +} + +} // namespace doris From 4d28bde19dd6073cc8127274c237c27149c0de2f Mon Sep 17 00:00:00 2001 From: happenlee Date: Mon, 10 Aug 2026 17:33:42 +0800 Subject: [PATCH 4/4] [fix](be) Reduce complex array_agg serialization memory ### What problem does this PR solve? Issue Number: close #65976 Related PR: #66601 Problem Summary: Complex array_agg states stored immutable type and execution-version metadata in every group and foreach position, adding per-state overhead. Serialization also allocated an untracked upper-bound string and copied it into the tracked destination, doubling peak payload memory. Keep serializer metadata on the aggregate function, serialize once into tracked padded destination storage with rollback on failure, and retain exact logical framing for zero-copy reads. ### Release note Reduce memory usage for complex array_agg and array_agg_foreach state serialization. ### Check List (For Author) - Test: Unit Test - Unit Test: GLIBC_COMPATIBILITY=OFF ./run-be-ut.sh -j 48 --run --filter=AggregateFunctionArrayAggTest.*:AggregateFunctionExceptionTest.* - Behavior changed: Yes. Complex array_agg state serialization removes redundant per-state metadata and tracked scratch duplication without changing query results. - Does this need documentation: No --- .../aggregate/aggregate_function_array_agg.h | 64 ++++++++++-------- .../exprs/aggregate/agg_array_agg_test.cpp | 66 +++++++++++++++++++ 2 files changed, 102 insertions(+), 28 deletions(-) diff --git a/be/src/exprs/aggregate/aggregate_function_array_agg.h b/be/src/exprs/aggregate/aggregate_function_array_agg.h index 40e7426c0c9112..e184f37c6604a3 100644 --- a/be/src/exprs/aggregate/aggregate_function_array_agg.h +++ b/be/src/exprs/aggregate/aggregate_function_array_agg.h @@ -38,6 +38,7 @@ class Arena; template struct AggregateFunctionArrayAggData { static constexpr PrimitiveType PType = T; + static constexpr bool use_native_serde = false; using ElementType = typename PrimitiveTypeTraits::CppType; using ColVecType = typename PrimitiveTypeTraits::ColumnType; using Self = AggregateFunctionArrayAggData; @@ -137,6 +138,7 @@ template requires(is_string_type(T)) struct AggregateFunctionArrayAggData { static constexpr PrimitiveType PType = T; + static constexpr bool use_native_serde = false; using ElementType = StringRef; using ColVecType = ColumnString; using Self = AggregateFunctionArrayAggData; @@ -232,16 +234,13 @@ template !is_date_type(T) && !is_ip(T)) struct AggregateFunctionArrayAggData { static constexpr PrimitiveType PType = T; + static constexpr bool use_native_serde = true; using ElementType = StringRef; using Self = AggregateFunctionArrayAggData; MutableColumnPtr column_data; - DataTypePtr column_type; - int be_exec_version; - AggregateFunctionArrayAggData(const DataTypes& argument_types, int be_exec_version_ = 0) - : column_data(argument_types[0]->create_column()), - column_type(argument_types[0]), - be_exec_version(be_exec_version_) {} + AggregateFunctionArrayAggData(const DataTypes& argument_types) + : column_data(argument_types[0]->create_column()) {} void add(const IColumn& column, size_t row_num) { column_data->insert_from(column, row_num); } @@ -267,23 +266,32 @@ struct AggregateFunctionArrayAggData { to_arr.get_offsets().push_back(to_nested_col.size()); } - void write(BufferWritable& buf) const { - const auto serialized_bytes = - column_type->get_uncompressed_serialized_bytes(*column_data, be_exec_version); - std::string serialized_buffer(serialized_bytes, '\0'); - const auto* end = - column_type->serialize(*column_data, serialized_buffer.data(), be_exec_version); - DORIS_CHECK_LE(end, serialized_buffer.data() + serialized_bytes); - serialized_buffer.resize(end - serialized_buffer.data()); - buf.write_binary(serialized_buffer); + void write(BufferWritable& buf, const IDataType& column_type, int be_exec_version) const { + const auto max_serialized_bytes = cast_set( + column_type.get_uncompressed_serialized_bytes(*column_data, be_exec_version)); + buf.resize(sizeof(UInt64) + max_serialized_bytes); + auto* serialized_data = buf.data() + sizeof(UInt64); + const char* end = nullptr; + try { + end = column_type.serialize(*column_data, serialized_data, be_exec_version); + } catch (...) { + buf.resize(0); + throw; + } + DORIS_CHECK_LE(end, serialized_data + max_serialized_bytes); + const auto serialized_bytes = static_cast(end - serialized_data); + memcpy(buf.data(), &serialized_bytes, sizeof(serialized_bytes)); + const auto frame_bytes = sizeof(serialized_bytes) + cast_set(serialized_bytes); + buf.resize(frame_bytes); + buf.add_offset(frame_bytes); } - void read(BufferReadable& buf) { + void read(BufferReadable& buf, const IDataType& column_type, int be_exec_version) { DORIS_CHECK(column_data->empty()); UInt64 serialized_bytes = 0; - buf.read_var_uint(serialized_bytes); + buf.read_binary(serialized_bytes); const auto* serialized_data = buf.data(); - const auto* end = column_type->deserialize(serialized_data, &column_data, be_exec_version); + const auto* end = column_type.deserialize(serialized_data, &column_data, be_exec_version); DORIS_CHECK_EQ(end, serialized_data + serialized_bytes); buf.add_offset(serialized_bytes); } @@ -311,14 +319,6 @@ class AggregateFunctionArrayAgg final DataTypePtr get_return_type() const override { return return_type; } - void create(AggregateDataPtr __restrict place) const override { - if constexpr (Data::PType == INVALID_TYPE) { - new (place) Data(this->argument_types, this->version); - } else { - Base::create(place); - } - } - void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num, Arena& arena) const override { this->data(place).add(*columns[0], row_num); @@ -330,12 +330,20 @@ class AggregateFunctionArrayAgg final } void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override { - this->data(place).write(buf); + if constexpr (Data::use_native_serde) { + this->data(place).write(buf, *this->argument_types[0], this->version); + } else { + this->data(place).write(buf); + } } void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf, Arena&) const override { - this->data(place).read(buf); + if constexpr (Data::use_native_serde) { + this->data(place).read(buf, *this->argument_types[0], this->version); + } else { + this->data(place).read(buf); + } } void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override { diff --git a/be/test/exprs/aggregate/agg_array_agg_test.cpp b/be/test/exprs/aggregate/agg_array_agg_test.cpp index 3f2d0c693abae3..84a3d823f68ab0 100644 --- a/be/test/exprs/aggregate/agg_array_agg_test.cpp +++ b/be/test/exprs/aggregate/agg_array_agg_test.cpp @@ -44,9 +44,12 @@ #include "core/types.h" #include "exprs/aggregate/agg_function_test.h" #include "exprs/aggregate/aggregate_function.h" +#include "exprs/aggregate/aggregate_function_array_agg.h" #include "exprs/aggregate/aggregate_function_simple_factory.h" #include "exprs/aggregate/aggregate_function_sort.h" #include "gtest/gtest_pred_impl.h" +#include "runtime/memory/mem_tracker_limiter.h" +#include "runtime/thread_context.h" namespace doris { class IColumn; @@ -121,6 +124,16 @@ void check_complex_array_agg_state(const DataTypePtr& data_type, EXPECT_TRUE(ColumnHelper::column_equal(std::move(result_column), std::move(expected_column))); } +class ThrowOnSerializeDataType final : public DataTypeString { +public: + char* serialize(const IColumn&, char* buf, int) const override { + *buf = 1; + throw Exception(ErrorCode::MEM_ALLOC_FAILED, "mock serialize allocation failure"); + } +}; + +static_assert(sizeof(AggregateFunctionArrayAggData) == sizeof(MutableColumnPtr)); + } // namespace TEST_F(AggregateFunctionArrayAggTest, test_array_agg_aint64) { @@ -350,6 +363,59 @@ TEST_F(AggregateFunctionArrayAggTest, foreach_complex_type_state_growth_and_roun EXPECT_TRUE(ColumnHelper::column_equal(std::move(result_column), std::move(expected_column))); } +TEST_F(AggregateFunctionArrayAggTest, complex_type_state_write_rolls_back_on_failure) { + auto data_type = std::make_shared(); + AggregateFunctionArrayAggData data(DataTypes {data_type}); + auto serialized_column = ColumnString::create(); + BufferWritable writer(*serialized_column); + writer.write_c_string("prefix"); + + EXPECT_THROW(data.write(writer, *data_type, BeExecVersionManager::get_newest_version()), + Exception); + EXPECT_EQ(serialized_column->get_chars().size(), 6); + + writer.commit(); + EXPECT_EQ(serialized_column->get_data_at(0).to_string(), "prefix"); +} + +TEST_F(AggregateFunctionArrayAggTest, large_complex_type_state_serialization_is_tracked) { + constexpr size_t PAYLOAD_BYTES = 2 * 1024 * 1024; + auto nullable_string = make_nullable(std::make_shared()); + auto nullable_inner_array = make_nullable(std::make_shared(nullable_string)); + auto input_type = std::make_shared(nullable_inner_array); + auto function = AggregateFunctionSimpleFactory::instance().get( + "array_agg_foreach", {input_type}, input_type, false, + BeExecVersionManager::get_newest_version(), {.is_foreach = true, .column_names = {}}); + ASSERT_NE(function, nullptr); + function->set_version(BeExecVersionManager::get_newest_version()); + + auto input_column = input_type->create_column(); + input_column->insert(array_field( + {array_field({Field::create_field(String(PAYLOAD_BYTES, 'x'))})})); + Arena arena; + AggregateFunctionGuard state(function.get()); + add_column_to_state(*function, state.data(), *input_column, arena); + + auto tracker = MemTrackerLimiter::create_shared(MemTrackerLimiter::Type::OTHER, + "ArrayAggLargeStateSerializationTest"); + auto switch_tracker = SwitchThreadMemTrackerLimiter(tracker); + thread_context()->thread_mem_tracker_mgr->flush_untracked_mem(); + const auto baseline = tracker->consumption(); + { + auto serialized_column = ColumnString::create(); + BufferWritable writer(*serialized_column); + function->serialize(state.data(), writer); + writer.commit(); + thread_context()->thread_mem_tracker_mgr->flush_untracked_mem(); + + const auto tracked_bytes = tracker->consumption() - baseline; + EXPECT_GT(tracked_bytes, PAYLOAD_BYTES); + EXPECT_GE(tracked_bytes, serialized_column->allocated_bytes()); + } + thread_context()->thread_mem_tracker_mgr->flush_untracked_mem(); + EXPECT_EQ(tracker->consumption(), baseline); +} + TEST(AggregateFunctionSortDataTest, merge_does_not_share_rhs_block) { auto data_type = std::make_shared(); Block prototype({ColumnWithTypeAndName(data_type->create_column(), data_type, "value"),