diff --git a/be/src/exec/operator/aggregation_sink_operator.cpp b/be/src/exec/operator/aggregation_sink_operator.cpp index f4e74be9995c20..9e54518e348585 100644 --- a/be/src/exec/operator/aggregation_sink_operator.cpp +++ b/be/src/exec/operator/aggregation_sink_operator.cpp @@ -28,6 +28,7 @@ #include "exprs/aggregate/aggregate_function_count.h" #include "exprs/aggregate/aggregate_function_simple_factory.h" #include "exprs/vectorized_agg_fn.h" +#include "runtime/query_context.h" #include "runtime/runtime_profile.h" #include "runtime/thread_context.h" @@ -77,6 +78,7 @@ Status AggSinkLocalState::init(RuntimeState* state, LocalSinkStateInfo& info) { _memory_usage_container = ADD_COUNTER(custom_profile(), "MemoryUsageContainer", TUnit::BYTES); _memory_usage_arena = ADD_COUNTER(custom_profile(), "MemoryUsageArena", TUnit::BYTES); + _update_runtime_predicate_timer = ADD_TIMER(custom_profile(), "UpdateRuntimePredicateTime"); return Status::OK(); } @@ -225,6 +227,31 @@ Status AggSinkLocalState::_merge_with_serialized_key(Block* block) { } } +Status AggSinkLocalState::_update_runtime_predicate(RuntimeState* state) { + auto& p = Base::_parent->template cast(); + auto* query_ctx = state->get_query_ctx(); + if (!query_ctx->has_runtime_predicate(p.node_id()) || !Base::_shared_state->reach_limit) { + return Status::OK(); + } + + SCOPED_TIMER(_update_runtime_predicate_timer); + auto& predicate = query_ctx->get_runtime_predicate(p.node_id()); + if (!predicate.enable()) { + return Status::OK(); + } + + DCHECK(Base::_shared_state->do_sort_limit); + DCHECK(!Base::_shared_state->limit_columns.empty()); + DCHECK_GE(Base::_shared_state->limit_columns_min, 0); + Field new_top {PrimitiveType::TYPE_NULL}; + Base::_shared_state->limit_columns[0]->get(Base::_shared_state->limit_columns_min, new_top); + if (!new_top.is_null() && new_top != _old_top) { + RETURN_IF_ERROR(predicate.update(new_top)); + _old_top = std::move(new_top); + } + return Status::OK(); +} + size_t AggSinkLocalState::_memory_usage() const { if (0 == get_hash_table_size()) { return 0; @@ -895,6 +922,13 @@ Status AggSinkOperatorX::init(const TPlanNode& tnode, RuntimeState* state) { } } + auto* query_ctx = state->get_query_ctx(); + if (query_ctx->has_runtime_predicate(_node_id)) { + DORIS_CHECK(_do_sort_limit); + DORIS_CHECK_GT(_limit, 0); + query_ctx->get_runtime_predicate(_node_id).set_detected_source(); + } + return Status::OK(); } @@ -991,6 +1025,7 @@ Status AggSinkOperatorX::sink_impl(doris::RuntimeState* state, Block* in_block, local_state._shared_state->input_num_rows += in_block->rows(); if (in_block->rows() > 0) { RETURN_IF_ERROR(local_state._executor->execute(&local_state, in_block)); + RETURN_IF_ERROR(local_state._update_runtime_predicate(state)); local_state._executor->update_memusage(&local_state); COUNTER_SET(local_state._hash_table_size_counter, (int64_t)local_state.get_hash_table_size()); diff --git a/be/src/exec/operator/aggregation_sink_operator.h b/be/src/exec/operator/aggregation_sink_operator.h index 50b32ecb3645ac..da510fca84e84d 100644 --- a/be/src/exec/operator/aggregation_sink_operator.h +++ b/be/src/exec/operator/aggregation_sink_operator.h @@ -19,6 +19,7 @@ #include +#include "core/field.h" #include "exec/operator/operator.h" #include "runtime/exec_env.h" #include "runtime/runtime_profile.h" @@ -80,6 +81,7 @@ class AggSinkLocalState : public PipelineXSinkLocalState { Status _execute_with_serialized_key(Block* block); Status _merge_with_serialized_key(Block* block); void _update_memusage_with_serialized_key(); + Status _update_runtime_predicate(RuntimeState* state); template Status _execute_with_serialized_key_helper(Block* block); @@ -117,6 +119,7 @@ class AggSinkLocalState : public PipelineXSinkLocalState { RuntimeProfile::Counter* _serialize_key_arena_memory_usage = nullptr; RuntimeProfile::Counter* _memory_usage_container = nullptr; RuntimeProfile::Counter* _memory_usage_arena = nullptr; + RuntimeProfile::Counter* _update_runtime_predicate_timer = nullptr; bool _should_limit_output = false; @@ -130,6 +133,7 @@ class AggSinkLocalState : public PipelineXSinkLocalState { std::unique_ptr _executor = nullptr; int64_t _memory_usage_last_executing = 0; + Field _old_top {PrimitiveType::TYPE_NULL}; }; class AggSinkOperatorX MOCK_REMOVE(final) : public DataSinkOperatorX { diff --git a/be/src/exec/operator/streaming_aggregation_operator.cpp b/be/src/exec/operator/streaming_aggregation_operator.cpp index 31ae0d843e4385..0ac7fae2f61cfa 100644 --- a/be/src/exec/operator/streaming_aggregation_operator.cpp +++ b/be/src/exec/operator/streaming_aggregation_operator.cpp @@ -31,6 +31,7 @@ #include "exprs/aggregate/aggregate_function_simple_factory.h" #include "exprs/vectorized_agg_fn.h" #include "exprs/vslot_ref.h" +#include "runtime/query_context.h" namespace doris { class RuntimeState; @@ -72,6 +73,7 @@ Status StreamingAggLocalState::init(RuntimeState* state, LocalStateInfo& info) { _get_results_timer = ADD_TIMER(custom_profile(), "GetResultsTime"); _hash_table_iterate_timer = ADD_TIMER(custom_profile(), "HashTableIterateTime"); _insert_keys_to_column_timer = ADD_TIMER(custom_profile(), "InsertKeysToColumnTime"); + _update_runtime_predicate_timer = ADD_TIMER(custom_profile(), "UpdateRuntimePredicateTime"); return Status::OK(); } @@ -192,6 +194,31 @@ Status StreamingAggLocalState::do_pre_agg(RuntimeState* state, Block* input_bloc return Status::OK(); } +Status StreamingAggLocalState::_update_runtime_predicate(RuntimeState* state) { + auto& p = Base::_parent->template cast(); + auto* query_ctx = state->get_query_ctx(); + if (!query_ctx->has_runtime_predicate(p.node_id()) || need_do_sort_limit != 1) { + return Status::OK(); + } + + SCOPED_TIMER(_update_runtime_predicate_timer); + auto& predicate = query_ctx->get_runtime_predicate(p.node_id()); + if (!predicate.enable()) { + return Status::OK(); + } + + DCHECK(do_sort_limit); + DCHECK(!limit_columns.empty()); + DCHECK_GE(limit_columns_min, 0); + Field new_top {PrimitiveType::TYPE_NULL}; + limit_columns[0]->get(limit_columns_min, new_top); + if (!new_top.is_null() && new_top != _old_top) { + RETURN_IF_ERROR(predicate.update(new_top)); + _old_top = std::move(new_top); + } + return Status::OK(); +} + bool StreamingAggLocalState::_should_expand_preagg_hash_tables() { if (!_should_expand_hash_table) { return false; @@ -992,6 +1019,13 @@ Status StreamingAggOperatorX::init(const TPlanNode& tnode, RuntimeState* state) } } + auto* query_ctx = state->get_query_ctx(); + if (query_ctx->has_runtime_predicate(_node_id)) { + DORIS_CHECK(_do_sort_limit); + DORIS_CHECK_GT(_sort_limit, 0); + query_ctx->get_runtime_predicate(_node_id).set_detected_source(); + } + _op_name = "STREAMING_AGGREGATION_OPERATOR"; return Status::OK(); } @@ -1124,6 +1158,7 @@ Status StreamingAggOperatorX::push(RuntimeState* state, Block* in_block, bool eo if (in_block->rows() > 0) { RETURN_IF_ERROR( local_state.do_pre_agg(state, in_block, local_state._pre_aggregated_block.get())); + RETURN_IF_ERROR(local_state._update_runtime_predicate(state)); } in_block->clear_column_data(_child->row_desc().num_materialized_slots()); return Status::OK(); diff --git a/be/src/exec/operator/streaming_aggregation_operator.h b/be/src/exec/operator/streaming_aggregation_operator.h index b23c79477a1a2b..f9dcf24fc62d17 100644 --- a/be/src/exec/operator/streaming_aggregation_operator.h +++ b/be/src/exec/operator/streaming_aggregation_operator.h @@ -23,6 +23,7 @@ #include "common/status.h" #include "core/block/block.h" +#include "core/field.h" #include "exec/operator/operator.h" #include "runtime/runtime_profile.h" @@ -55,6 +56,7 @@ class StreamingAggLocalState MOCK_REMOVE(final) : public PipelineXLocalState need_computes; std::vector cmp_res; std::vector order_directions; diff --git a/be/test/exec/operator/agg_operator_group_by_limit_opt_test.cpp b/be/test/exec/operator/agg_operator_group_by_limit_opt_test.cpp index 0e5d2f83248422..d380c89d7a36e7 100644 --- a/be/test/exec/operator/agg_operator_group_by_limit_opt_test.cpp +++ b/be/test/exec/operator/agg_operator_group_by_limit_opt_test.cpp @@ -27,16 +27,43 @@ #include "exec/operator/assert_num_rows_operator.h" #include "exec/operator/mock_operator.h" #include "exec/operator/operator_helper.h" +#include "exec/pipeline/thrift_builder.h" +#include "runtime/query_context.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_agg_fn_evaluator.h" #include "testutil/mock/mock_slot_ref.h" namespace doris { +namespace { + +constexpr TPlanNodeId TOPN_FILTER_TARGET_NODE_ID = 20; + +void init_topn_runtime_predicate(MockRuntimeState& state, TPlanNodeId source_node_id, bool is_asc) { + auto target_expr = TRuntimeFilterDescBuilder::get_default_expr(); + target_expr.nodes[0].__set_type(create_type_desc(TYPE_BIGINT)); + + TTopnFilterDesc desc; + desc.__set_source_node_id(source_node_id); + desc.__set_is_asc(is_asc); + desc.__set_null_first(false); + desc.__set_target_node_id_to_target_expr({{TOPN_FILTER_TARGET_NODE_ID, target_expr}}); + state.get_query_ctx()->init_runtime_predicates({desc}); + + auto& predicate = state.get_query_ctx()->get_runtime_predicate(source_node_id); + predicate.set_detected_source(); + DORIS_CHECK(predicate.init_target(TOPN_FILTER_TARGET_NODE_ID, {}, -1).ok()); +} + +Field get_topn_runtime_predicate_value(MockRuntimeState& state, TPlanNodeId source_node_id) { + return state.get_query_ctx()->get_runtime_predicate(source_node_id).get_value(); +} + +} // namespace auto static init_sink_and_source(std::shared_ptr sink_op, std::shared_ptr source_op, OperatorContext& ctx) { - auto shared_state = sink_op->create_shared_state(); + auto shared_state = std::static_pointer_cast(sink_op->create_shared_state()); { auto local_state = AggSinkOperatorX::LocalState ::create_unique(sink_op.get(), &ctx.state); LocalSinkStateInfo info {.task_idx = 0, @@ -194,6 +221,90 @@ TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_need_finalize_without_ord } } +TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_desc_order_by_updates_runtime_predicate) { + OperatorContext ctx; + auto sink_op = std::make_shared(); + sink_op->_aggregate_evaluators.push_back(create_mock_agg_fn_evaluator( + ctx.pool, MockSlotRef::create_mock_contexts(1, std::make_shared()), + false, false)); + sink_op->_pool = &ctx.pool; + sink_op->_limit = 2; + sink_op->_probe_expr_ctxs = + MockSlotRef::create_mock_contexts(0, std::make_shared()); + sink_op->_order_directions = {-1}; + sink_op->_null_directions = {-1}; + sink_op->_do_sort_limit = true; + ASSERT_TRUE(sink_op->prepare(&ctx.state).ok()); + + auto source_op = std::make_shared(); + source_op->mock_row_descriptor.reset(new MockRowDescriptor { + {std::make_shared(), std::make_shared()}, &ctx.pool}); + source_op->_without_key = false; + source_op->_needs_finalize = true; + ASSERT_TRUE(source_op->prepare(&ctx.state).ok()); + + auto shared_state = init_sink_and_source(sink_op, source_op, ctx); + init_topn_runtime_predicate(ctx.state, sink_op->node_id(), false); + + Block first_block {ColumnHelper::create_column_with_name({1, 2, 3, 4}), + ColumnHelper::create_column_with_name({1, 2, 3, 4})}; + ASSERT_TRUE(sink_op->sink(&ctx.state, &first_block, false).ok()); + EXPECT_EQ(get_topn_runtime_predicate_value(ctx.state, sink_op->node_id()), + Field::create_field(3)); + + Block second_block {ColumnHelper::create_column_with_name({2, 5}), + ColumnHelper::create_column_with_name({2, 5})}; + ASSERT_TRUE(sink_op->sink(&ctx.state, &second_block, true).ok()); + EXPECT_EQ(get_topn_runtime_predicate_value(ctx.state, sink_op->node_id()), + Field::create_field(4)); +} + +TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, + test_null_boundary_does_not_update_runtime_predicate) { + OperatorContext ctx; + auto sink_op = std::make_shared(); + sink_op->_aggregate_evaluators.push_back(create_mock_agg_fn_evaluator( + ctx.pool, MockSlotRef::create_mock_contexts(1, std::make_shared()), + false, false)); + sink_op->_pool = &ctx.pool; + sink_op->_limit = 1; + sink_op->_probe_expr_ctxs = MockSlotRef::create_mock_contexts( + 0, std::make_shared(std::make_shared())); + sink_op->_order_directions = {1}; + sink_op->_null_directions = {-1}; + sink_op->_do_sort_limit = true; + ASSERT_TRUE(sink_op->prepare(&ctx.state).ok()); + + auto source_op = std::make_shared(); + source_op->mock_row_descriptor.reset(new MockRowDescriptor { + {std::make_shared(std::make_shared()), + std::make_shared()}, + &ctx.pool}); + source_op->_without_key = false; + source_op->_needs_finalize = true; + ASSERT_TRUE(source_op->prepare(&ctx.state).ok()); + + auto shared_state = init_sink_and_source(sink_op, source_op, ctx); + init_topn_runtime_predicate(ctx.state, sink_op->node_id(), true); + shared_state->reach_limit = true; + shared_state->do_sort_limit = true; + shared_state->limit_columns.emplace_back( + std::make_shared(std::make_shared())->create_column()); + shared_state->limit_columns[0]->insert_default(); + shared_state->limit_columns_min = 0; + + auto* local_state = static_cast(ctx.state.get_sink_local_state()); + ASSERT_TRUE(local_state->_update_runtime_predicate(&ctx.state).ok()); + auto& predicate = ctx.state.get_query_ctx()->get_runtime_predicate(sink_op->node_id()); + EXPECT_FALSE(predicate.has_value()); + + shared_state->limit_columns[0]->insert(Field::create_field(7)); + shared_state->limit_columns_min = 1; + ASSERT_TRUE(local_state->_update_runtime_predicate(&ctx.state).ok()); + ASSERT_TRUE(predicate.has_value()); + EXPECT_EQ(predicate.get_value(), Field::create_field(7)); +} + TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_need_finalize_with_order_by) { /* select column1, sum(column2) from test_table group by column1 order by column1 limit 3; @@ -258,6 +369,7 @@ TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_need_finalize_with_order_ EXPECT_TRUE(source_op->prepare(&ctx.state).ok()); auto shared_state = init_sink_and_source(sink_op, source_op, ctx); + init_topn_runtime_predicate(ctx.state, sink_op->node_id(), true); { Block block {ColumnHelper::create_column_with_name({1, 2, 3, 4, 5, 6}), @@ -266,6 +378,8 @@ TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_need_finalize_with_order_ std::cout << block.dump_data() << std::endl; auto st = sink_op->sink(&ctx.state, &block, false); EXPECT_TRUE(st.ok()) << st.msg(); + EXPECT_EQ(get_topn_runtime_predicate_value(ctx.state, sink_op->node_id()), + Field::create_field(3)); } { @@ -275,6 +389,8 @@ TEST_F(AggOperatorGroupByLimitOptTestWithGroupBy, test_need_finalize_with_order_ std::cout << block.dump_data() << std::endl; auto st = sink_op->sink(&ctx.state, &block, true); EXPECT_TRUE(st.ok()) << st.msg(); + EXPECT_EQ(get_topn_runtime_predicate_value(ctx.state, sink_op->node_id()), + Field::create_field(2)); } { diff --git a/be/test/exec/operator/streaming_agg_operator_test.cpp b/be/test/exec/operator/streaming_agg_operator_test.cpp index 6ab337513fd430..ca92fcdf75bd5b 100644 --- a/be/test/exec/operator/streaming_agg_operator_test.cpp +++ b/be/test/exec/operator/streaming_agg_operator_test.cpp @@ -29,6 +29,8 @@ #include "exec/operator/mock_operator.h" #include "exec/operator/operator_helper.h" #include "exec/operator/streaming_aggregation_operator.h" +#include "exec/pipeline/thrift_builder.h" +#include "runtime/query_context.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_agg_fn_evaluator.h" #include "testutil/mock/mock_runtime_state.h" @@ -36,6 +38,27 @@ #include "util/jsonb_document.h" namespace doris { +namespace { + +constexpr TPlanNodeId TOPN_FILTER_TARGET_NODE_ID = 20; + +void init_topn_runtime_predicate(MockRuntimeState& state, TPlanNodeId source_node_id) { + auto target_expr = TRuntimeFilterDescBuilder::get_default_expr(); + target_expr.nodes[0].__set_type(create_type_desc(TYPE_BIGINT)); + + TTopnFilterDesc desc; + desc.__set_source_node_id(source_node_id); + desc.__set_is_asc(true); + desc.__set_null_first(false); + desc.__set_target_node_id_to_target_expr({{TOPN_FILTER_TARGET_NODE_ID, target_expr}}); + state.get_query_ctx()->init_runtime_predicates({desc}); + + auto& predicate = state.get_query_ctx()->get_runtime_predicate(source_node_id); + predicate.set_detected_source(); + DORIS_CHECK(predicate.init_target(TOPN_FILTER_TARGET_NODE_ID, {}, -1).ok()); +} + +} // namespace struct MockStreamingAggOperatorX : public StreamingAggOperatorX { MockStreamingAggOperatorX() = default; @@ -155,6 +178,51 @@ TEST_F(StreamingAggOperatorTest, test1) { { EXPECT_TRUE(local_state->close(state.get()).ok()); } } +TEST_F(StreamingAggOperatorTest, topn_runtime_predicate_updates_in_hash_and_passthrough_paths) { + op->_aggregate_evaluators.push_back(create_mock_agg_fn_evaluator( + pool, MockSlotRef::create_mock_contexts(1, std::make_shared()), false, + false)); + op->_pool = &pool; + op->_needs_finalize = false; + op->_sort_limit = 3; + op->_do_sort_limit = true; + op->_order_directions = {1}; + op->_null_directions = {1}; + + ASSERT_TRUE(op->set_child(child_op)); + ASSERT_TRUE(op->prepare(state.get()).ok()); + op->_probe_expr_ctxs = MockSlotRef::create_mock_contexts(0, std::make_shared()); + + auto local_state_ptr = std::make_unique(state.get(), op.get()); + LocalStateInfo info {.parent_profile = &profile, + .scan_ranges = {}, + .shared_state = nullptr, + .shared_state_map = {}, + .task_idx = 0}; + ASSERT_TRUE(local_state_ptr->init(state.get(), info).ok()); + state->resize_op_id_to_local_state(-100); + state->emplace_local_state(op->operator_id(), std::move(local_state_ptr)); + local_state = + static_cast(state->get_local_state(op->operator_id())); + ASSERT_TRUE(local_state->open(state.get()).ok()); + init_topn_runtime_predicate(*state, op->node_id()); + + Block first_block {ColumnHelper::create_column_with_name({1, 2, 3, 4, 5, 6}), + ColumnHelper::create_column_with_name({1, 2, 3, 4, 5, 6})}; + ASSERT_TRUE(op->push(state.get(), &first_block, false).ok()); + auto& predicate = state->get_query_ctx()->get_runtime_predicate(op->node_id()); + ASSERT_TRUE(predicate.has_value()); + EXPECT_EQ(predicate.get_value(), Field::create_field(3)); + + local_state->should_not_do_pre_agg = true; + Block second_block {ColumnHelper::create_column_with_name({0, 2, 7, 8, 9, 10}), + ColumnHelper::create_column_with_name({0, 2, 7, 8, 9, 10})}; + ASSERT_TRUE(op->push(state.get(), &second_block, true).ok()); + EXPECT_EQ(predicate.get_value(), Field::create_field(2)); + + ASSERT_TRUE(local_state->close(state.get()).ok()); +} + TEST_F(StreamingAggOperatorTest, require_hash_shuffle_after_non_hash_local_exchange) { state->_query_options.__set_enable_local_exchange_before_agg(false); op->_needs_finalize = false; diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/FileScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/FileScanNode.java index 6874d929b67a0a..bcfae6406e1d48 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/FileScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/FileScanNode.java @@ -189,7 +189,7 @@ public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { if (useTopnFilter()) { String topnFilterSources = String.join(",", - topnFilterSortNodes.stream() + topnFilterSourceNodes.stream() .map(node -> node.getId().asInt() + "").collect(Collectors.toList())); output.append(prefix).append("TOPN OPT:").append(topnFilterSources).append("\n"); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/PluginDrivenScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/PluginDrivenScanNode.java index 7710b46224cf67..86fb0f92ad112a 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/PluginDrivenScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/scan/PluginDrivenScanNode.java @@ -571,7 +571,7 @@ public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { } if (useTopnFilter()) { String topnFilterSources = String.join(",", - topnFilterSortNodes.stream() + topnFilterSourceNodes.stream() .map(node -> node.getId().asInt() + "").collect(Collectors.toList())); output.append(prefix).append("TOPN OPT:").append(topnFilterSources).append("\n"); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/glue/translator/PhysicalPlanTranslator.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/glue/translator/PhysicalPlanTranslator.java index 09c8bc2462835e..18a5a15ec00bf7 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/glue/translator/PhysicalPlanTranslator.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/glue/translator/PhysicalPlanTranslator.java @@ -1304,6 +1304,11 @@ public PlanFragment visitPhysicalHashAggregate( } else { aggregationNode.setSortByGroupKey(null); } + if (context.getTopnFilterContext().isTopnFilterSource(aggregate)) { + context.getTopnFilterContext().translateSource(aggregate, aggregationNode); + TopnFilter filter = context.getTopnFilterContext().getTopnFilter(aggregate); + aggregationNode.setTopnFilterTargets(translateTopnFilterTargets(filter)); + } setPlanRoot(inputPlanFragment, aggregationNode, aggregate); if (aggregate.getStats() != null) { aggregationNode.setCardinality((long) aggregate.getStats().getRowCount()); @@ -2644,18 +2649,7 @@ public PlanFragment visitPhysicalTopN(PhysicalTopN topN, PlanTra if (context.getTopnFilterContext().isTopnFilterSource(topN)) { context.getTopnFilterContext().translateSource(topN, sortNode); TopnFilter filter = context.getTopnFilterContext().getTopnFilter(topN); - List> targets = new ArrayList<>(); - for (Entry entry : filter.legacyTargets.entrySet()) { - Set inputSlots = entry.getValue().getInputSlotRef(); - if (inputSlots.size() != 1) { - LOG.warn("topn filter targets error: " + inputSlots); - } else { - SlotRef slot = inputSlots.iterator().next(); - targets.add(Pair.of(entry.getKey().getId().asInt(), - (slot.getDesc().getId().asInt()))); - } - } - sortNode.setTopnFilterTargets(targets); + sortNode.setTopnFilterTargets(translateTopnFilterTargets(filter)); } // push sort to scan opt if (sortNode.getChild(0) instanceof OlapScanNode) { @@ -2708,6 +2702,20 @@ public PlanFragment visitPhysicalTopN(PhysicalTopN topN, PlanTra return inputFragment; } + private List> translateTopnFilterTargets(TopnFilter filter) { + List> targets = new ArrayList<>(); + for (Entry entry : filter.legacyTargets.entrySet()) { + Set inputSlots = entry.getValue().getInputSlotRef(); + if (inputSlots.size() != 1) { + LOG.warn("topn filter targets error: " + inputSlots); + } else { + SlotRef slot = inputSlots.iterator().next(); + targets.add(Pair.of(entry.getKey().getId().asInt(), slot.getDesc().getId().asInt())); + } + } + return targets; + } + @Override public PlanFragment visitPhysicalRepeat(PhysicalRepeat repeat, PlanTranslatorContext context) { PlanFragment inputPlanFragment = repeat.child(0).accept(this, context); diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopNScanOpt.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopNScanOpt.java index 9dcbfea57dd497..7df5a8229bb7a4 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopNScanOpt.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopNScanOpt.java @@ -19,9 +19,14 @@ import org.apache.doris.nereids.CascadesContext; import org.apache.doris.nereids.processor.post.TopnFilterPushDownVisitor.PushDownContext; +import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.nereids.trees.plans.AbstractPlan; import org.apache.doris.nereids.trees.plans.Plan; import org.apache.doris.nereids.trees.plans.SortPhase; import org.apache.doris.nereids.trees.plans.algebra.TopN; +import org.apache.doris.nereids.trees.plans.physical.PhysicalDistribute; +import org.apache.doris.nereids.trees.plans.physical.PhysicalHashAggregate; +import org.apache.doris.nereids.trees.plans.physical.PhysicalProject; import org.apache.doris.nereids.trees.plans.physical.PhysicalTopN; import org.apache.doris.nereids.types.DataType; @@ -40,14 +45,56 @@ public PhysicalTopN visitPhysicalTopN(PhysicalTopN aggregateSource = findAggregateSource(topN); + AbstractPlan source = aggregateSource == null ? topN : aggregateSource; + Expression probeExpr = aggregateSource == null + ? topN.getOrderKeys().get(0).getExpr() + : aggregateSource.getGroupByExpressions().get(0); TopnFilterPushDownVisitor.PushDownContext pushdownContext = new PushDownContext(topN, - topN.getOrderKeys().get(0).getExpr(), + source, probeExpr, topN.getOrderKeys().get(0).isNullFirst()); - topN.accept(pusher, pushdownContext); + boolean pushed = source.accept(pusher, pushdownContext); + if (!pushed && aggregateSource != null) { + pushdownContext = new PushDownContext(topN, topN, + topN.getOrderKeys().get(0).getExpr(), + topN.getOrderKeys().get(0).isNullFirst()); + topN.accept(pusher, pushdownContext); + } } return topN; } + private PhysicalHashAggregate findAggregateSource( + PhysicalTopN topN) { + Plan topNChild = topN.child(); + if (topNChild instanceof PhysicalProject) { + topNChild = topNChild.child(0); + } + if (!(topNChild instanceof PhysicalHashAggregate)) { + return null; + } + + PhysicalHashAggregate upperAggregate = + (PhysicalHashAggregate) topNChild; + if (upperAggregate.getTopnPushInfo() == null) { + return null; + } + + Plan aggregateChild = upperAggregate.child(); + if (aggregateChild instanceof PhysicalDistribute) { + aggregateChild = aggregateChild.child(0); + } + if (aggregateChild instanceof PhysicalHashAggregate) { + PhysicalHashAggregate lowerAggregate = + (PhysicalHashAggregate) aggregateChild; + if (lowerAggregate.getTopnPushInfo() != null) { + return lowerAggregate; + } + return null; + } + return upperAggregate; + } + boolean checkTopN(TopN topN) { if (!(topN instanceof PhysicalTopN)) { return false; diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterContext.java index d52c5b71ac4367..08fc5a46abff6c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterContext.java @@ -21,11 +21,12 @@ import org.apache.doris.nereids.glue.translator.ExpressionTranslator; import org.apache.doris.nereids.glue.translator.PlanTranslatorContext; import org.apache.doris.nereids.trees.expressions.Expression; -import org.apache.doris.nereids.trees.plans.ObjectId; +import org.apache.doris.nereids.trees.plans.AbstractPlan; import org.apache.doris.nereids.trees.plans.algebra.TopN; import org.apache.doris.nereids.trees.plans.physical.PhysicalLazyMaterializeOlapScan; import org.apache.doris.nereids.trees.plans.physical.PhysicalRelation; import org.apache.doris.nereids.trees.plans.physical.TopnFilter; +import org.apache.doris.planner.PlanNode; import org.apache.doris.planner.ScanNode; import org.apache.doris.planner.SortNode; @@ -40,26 +41,26 @@ * topN runtime filter context */ public class TopnFilterContext { - private final Map filters = Maps.newHashMap(); + private final Map filters = Maps.newHashMap(); /** * add topN filter */ - public void addTopnFilter(TopN topn, PhysicalRelation scan, Expression expr) { - TopnFilter filter = filters.get(topn.getObjectId()); + public void addTopnFilter(TopN topn, AbstractPlan source, PhysicalRelation scan, Expression expr) { + TopnFilter filter = filters.get(source.getId()); if (filter == null) { - filters.put(topn.getObjectId(), new TopnFilter(topn, scan, expr)); + filters.put(source.getId(), new TopnFilter(topn, source, scan, expr)); } else { filter.addTarget(scan, expr); } } - public boolean isTopnFilterSource(TopN topn) { - return filters.containsKey(topn.getObjectId()); + public boolean isTopnFilterSource(AbstractPlan source) { + return filters.containsKey(source.getId()); } - public TopnFilter getTopnFilter(TopN topn) { - return filters.get(topn.getObjectId()); + public TopnFilter getTopnFilter(AbstractPlan source) { + return filters.get(source.getId()); } public List getTopnFilters() { @@ -91,16 +92,18 @@ private void translateTarget(TopnFilter filter, PhysicalRelation relation, ScanN /** * translate topn-filter */ - public void translateSource(TopN topn, SortNode sortNode) { - TopnFilter filter = filters.get(topn.getObjectId()); + public void translateSource(AbstractPlan source, PlanNode legacySourceNode) { + TopnFilter filter = filters.get(source.getId()); if (filter == null) { return; } - filter.legacySortNode = sortNode; - sortNode.setUseTopnOpt(true); + filter.legacySourceNode = legacySourceNode; + if (legacySourceNode instanceof SortNode) { + ((SortNode) legacySourceNode).setUseTopnOpt(true); + } Preconditions.checkArgument(!filter.legacyTargets.isEmpty(), "missing targets: " + filter); for (ScanNode scan : filter.legacyTargets.keySet()) { - scan.addTopnFilterSortNode(sortNode); + scan.addTopnFilterSourceNode(legacySourceNode); } } @@ -112,8 +115,8 @@ public String toString() { String indent = " "; String arrow = " -> "; builder.append("filters:\n"); - for (ObjectId topnId : filters.keySet()) { - builder.append(indent).append(arrow).append(filters.get(topnId)).append("\n"); + for (Integer sourceId : filters.keySet()) { + builder.append(indent).append(arrow).append(filters.get(sourceId)).append("\n"); } return builder.toString(); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterPushDownVisitor.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterPushDownVisitor.java index bd72b8aaaefce6..b58399dd031952 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterPushDownVisitor.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/TopnFilterPushDownVisitor.java @@ -24,12 +24,14 @@ import org.apache.doris.nereids.trees.expressions.Slot; import org.apache.doris.nereids.trees.expressions.functions.scalar.Nullable; import org.apache.doris.nereids.trees.expressions.visitor.ExpressionVisitors; +import org.apache.doris.nereids.trees.plans.AbstractPlan; import org.apache.doris.nereids.trees.plans.Plan; import org.apache.doris.nereids.trees.plans.algebra.TopN; import org.apache.doris.nereids.trees.plans.algebra.Union; import org.apache.doris.nereids.trees.plans.physical.PhysicalCTEAnchor; import org.apache.doris.nereids.trees.plans.physical.PhysicalCTEProducer; import org.apache.doris.nereids.trees.plans.physical.PhysicalFileScan; +import org.apache.doris.nereids.trees.plans.physical.PhysicalHashAggregate; import org.apache.doris.nereids.trees.plans.physical.PhysicalHashJoin; import org.apache.doris.nereids.trees.plans.physical.PhysicalLazyMaterializeOlapScan; import org.apache.doris.nereids.trees.plans.physical.PhysicalNestedLoopJoin; @@ -66,16 +68,18 @@ public TopnFilterPushDownVisitor(TopnFilterContext topnFilterContext) { public static class PushDownContext { final Expression probeExpr; final TopN topn; + final AbstractPlan source; final boolean nullsFirst; - public PushDownContext(TopN topn, Expression probeExpr, boolean nullsFirst) { + public PushDownContext(TopN topn, AbstractPlan source, Expression probeExpr, boolean nullsFirst) { this.topn = topn; + this.source = source; this.probeExpr = probeExpr; this.nullsFirst = nullsFirst; } public PushDownContext withNewProbeExpression(Expression newProbe) { - return new PushDownContext(topn, newProbe, nullsFirst); + return new PushDownContext(topn, source, newProbe, nullsFirst); } } @@ -158,12 +162,24 @@ public Boolean visitPhysicalCTEProducer(PhysicalCTEProducer anch @Override public Boolean visitPhysicalTopN(PhysicalTopN topn, PushDownContext ctx) { - if (topn.equals(ctx.topn)) { + if (topn.getId() == ctx.source.getId()) { return topn.child().accept(this, ctx); } return false; } + @Override + public Boolean visitPhysicalHashAggregate( + PhysicalHashAggregate aggregate, PushDownContext ctx) { + if (ctx.source instanceof PhysicalHashAggregate) { + if (aggregate.getId() == ctx.source.getId()) { + return aggregate.child().accept(this, ctx); + } + return false; + } + return visit(aggregate, ctx); + } + @Override public Boolean visitPhysicalWindow(PhysicalWindow window, PushDownContext ctx) { return false; @@ -248,7 +264,7 @@ public Boolean visitPhysicalRelation(PhysicalRelation relation, PushDownContext || Math.max(relation.getStats().getRowCount(), 1) * ConnectContext.get().getSessionVariable().topnFilterRatio > ctx.topn.getLimit() + ctx.topn.getOffset()) { - topnFilterContext.addTopnFilter(ctx.topn, relation, ctx.probeExpr); + topnFilterContext.addTopnFilter(ctx.topn, ctx.source, relation, ctx.probeExpr); return true; } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/physical/TopnFilter.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/physical/TopnFilter.java index ab221e55506527..9ff48a1065cba0 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/physical/TopnFilter.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/physical/TopnFilter.java @@ -20,9 +20,10 @@ import org.apache.doris.analysis.Expr; import org.apache.doris.analysis.ExprToThriftVisitor; import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.nereids.trees.plans.AbstractPlan; import org.apache.doris.nereids.trees.plans.algebra.TopN; +import org.apache.doris.planner.PlanNode; import org.apache.doris.planner.ScanNode; -import org.apache.doris.planner.SortNode; import org.apache.doris.thrift.TTopnFilterDesc; import com.google.common.collect.Maps; @@ -34,12 +35,14 @@ */ public class TopnFilter { public TopN topn; - public SortNode legacySortNode; + public AbstractPlan source; + public PlanNode legacySourceNode; public Map targets = Maps.newHashMap(); public Map legacyTargets = Maps.newHashMap(); - public TopnFilter(TopN topn, PhysicalRelation rel, Expression expr) { + public TopnFilter(TopN topn, AbstractPlan source, PhysicalRelation rel, Expression expr) { this.topn = topn; + this.source = source; targets.put(rel, expr); } @@ -54,7 +57,7 @@ public boolean hasTargetRelation(PhysicalRelation rel) { @Override public String toString() { StringBuilder builder = new StringBuilder(); - builder.append(topn).append("->[ "); + builder.append(source).append("->[ "); for (PhysicalRelation rel : targets.keySet()) { builder.append("(").append(rel).append(":").append(targets.get(rel)).append(") "); } @@ -67,7 +70,7 @@ public String toString() { */ public TTopnFilterDesc toThrift() { TTopnFilterDesc tFilter = new TTopnFilterDesc(); - tFilter.setSourceNodeId(legacySortNode.getId().asInt()); + tFilter.setSourceNodeId(legacySourceNode.getId().asInt()); tFilter.setIsAsc(topn.getOrderKeys().get(0).isAsc()); tFilter.setNullFirst(topn.getOrderKeys().get(0).isNullFirst()); for (ScanNode scan : legacyTargets.keySet()) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/AggregationNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/AggregationNode.java index 786e026660af39..442c59a9edd609 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/AggregationNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/AggregationNode.java @@ -66,6 +66,7 @@ public class AggregationNode extends PlanNode { private boolean useStreamingPreagg; private SortInfo sortByGroupKey; + private List> topnFilterTargets; private boolean queryCacheCandidate; @@ -242,6 +243,9 @@ public String getNodeExplainString(String detailPrefix, TExplainLevel detailLeve output.append(detailPrefix).append("having: ").append(getExplainString(conjuncts)).append("\n"); } output.append(detailPrefix).append("sortByGroupKey:").append(sortByGroupKey != null).append("\n"); + if (topnFilterTargets != null) { + output.append(detailPrefix).append("TOPN filter targets: ").append(topnFilterTargets).append("\n"); + } output.append(detailPrefix).append(String.format( "cardinality=%,d", cardinality)).append("\n"); return output.toString(); @@ -265,6 +269,10 @@ public void setSortByGroupKey(SortInfo sortByGroupKey) { this.sortByGroupKey = sortByGroupKey; } + public void setTopnFilterTargets(List> topnFilterTargets) { + this.topnFilterTargets = topnFilterTargets; + } + public boolean isQueryCacheCandidate() { return queryCacheCandidate; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/DataGenScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/DataGenScanNode.java index d17981c9177324..db8d76027fcaf1 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/DataGenScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/DataGenScanNode.java @@ -121,7 +121,7 @@ public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { output.append(prefix).append("table value function: ").append(tvf.getDataGenFunctionName()).append("\n"); if (useTopnFilter()) { String topnFilterSources = String.join(",", - topnFilterSortNodes.stream() + topnFilterSourceNodes.stream() .map(node -> node.getId().asInt() + "").collect(Collectors.toList())); output.append(prefix).append("TOPN OPT:").append(topnFilterSources).append("\n"); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java index d17c8b7dd63ba2..f5130cc9f1a0b2 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java @@ -1126,7 +1126,7 @@ public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { if (useTopnFilter()) { String topnFilterSources = String.join(",", - topnFilterSortNodes.stream() + topnFilterSourceNodes.stream() .map(node -> node.getId().asInt() + "").collect(Collectors.toList())); output.append(prefix).append("TOPN OPT:").append(topnFilterSources).append("\n"); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java index 1b075ee051f3b1..ab4e363fcd59e0 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java @@ -104,7 +104,7 @@ public abstract class ScanNode extends PlanNode implements SplitGenerator { private boolean hasPartitionPredicate = false; // support multi topn filter - protected final List topnFilterSortNodes = Lists.newArrayList(); + protected final List topnFilterSourceNodes = Lists.newArrayList(); protected TableSnapshot tableSnapshot; protected List columns; @@ -716,24 +716,24 @@ public static void setVisibleVersionForOlapScanNodes(List scanNodes) t protected void toThrift(TPlanNode msg) { // topn filter if (useTopnFilter()) { - List topnFilterSourceNodeIds = getTopnFilterSortNodes() + List topnFilterSourceNodeIds = getTopnFilterSourceNodes() .stream() - .map(sortNode -> sortNode.getId().asInt()) + .map(sourceNode -> sourceNode.getId().asInt()) .collect(Collectors.toList()); msg.setTopnFilterSourceNodeIds(topnFilterSourceNodeIds); } } - public void addTopnFilterSortNode(SortNode sortNode) { - topnFilterSortNodes.add(sortNode); + public void addTopnFilterSourceNode(PlanNode sourceNode) { + topnFilterSourceNodes.add(sourceNode); } - public List getTopnFilterSortNodes() { - return topnFilterSortNodes; + public List getTopnFilterSourceNodes() { + return topnFilterSourceNodes; } public boolean useTopnFilter() { - return !topnFilterSortNodes.isEmpty(); + return !topnFilterSourceNodes.isEmpty(); } public long getSelectedPartitionNum() { diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java b/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java index 82979d901db526..844d0ee4fa2356 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java @@ -3341,7 +3341,10 @@ public FragmentExecParams(PlanFragment fragment) { Map toThrift(int backendNum) { Set topnSortNodes = scanNodes.stream() - .flatMap(scanNode -> scanNode.getTopnFilterSortNodes().stream()).collect(Collectors.toSet()); + .flatMap(scanNode -> scanNode.getTopnFilterSourceNodes().stream()) + .filter(SortNode.class::isInstance) + .map(SortNode.class::cast) + .collect(Collectors.toSet()); topnSortNodes.forEach(SortNode::setHasRuntimePredicate); long memLimit = queryOptions.getMemLimit(); diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java b/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java index 1a239e3122a365..627cfb99bc8724 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java @@ -199,8 +199,10 @@ public static Map plansToThr static void setRuntimePredicateIfNeed(Collection scanNodes) { for (ScanNode scanNode : scanNodes) { - for (SortNode topnFilterSortNode : scanNode.getTopnFilterSortNodes()) { - topnFilterSortNode.setHasRuntimePredicate(); + for (PlanNode topnFilterSourceNode : scanNode.getTopnFilterSourceNodes()) { + if (topnFilterSourceNode instanceof SortNode) { + ((SortNode) topnFilterSourceNode).setHasRuntimePredicate(); + } } } } diff --git a/fe/fe-core/src/test/java/org/apache/doris/datasource/scan/PluginDrivenScanNodeVerboseExplainTest.java b/fe/fe-core/src/test/java/org/apache/doris/datasource/scan/PluginDrivenScanNodeVerboseExplainTest.java index 8146d902224d8c..6f489788d846a2 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/datasource/scan/PluginDrivenScanNodeVerboseExplainTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/datasource/scan/PluginDrivenScanNodeVerboseExplainTest.java @@ -93,7 +93,7 @@ private static PluginDrivenScanNode nodeForCatalogType(String catalogType) { Deencapsulation.setField(node, "conjuncts", new ArrayList<>()); Deencapsulation.setField(node, "scanRangeLocations", new ArrayList<>()); // useTopnFilter() runs at the method tail (common to both EXPLAIN paths) and derefs this list. - Deencapsulation.setField(node, "topnFilterSortNodes", new ArrayList<>()); + Deencapsulation.setField(node, "topnFilterSourceNodes", new ArrayList<>()); // Pre-seed the cache so getOrLoadScanNodeProperties() returns it without contacting the connector. Deencapsulation.setField(node, "scanNodeProperties", Collections.emptyMap()); // Pre-seed the isBatchMode cache so the gate's !isBatchMode() is deterministic (no computeBatchMode). diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/postprocess/TopNRuntimeFilterTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/postprocess/TopNRuntimeFilterTest.java index 84868c93dd7226..b8fbf18e1d8732 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/postprocess/TopNRuntimeFilterTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/postprocess/TopNRuntimeFilterTest.java @@ -18,6 +18,7 @@ package org.apache.doris.nereids.postprocess; import org.apache.doris.nereids.datasets.ssb.SSBTestBase; +import org.apache.doris.nereids.glue.translator.PhysicalPlanTranslator; import org.apache.doris.nereids.glue.translator.PlanTranslatorContext; import org.apache.doris.nereids.processor.post.PlanPostProcessors; import org.apache.doris.nereids.processor.post.TopnFilterContext; @@ -29,6 +30,7 @@ import org.apache.doris.nereids.trees.expressions.literal.IntegerLiteral; import org.apache.doris.nereids.trees.plans.Plan; import org.apache.doris.nereids.trees.plans.SortPhase; +import org.apache.doris.nereids.trees.plans.physical.PhysicalHashAggregate; import org.apache.doris.nereids.trees.plans.physical.PhysicalLazyMaterializeOlapScan; import org.apache.doris.nereids.trees.plans.physical.PhysicalPlan; import org.apache.doris.nereids.trees.plans.physical.PhysicalRelation; @@ -36,6 +38,8 @@ import org.apache.doris.nereids.trees.plans.physical.TopnFilter; import org.apache.doris.nereids.util.MemoPatternMatchSupported; import org.apache.doris.nereids.util.PlanChecker; +import org.apache.doris.planner.AggregationNode; +import org.apache.doris.planner.PlanFragment; import org.apache.doris.planner.ScanNode; import org.junit.jupiter.api.Assertions; @@ -87,14 +91,14 @@ public void testTranslateTopNFilterTargetThroughLazyMaterializeOlapScan() { Expression probeExpr = filter.targets.get(lazyScan); TopnFilterContext topnFilterContext = new TopnFilterContext(); - topnFilterContext.addTopnFilter(filter.topn, lazyScan.getScan(), probeExpr); + topnFilterContext.addTopnFilter(filter.topn, filter.source, lazyScan.getScan(), probeExpr); PlanTranslatorContext translatorContext = new PlanTranslatorContext(); probeExpr.getInputSlots().forEach(slot -> translatorContext .createSlotDesc(translatorContext.generateTupleDesc(), (SlotReference) slot)); ScanNode legacyScan = Mockito.mock(ScanNode.class); topnFilterContext.translateTarget(lazyScan, legacyScan, translatorContext); - Assertions.assertTrue(topnFilterContext.getTopnFilter(filter.topn).legacyTargets.containsKey(legacyScan)); + Assertions.assertTrue(topnFilterContext.getTopnFilter(filter.source).legacyTargets.containsKey(legacyScan)); } @Test @@ -111,6 +115,83 @@ public void testUseTopNRfForComplexCase() { Assertions.assertTrue(checker.getCascadesContext().getTopnFilterContext().isTopnFilterSource(localTopN)); } + @Test + public void testUseLowestAggregateAsTopNRuntimeFilterSource() { + String sql = "select c_nation, count(*) from customer group by c_nation limit 5"; + PlanChecker checker = PlanChecker.from(connectContext).analyze(sql) + .rewrite() + .implement(); + PhysicalPlan plan = checker.getPhysicalPlan(); + plan = new PlanPostProcessors(checker.getCascadesContext()).process(plan); + + TopnFilterContext filterContext = checker.getCascadesContext().getTopnFilterContext(); + TopnFilter filter = filterContext.getTopnFilters().stream() + .filter(f -> f.source instanceof PhysicalHashAggregate) + .findFirst() + .orElseThrow(() -> new AssertionError("aggregate topn filter source not found")); + PhysicalHashAggregate aggregateSource = (PhysicalHashAggregate) filter.source; + List> pushedAggregates = plan.collectToList( + node -> node instanceof PhysicalHashAggregate + && ((PhysicalHashAggregate) node).getTopnPushInfo() != null); + + Assertions.assertTrue(pushedAggregates.size() >= 2, plan.treeString()); + int lowestDepth = pushedAggregates.stream().mapToInt(Plan::depth).min().orElseThrow(); + Assertions.assertEquals(lowestDepth, aggregateSource.depth()); + Assertions.assertTrue(filterContext.isTopnFilterSource(aggregateSource)); + } + + @Test + public void testTranslateAggregateTopNRuntimeFilterSource() { + String sql = "select c_nation, count(*) from customer " + + "group by c_nation order by c_nation limit 5"; + PlanChecker checker = PlanChecker.from(connectContext).analyze(sql) + .rewrite() + .implement(); + PhysicalPlan plan = checker.getPhysicalPlan(); + plan = new PlanPostProcessors(checker.getCascadesContext()).process(plan); + + TopnFilter filter = checker.getCascadesContext().getTopnFilterContext().getTopnFilters().stream() + .filter(f -> f.source instanceof PhysicalHashAggregate) + .findFirst() + .orElseThrow(() -> new AssertionError("aggregate topn filter source not found")); + PlanFragment fragment = new PhysicalPlanTranslator( + new PlanTranslatorContext(checker.getCascadesContext())).translatePlan(plan); + + Assertions.assertNotNull(fragment); + Assertions.assertInstanceOf(AggregationNode.class, filter.legacySourceNode); + Assertions.assertEquals(filter.legacySourceNode.getId().asInt(), + filter.toThrift().getSourceNodeId()); + Assertions.assertFalse(filter.legacyTargets.isEmpty()); + for (ScanNode scanNode : filter.legacyTargets.keySet()) { + Assertions.assertTrue(scanNode.getTopnFilterSourceNodes().contains(filter.legacySourceNode)); + } + } + + @Test + public void testFallbackToSortSourceWhenPushTopNToAggregateDisabled() { + String sql = "select c_nation, count(*) from customer " + + "group by c_nation order by c_nation limit 5"; + boolean originalPushTopnToAgg = connectContext.getSessionVariable().pushTopnToAgg; + long originalTopnOptLimitThreshold = connectContext.getSessionVariable().topnOptLimitThreshold; + connectContext.getSessionVariable().pushTopnToAgg = false; + connectContext.getSessionVariable().topnOptLimitThreshold = 0; + try { + PlanChecker checker = PlanChecker.from(connectContext).analyze(sql) + .rewrite() + .implement(); + PhysicalPlan plan = checker.getPhysicalPlan(); + new PlanPostProcessors(checker.getCascadesContext()).process(plan); + + Assertions.assertFalse(checker.getCascadesContext().getTopnFilterContext() + .getTopnFilters().isEmpty()); + checker.getCascadesContext().getTopnFilterContext().getTopnFilters() + .forEach(filter -> Assertions.assertInstanceOf(PhysicalTopN.class, filter.source)); + } finally { + connectContext.getSessionVariable().pushTopnToAgg = originalPushTopnToAgg; + connectContext.getSessionVariable().topnOptLimitThreshold = originalTopnOptLimitThreshold; + } + } + @Test public void testNotUseTopNRfOnWindow() { String sql = "select rank() over (partition by c_nation order by c_custkey) " diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/CoordinatorTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/CoordinatorTest.java index 1cdc3163d909f1..3f47829b9e8988 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/CoordinatorTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/CoordinatorTest.java @@ -23,6 +23,7 @@ import org.apache.doris.nereids.NereidsPlanner; import org.apache.doris.nereids.trees.plans.physical.TopnFilter; import org.apache.doris.nereids.util.PlanChecker; +import org.apache.doris.planner.AggregationNode; import org.apache.doris.planner.PlanFragment; import org.apache.doris.planner.PlanFragmentId; import org.apache.doris.planner.ScanNode; @@ -43,6 +44,7 @@ import java.io.IOException; import java.lang.reflect.Field; +import java.util.Arrays; import java.util.Collections; import java.util.List; import java.util.Map; @@ -132,7 +134,9 @@ public void testTopnFilterDescsSharedAmongInstances() throws Exception { public void testFragmentExecParamsMarksNonOlapTopnFilterSource() { ScanNode scanNode = Mockito.mock(ScanNode.class); SortNode sortNode = Mockito.mock(SortNode.class); - Mockito.when(scanNode.getTopnFilterSortNodes()).thenReturn(Collections.singletonList(sortNode)); + AggregationNode aggregationNode = Mockito.mock(AggregationNode.class); + Mockito.when(scanNode.getTopnFilterSourceNodes()) + .thenReturn(Arrays.asList(sortNode, aggregationNode)); PlanFragment fragment = Mockito.mock(PlanFragment.class); Mockito.when(fragment.getFragmentId()).thenReturn(new PlanFragmentId(0)); @@ -149,6 +153,7 @@ public void testFragmentExecParamsMarksNonOlapTopnFilterSource() { fragParams.toThrift(0); Mockito.verify(sortNode).setHasRuntimePredicate(); + Mockito.verifyNoInteractions(aggregationNode); } private NereidsPlanner plan(String sql) throws IOException { diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/runtime/ThriftPlansBuilderTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/runtime/ThriftPlansBuilderTest.java index 62a83172f94194..850ba100e2af9f 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/runtime/ThriftPlansBuilderTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/runtime/ThriftPlansBuilderTest.java @@ -24,6 +24,7 @@ import org.apache.doris.nereids.trees.plans.distribute.worker.job.DefaultScanSource; import org.apache.doris.nereids.trees.plans.distribute.worker.job.LocalShuffleAssignedJob; import org.apache.doris.nereids.trees.plans.distribute.worker.job.UnassignedJob; +import org.apache.doris.planner.AggregationNode; import org.apache.doris.planner.PlanNodeId; import org.apache.doris.planner.RecursiveCteScanNode; import org.apache.doris.planner.ScanNode; @@ -44,11 +45,14 @@ public class ThriftPlansBuilderTest { public void testSetRuntimePredicateForNonOlapScanNode() { ScanNode scanNode = Mockito.mock(ScanNode.class); SortNode sortNode = Mockito.mock(SortNode.class); - Mockito.when(scanNode.getTopnFilterSortNodes()).thenReturn(Collections.singletonList(sortNode)); + AggregationNode aggregationNode = Mockito.mock(AggregationNode.class); + Mockito.when(scanNode.getTopnFilterSourceNodes()) + .thenReturn(Arrays.asList(sortNode, aggregationNode)); ThriftPlansBuilder.setRuntimePredicateIfNeed(Collections.singletonList(scanNode)); Mockito.verify(sortNode).setHasRuntimePredicate(); + Mockito.verifyNoInteractions(aggregationNode); } @Test diff --git a/regression-test/data/nereids_tpch_p0/tpch/push_topn_to_agg.out b/regression-test/data/nereids_tpch_p0/tpch/push_topn_to_agg.out index 0008027785ac5c..e0f582cfad0e8e 100644 --- a/regression-test/data/nereids_tpch_p0/tpch/push_topn_to_agg.out +++ b/regression-test/data/nereids_tpch_p0/tpch/push_topn_to_agg.out @@ -1,4 +1,24 @@ -- This file is automatically generated. You should know what you did if you want to edit this +-- !topn_runtime_filter_asc -- +1 0 +10 0 +11 0 +2 0 +4 0 +5 0 +7 0 +8 0 + +-- !topn_runtime_filter_desc -- +14990 0 +14991 0 +14993 0 +14994 0 +14996 0 +14997 0 +14999 0 +15000 0 + -- !shape_distinct_agg -- PhysicalResultSink --PhysicalLimit[GLOBAL] diff --git a/regression-test/suites/nereids_tpch_p0/tpch/push_topn_to_agg.groovy b/regression-test/suites/nereids_tpch_p0/tpch/push_topn_to_agg.groovy index 15e903d85e93de..e7bcdc61c6972b 100644 --- a/regression-test/suites/nereids_tpch_p0/tpch/push_topn_to_agg.groovy +++ b/regression-test/suites/nereids_tpch_p0/tpch/push_topn_to_agg.groovy @@ -18,9 +18,22 @@ */ suite("push_topn_to_agg") { + def checkAggregateTopnFilter = { String explainStr -> + def sourceMatcher = explainStr =~ /(?m)^\s*(\d+):VAGGREGATE \(update serialize\).*$/ + assertTrue(sourceMatcher.find()) + String sourceNodeId = sourceMatcher.group(1) + int scanStart = explainStr.indexOf(":VOlapScanNode", sourceMatcher.start()) + assertTrue(scanStart > sourceMatcher.start()) + assertTrue(explainStr.substring(sourceMatcher.start(), scanStart) + .contains("TOPN filter targets:")) + assertTrue(explainStr.contains("TOPN OPT:${sourceNodeId}")) + } + sql "set parallel_pipeline_task_num=2" String db = context.config.getDbNameByFile(new File(context.file.parent)) sql "use ${db}" + sql "set enable_bucketed_hash_agg=false" + sql "analyze table orders with sync" sql "set topn_opt_limit_threshold=1024" // limit -> agg // verify switch @@ -33,6 +46,7 @@ suite("push_topn_to_agg") { explain{ sql "select o_custkey, sum(o_shippriority) from orders group by o_custkey limit 4;" multiContains ("sortByGroupKey:true", 2) + check(checkAggregateTopnFilter) } // when apply this opt, trun off STREAMING @@ -46,8 +60,25 @@ suite("push_topn_to_agg") { explain{ sql "select o_custkey, sum(o_shippriority) from orders group by o_custkey order by o_custkey limit 8;" multiContains ("sortByGroupKey:true", 2) + check(checkAggregateTopnFilter) } + order_qt_topn_runtime_filter_asc """ + select o_custkey, sum(o_shippriority) + from orders + group by o_custkey + order by o_custkey + limit 8 + """ + + order_qt_topn_runtime_filter_desc """ + select o_custkey + 1, sum(o_shippriority) + from orders + group by o_custkey + 1 + order by o_custkey + 1 desc + limit 8 + """ + // order keys are part of group keys, // 1. adjust group keys (o_custkey, o_clerk) -> o_clerk, o_custkey // 2. append o_custkey to order key diff --git a/topn_agg_runtime_filter_implementation_plan.md b/topn_agg_runtime_filter_implementation_plan.md new file mode 100644 index 00000000000000..9675e676890f31 --- /dev/null +++ b/topn_agg_runtime_filter_implementation_plan.md @@ -0,0 +1,107 @@ +# GROUP BY AGG LIMIT 生成 TopN Runtime Filter 实施计划 + +## 背景与问题 + +普通 TopN 查询由 Sort Sink 在消费数据时维护 TopN 堆,并将第一排序列的当前边界发布为 Runtime Predicate。Scan 侧使用这个动态边界过滤数据,从而减少后续算子的输入量。 + +`GROUP BY ... LIMIT` 及 `GROUP BY ... ORDER BY group_key LIMIT` 已有另一条优化链路: + +1. `LimitAggToTopNAgg` 将符合条件的 `Limit + Aggregate` 改写为 `TopN + Aggregate`。 +2. `PushTopnToAgg` 把排序键和 `limit + offset` 下推给一阶段或本地/全局 Hash Aggregate。 +3. BE Aggregate 在消费输入时维护有界堆,并过滤不可能进入最终结果的新分组键。 + +当前 TopN Runtime Filter 的生产端固定为 SortNode。对于以下执行链路,SortNode 必须等阻塞式聚合结束后才能得到边界,此时 Scan 已经结束,无法通过 Runtime Predicate 减少 Scan 数据: + +```text +Scan -> Local Aggregate -> Exchange -> Global Aggregate -> TopN Sort +``` + +目标是复用 Aggregate 已有的有界堆,让最靠近 Scan 的一阶段或本地 Aggregate 直接成为 TopN Runtime Filter 生产端: + +```text +Scan <- Runtime Predicate <- Local/One-Phase Aggregate bounded heap +``` + +## 实现范围 + +- 只覆盖已经满足 `PushTopnToAgg` 条件、且 Aggregate 已持有 `TopnPushInfo` 的查询。 +- 不扩大 `LimitAggToTopNAgg` 的 SQL 适用范围。 +- 普通 TopN 查询继续使用 SortNode 作为 Runtime Filter 生产端。 +- Aggregate 场景只选择最靠近 Scan 的有效 Aggregate,避免同时生成无效的 Sort Filter。 +- 沿用现有 TopN Runtime Filter 的类型白名单、`topn_filter_ratio`、表达式下推、ASC/DESC 和 NULL 排序语义。 +- 多列排序只发布第一列的包含边界;第一列相等的行全部保留,由 Aggregate/TopN 继续处理后续排序列。 +- 同时支持普通 Hash Aggregate 和 Streaming Aggregate。 +- 保持 `TTopnFilterDesc` 协议不变,仅使用 Aggregate 节点 ID 作为 `source_node_id`。 + +## FE 修改 + +### 1. 泛化 TopN Filter 生产端 + +- 将 `TopnFilter` 和 `TopnFilterContext` 从仅支持 `PhysicalTopN -> SortNode` 泛化为支持物理 TopN 或 Physical Hash Aggregate 对应的 Legacy PlanNode。 +- Filter 仍保存原始 TopN 的排序方向、NULL 顺序和 limit 信息,生产端单独记录。 +- ScanNode 保存通用的 TopN Filter source node,并继续通过 `topn_filter_source_node_ids` 下发节点 ID。 + +### 2. 选择 Aggregate 生产端 + +- 在 `PushTopnToAgg` 完成 `TopnPushInfo` 标记后,由 `TopNScanOpt` 检查 TopN 下方聚合结构。 +- 两阶段聚合优先选择 Local Aggregate;单阶段聚合选择该 Aggregate。 +- 使用 Aggregate 第一 GROUP BY Key 作为下推起点。 +- 如果找不到符合条件的 Aggregate,则保持现有 SortNode 生产端行为。 + +### 3. 翻译和 Explain + +- 在 `PhysicalPlanTranslator` 翻译 Aggregate 时,将物理 Aggregate 与 Legacy `AggregationNode` 绑定为 Filter source。 +- Scan target 翻译逻辑继续复用现有表达式翻译流程。 +- SortNode source 仍强制选择 Heap Sort;AggregationNode 不需要该处理,因为其 TopN 堆已经由 `sortByGroupKey` 启用。 +- Explain 中展示 Scan 对应的 source node ID,并在 Aggregate 上展示 Filter targets,方便确认生产端和消费端。 + +## BE 修改 + +### 1. Hash Aggregate + +- Aggregate 初始化时,如果 QueryContext 中存在以当前 Aggregate 节点 ID 注册的 Runtime Predicate,则声明当前节点为生产端。 +- Aggregate TopN 堆首次有效后,从第一排序/分组列读取当前堆顶边界。 +- 每批输入处理完成后,仅在边界发生变化时调用现有 `RuntimePredicate::update`。 +- 使用现有 Runtime Predicate 的锁和单调收紧逻辑处理多个 Pipeline Instance 并发更新。 + +### 2. Streaming Aggregate + +- 使用与 Hash Aggregate 相同的生产端注册和边界发布语义。 +- 同时覆盖 Hash Table 聚合和 Streaming passthrough 分支维护的 TopN 堆。 +- 在堆尚未构建或边界为 NULL 时不发布过滤值,保持 Scan 侧全量通过。 + +### 3. 可观测性 + +- 为 Aggregate 增加 Runtime Predicate 更新时间指标,名称与 Sort Sink 的现有指标保持一致。 +- 保留 Scan 侧已有 TopN Filter source IDs 与过滤行数指标。 + +## 测试计划 + +### FE 单元测试 + +- `LIMIT -> AGG` 选择本地 Aggregate 而不是 SortNode 作为生产端。 +- 显式 `ORDER BY group_key LIMIT`、Project、单阶段和两阶段聚合。 +- 不兼容排序键、未支持数据类型、关闭 `push_topn_to_agg` 时保持原行为。 +- 验证翻译后的 descriptor source ID、Scan source ID 和 Aggregate targets。 + +### BE 单元测试 + +- Hash Aggregate 在多批输入后发布并收紧 ASC/DESC 边界。 +- Nullable Key 与 NULLS FIRST/LAST 不产生错误过滤。 +- Streaming Aggregate 发布相同边界。 +- 验证 Aggregate 输出结果与未启用 Runtime Filter 时一致。 + +### 回归测试 + +- 扩展 `nereids_tpch_p0/tpch/push_topn_to_agg`,检查 Explain 中 Aggregate source 和 Scan `TOPN OPT` 关联。 +- 使用有序查询生成确定性结果,覆盖 `LIMIT`、显式 `ORDER BY` 和表达式 GROUP BY。 +- `.out` 仅通过 `run-regression-test.sh` 自动生成。 + +## 验证与提交 + +1. 使用指定 clang-format 16 执行 C++ 格式化,并运行格式检查。 +2. 运行相关 FE UT、BE UT 和定向回归测试。 +3. 使用 `build.sh` 完成必要的 FE/BE 编译。 +4. BE 编译产生 compilation database 后,对修改的 C++ 文件运行 clang-tidy。 +5. 检查正确性、并发安全、节点生命周期、兼容性和 Explain 可观测性。 +6. 只暂存并提交本任务相关文件,不包含本地环境或其他未跟踪文件。