Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 116 additions & 0 deletions cpp/src/arrow/acero/sorted_merge_node_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -16,16 +16,22 @@
// under the License.

#include <gtest/gtest.h>
#include <limits>
#include <memory>
#include <vector>

#include "arrow/acero/exec_plan.h"
#include "arrow/acero/map_node.h"
#include "arrow/acero/options.h"
#include "arrow/acero/test_nodes.h"
#include "arrow/array/builder_base.h"
#include "arrow/array/builder_primitive.h"
#include "arrow/array/concatenate.h"
#include "arrow/compute/ordering.h"
#include "arrow/record_batch.h"
#include "arrow/result.h"
#include "arrow/scalar.h"
#include "arrow/status.h"
#include "arrow/table.h"
#include "arrow/testing/generator.h"
#include "arrow/testing/gtest_util.h"
Expand Down Expand Up @@ -83,4 +89,114 @@ TEST(SortedMergeNode, Basic) {
AssertArraysEqual(*expected_ts, *output_ts);
}

TEST(SortedMergeNode, TestSortedMergeTwoInputsWithBool) {
const int64_t row_count = (16 << 10); // 16k rows per input
// Not a multiple of 8, so that batches start at non-byte-aligned offsets
const int64_t batch_size = 1001;

// Create schema with int column A and bool column B
auto test_schema = arrow::schema(
{arrow::field("col_a", arrow::int32()), arrow::field("col_b", arrow::boolean())});

// Helper lambda to create table source with specific pattern
auto make_source = [&](int64_t cnt, int offset) -> arrow::Result<Declaration> {
// Create column A (int) - values from offset to offset+cnt-1
arrow::Int32Builder col_a_builder;
std::vector<int32_t> col_a_values;
col_a_values.reserve(cnt);
for (int64_t i = 0; i < cnt; ++i) {
col_a_values.push_back(static_cast<int32_t>(offset + i));
}
ARROW_RETURN_NOT_OK(col_a_builder.AppendValues(col_a_values));
std::shared_ptr<arrow::Array> col_a_arr;
ARROW_RETURN_NOT_OK(col_a_builder.Finish(&col_a_arr));

// Create column B (bool) - pattern: null if col_a % 7 == 0, otherwise
// true if col_a % 5 == 0, false otherwise
arrow::BooleanBuilder col_b_builder;
for (int64_t i = 0; i < cnt; ++i) {
int32_t a_value = static_cast<int32_t>(offset + i);
if (a_value % 7 == 0) {
ARROW_RETURN_NOT_OK(col_b_builder.AppendNull());
continue;
}
bool b_value = (a_value % 5 == 0);
ARROW_RETURN_NOT_OK(col_b_builder.Append(b_value));
}
std::shared_ptr<arrow::Array> col_b_arr;
ARROW_RETURN_NOT_OK(col_b_builder.Finish(&col_b_arr));

auto table = arrow::Table::Make(test_schema, {col_a_arr, col_b_arr});
auto table_source =
Declaration("table_source", TableSourceNodeOptions(table, batch_size));
return table_source;
};

ASSERT_OK_AND_ASSIGN(auto source1, make_source(row_count, 0));
ASSERT_OK_AND_ASSIGN(auto source2, make_source(row_count, 8192));

// Create sorted merge by column A
auto ops = OrderByNodeOptions(compute::Ordering({compute::SortKey("col_a")}));
Declaration sorted_merge{"sorted_merge", {source1, source2}, ops};

// Execute plan and collect result
ASSERT_OK_AND_ASSIGN(auto result_table,
arrow::acero::DeclarationToTable(sorted_merge, false));

ASSERT_TRUE(result_table != nullptr);

// Verify results
auto col_a = result_table->GetColumnByName("col_a");
auto col_b = result_table->GetColumnByName("col_b");
ASSERT_TRUE(col_a != nullptr);
ASSERT_TRUE(col_b != nullptr);

// Verify sorting and bool values
int32_t last_a_value = std::numeric_limits<int32_t>::min();
int64_t total_rows_checked = 0;

for (int i = 0; i < col_a->num_chunks(); i++) {
auto a_chunk = std::static_pointer_cast<arrow::Int32Array>(col_a->chunk(i));
auto b_chunk = std::static_pointer_cast<arrow::BooleanArray>(col_b->chunk(i));

ASSERT_EQ(a_chunk->length(), b_chunk->length())
<< "Column A and B must have same length in chunk " << i;

for (int64_t j = 0; j < a_chunk->length(); j++) {
ASSERT_FALSE(a_chunk->IsNull(j)) << "Column A should not have null values";

int32_t a_value = a_chunk->Value(j);

// Verify sorting by column A
ASSERT_GE(a_value, last_a_value)
<< "Values not sorted at chunk " << i << ", row " << j
<< ": current=" << a_value << ", last=" << last_a_value;
last_a_value = a_value;

// Verify bool validity: should be null if a_value % 7 == 0
bool expected_b_null = (a_value % 7 == 0);
ASSERT_EQ(b_chunk->IsNull(j), expected_b_null)
<< "Bool validity incorrect at chunk " << i << ", row " << j
<< ": col_a=" << a_value;

if (!expected_b_null) {
// Verify bool value correctness: should be true if a_value % 5 == 0
bool b_value = b_chunk->Value(j);
bool expected_b_value = (a_value % 5 == 0);
ASSERT_EQ(b_value, expected_b_value)
<< "Bool value incorrect at chunk " << i << ", row " << j
<< ": col_a=" << a_value << ", col_b=" << b_value
<< ", expected=" << expected_b_value;
}

total_rows_checked++;
}
}

ASSERT_EQ(last_a_value, 24575);

ASSERT_EQ(total_rows_checked, row_count * 2)
<< "Expected " << row_count * 2 << " rows after merge";
}

} // namespace arrow::acero
5 changes: 4 additions & 1 deletion cpp/src/arrow/acero/unmaterialized_table_internal.h
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,10 @@ class UnmaterializedCompositeTable {
builder.UnsafeAppendNull();
return Status::OK();
}
builder.UnsafeAppend(bit_util::GetBit(source->template GetValues<uint8_t>(1), row));

const int64_t bit_offset = source->offset + static_cast<int64_t>(row);
builder.UnsafeAppend(
bit_util::GetBit(source->template GetValues<uint8_t>(1, 0), bit_offset));
return Status::OK();
}

Expand Down
Loading