|
30 | 30 | #include "gtest/gtest.h" |
31 | 31 | #include "paimon/common/types/data_field.h" |
32 | 32 | #include "paimon/common/utils/fields_comparator.h" |
| 33 | +#include "paimon/core/core_options.h" |
33 | 34 | #include "paimon/core/io/concat_key_value_record_reader.h" |
| 35 | +#include "paimon/core/io/key_value_in_memory_record_reader.h" |
34 | 36 | #include "paimon/core/io/key_value_record_reader.h" |
| 37 | +#include "paimon/core/io/merged_key_value_record_reader.h" |
35 | 38 | #include "paimon/core/key_value.h" |
| 39 | +#include "paimon/core/mergetree/compact/aggregate/aggregate_merge_function.h" |
36 | 40 | #include "paimon/core/mergetree/compact/deduplicate_merge_function.h" |
37 | 41 | #include "paimon/core/mergetree/compact/reducer_merge_function_wrapper.h" |
38 | 42 | #include "paimon/core/mergetree/compact/sort_merge_reader_with_loser_tree.h" |
@@ -125,6 +129,48 @@ class SortMergeReaderTest : public testing::Test { |
125 | 129 | } |
126 | 130 | } |
127 | 131 |
|
| 132 | + template <typename SortMergeReaderType> |
| 133 | + void CheckSortMergeResultForAggregate( |
| 134 | + const std::vector<std::shared_ptr<arrow::StructArray>>& src_array_vec, |
| 135 | + const std::shared_ptr<FieldsComparator>& user_key_comparator, |
| 136 | + const std::shared_ptr<FieldsComparator>& user_defined_seq_comparator, |
| 137 | + const std::shared_ptr<arrow::Schema>& key_schema, |
| 138 | + const std::shared_ptr<arrow::Schema>& value_schema, |
| 139 | + const std::vector<std::string>& user_defined_sequence_fields, |
| 140 | + const std::vector<std::string>& primary_keys, const CoreOptions& core_options, |
| 141 | + const std::vector<KeyValue>& expected) const { |
| 142 | + for (auto batch_size : {1, 2, 3, 4, 100}) { |
| 143 | + ASSERT_OK_AND_ASSIGN( |
| 144 | + std::unique_ptr<AggregateMergeFunction> mfunc, |
| 145 | + AggregateMergeFunction::Create(value_schema, primary_keys, core_options)); |
| 146 | + auto merge_function_wrapper = |
| 147 | + std::make_shared<ReducerMergeFunctionWrapper>(std::move(mfunc)); |
| 148 | + std::vector<std::unique_ptr<KeyValueRecordReader>> merged_readers; |
| 149 | + |
| 150 | + int64_t last_seq = 0; |
| 151 | + std::vector<std::unique_ptr<KeyValueRecordReader>> readers; |
| 152 | + for (const auto& src_array : src_array_vec) { |
| 153 | + auto in_memory_reader = std::make_unique<KeyValueInMemoryRecordReader>( |
| 154 | + last_seq, src_array, std::vector<RecordBatch::RowKind>{}, primary_keys, |
| 155 | + user_defined_sequence_fields, /*sequence_fields_ascending=*/true, |
| 156 | + user_key_comparator, pool_); |
| 157 | + last_seq += src_array->length(); |
| 158 | + merged_readers.push_back(std::make_unique<MergedKeyValueRecordReader>( |
| 159 | + std::move(in_memory_reader), user_key_comparator, merge_function_wrapper)); |
| 160 | + } |
| 161 | + |
| 162 | + auto sort_merge_reader = std::make_unique<SortMergeReaderType>( |
| 163 | + std::move(merged_readers), user_key_comparator, user_defined_seq_comparator, |
| 164 | + merge_function_wrapper); |
| 165 | + ASSERT_OK_AND_ASSIGN( |
| 166 | + std::vector<KeyValue> results, |
| 167 | + (ReadResultCollector::CollectKeyValueResult< |
| 168 | + SortMergeReader, SortMergeReader::Iterator>(sort_merge_reader.get()))); |
| 169 | + KeyValueChecker::CheckResult(expected, results, key_schema->num_fields(), |
| 170 | + value_schema->num_fields()); |
| 171 | + } |
| 172 | + } |
| 173 | + |
128 | 174 | private: |
129 | 175 | std::shared_ptr<MemoryPool> pool_; |
130 | 176 | }; |
@@ -585,4 +631,86 @@ TEST_F(SortMergeReaderTest, TestSortMergeIn3WaysWithUserDefinedSeq) { |
585 | 631 | user_defined_seq_comparator, key_schema, value_schema, expected); |
586 | 632 | } |
587 | 633 |
|
| 634 | +TEST_F(SortMergeReaderTest, TestSortMergeWithAggMergeFunction) { |
| 635 | + // key: k0, user defined sequence field: ts, value: v0 |
| 636 | + // Format: [_SEQUENCE_NUMBER, _VALUE_KIND, k0, ts, v0] |
| 637 | + // Using sum aggregation: k0 uses primary-key agg, ts use last_value agg and v0 use sum agg. |
| 638 | + // |
| 639 | + // Reader1 (SEQUENCE_NUMBER 0..5): |
| 640 | + // [key=1,ts=1,v=1], [key=1,ts=2,v=2], [key=1,ts=3,v=3] |
| 641 | + // [key=1,ts=4,v=4], [key=2,ts=4,v=40], [key=2,ts=5,v=50] |
| 642 | + // Reader2 (SEQUENCE_NUMBER 6..11): |
| 643 | + // [key=1,ts=5,v=5], [key=1,ts=6,v=6], [key=2,ts=1,v=10] |
| 644 | + // [key=2,ts=2,v=20], [key=2,ts=3,v=30], [key=2,ts=6,v=60] |
| 645 | + // |
| 646 | + // With user_defined_seq_comparator on ts field, sort by key asc, then ts asc within same key: |
| 647 | + // key=1: ts=1(v=1), ts=2(v=2), ts=3(v=3), ts=4(v=4), ts=5(v=5), ts=6(v=6) |
| 648 | + // key=2: ts=1(v=10), ts=2(v=20), ts=3(v=30), ts=4(v=40), ts=5(v=50), ts=6(v=60) |
| 649 | + // |
| 650 | + // After sum aggregation: |
| 651 | + // key=1: k0=1, ts=last_value(1,2,3,4,5,6)=6, v0=sum(1,2,3,4,5,6)=21, seq=7 |
| 652 | + // key=2: k0=2, ts=last_value(1,2,3,4,5,6)=6, v0=sum(10,20,30,40,50,60)=210, seq=11 |
| 653 | + |
| 654 | + arrow::FieldVector fields = {arrow::field("k0", arrow::int32()), |
| 655 | + arrow::field("ts", arrow::int32()), |
| 656 | + arrow::field("v0", arrow::int32())}; |
| 657 | + |
| 658 | + auto data_fields = CreateDataField(fields); |
| 659 | + std::shared_ptr<arrow::Schema> key_schema = arrow::schema(arrow::FieldVector({fields[0]})); |
| 660 | + std::shared_ptr<arrow::Schema> value_schema = |
| 661 | + arrow::schema(arrow::FieldVector({fields[0], fields[1], fields[2]})); |
| 662 | + std::shared_ptr<arrow::DataType> src_type = arrow::struct_(fields); |
| 663 | + |
| 664 | + auto src_array1 = std::dynamic_pointer_cast<arrow::StructArray>( |
| 665 | + arrow::ipc::internal::json::ArrayFromJSON(src_type, R"([ |
| 666 | + [1, 1, 1], |
| 667 | + [1, 2, 2], |
| 668 | + [1, 3, 3], |
| 669 | + [1, 4, 4], |
| 670 | + [2, 4, 40], |
| 671 | + [2, 5, 50] |
| 672 | + ])") |
| 673 | + .ValueOrDie()); |
| 674 | + |
| 675 | + auto src_array2 = std::dynamic_pointer_cast<arrow::StructArray>( |
| 676 | + arrow::ipc::internal::json::ArrayFromJSON(src_type, R"([ |
| 677 | + [1, 5, 5], |
| 678 | + [1, 6, 6], |
| 679 | + [2, 1, 10], |
| 680 | + [2, 2, 20], |
| 681 | + [2, 3, 30], |
| 682 | + [2, 6, 60] |
| 683 | + ])") |
| 684 | + .ValueOrDie()); |
| 685 | + |
| 686 | + ASSERT_OK_AND_ASSIGN(std::shared_ptr<FieldsComparator> user_key_comparator, |
| 687 | + FieldsComparator::Create({data_fields[0]}, std::vector<int32_t>({0}), |
| 688 | + /*is_ascending_order=*/true)); |
| 689 | + // user_defined_seq_comparator based on ts field (index 1 in value schema {k0, ts, v0}) |
| 690 | + ASSERT_OK_AND_ASSIGN(std::shared_ptr<FieldsComparator> user_defined_seq_comparator, |
| 691 | + FieldsComparator::Create(data_fields, std::vector<int32_t>({1}), |
| 692 | + /*is_ascending_order=*/true)); |
| 693 | + // Configure sum aggregation for all non-primary-key fields |
| 694 | + std::string user_defined_sequence_field = "ts"; |
| 695 | + ASSERT_OK_AND_ASSIGN( |
| 696 | + CoreOptions core_options, |
| 697 | + CoreOptions::FromMap({{Options::FIELDS_DEFAULT_AGG_FUNC, "sum"}, |
| 698 | + {Options::SEQUENCE_FIELD, user_defined_sequence_field}})); |
| 699 | + |
| 700 | + // After sum aggregation, same-key rows are merged: |
| 701 | + // key=1: seq=7, k0=1, ts=6, v0=21 |
| 702 | + // key=2: seq=11, k0=2, ts=6, v0=210 |
| 703 | + std::vector<KeyValue> expected = |
| 704 | + KeyValueChecker::GenerateKeyValues({7, 11}, {{1}, {2}}, {{1, 6, 21}, {2, 6, 210}}, pool_); |
| 705 | + for (auto& kv : expected) { |
| 706 | + kv.level = KeyValue::UNKNOWN_LEVEL; |
| 707 | + } |
| 708 | + CheckSortMergeResultForAggregate<SortMergeReaderWithLoserTree>( |
| 709 | + {src_array1, src_array2}, user_key_comparator, user_defined_seq_comparator, key_schema, |
| 710 | + value_schema, {user_defined_sequence_field}, {"k0"}, core_options, expected); |
| 711 | + CheckSortMergeResultForAggregate<SortMergeReaderWithMinHeap>( |
| 712 | + {src_array1, src_array2}, user_key_comparator, user_defined_seq_comparator, key_schema, |
| 713 | + value_schema, {user_defined_sequence_field}, {"k0"}, core_options, expected); |
| 714 | +} |
| 715 | + |
588 | 716 | } // namespace paimon::test |
0 commit comments