Skip to content

Commit 4e0c39c

Browse files
committed
add test
1 parent 57a5a9d commit 4e0c39c

1 file changed

Lines changed: 128 additions & 0 deletions

File tree

src/paimon/core/mergetree/compact/sort_merge_reader_test.cpp

Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,9 +30,13 @@
3030
#include "gtest/gtest.h"
3131
#include "paimon/common/types/data_field.h"
3232
#include "paimon/common/utils/fields_comparator.h"
33+
#include "paimon/core/core_options.h"
3334
#include "paimon/core/io/concat_key_value_record_reader.h"
35+
#include "paimon/core/io/key_value_in_memory_record_reader.h"
3436
#include "paimon/core/io/key_value_record_reader.h"
37+
#include "paimon/core/io/merged_key_value_record_reader.h"
3538
#include "paimon/core/key_value.h"
39+
#include "paimon/core/mergetree/compact/aggregate/aggregate_merge_function.h"
3640
#include "paimon/core/mergetree/compact/deduplicate_merge_function.h"
3741
#include "paimon/core/mergetree/compact/reducer_merge_function_wrapper.h"
3842
#include "paimon/core/mergetree/compact/sort_merge_reader_with_loser_tree.h"
@@ -125,6 +129,48 @@ class SortMergeReaderTest : public testing::Test {
125129
}
126130
}
127131

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+
128174
private:
129175
std::shared_ptr<MemoryPool> pool_;
130176
};
@@ -585,4 +631,86 @@ TEST_F(SortMergeReaderTest, TestSortMergeIn3WaysWithUserDefinedSeq) {
585631
user_defined_seq_comparator, key_schema, value_schema, expected);
586632
}
587633

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+
588716
} // namespace paimon::test

0 commit comments

Comments
 (0)