Coverage Report

Created: 2026-05-08 18:22

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
be/src/exprs/vectorized_agg_fn.cpp
Line
Count
Source
1
// Licensed to the Apache Software Foundation (ASF) under one
2
// or more contributor license agreements.  See the NOTICE file
3
// distributed with this work for additional information
4
// regarding copyright ownership.  The ASF licenses this file
5
// to you under the Apache License, Version 2.0 (the
6
// "License"); you may not use this file except in compliance
7
// with the License.  You may obtain a copy of the License at
8
//
9
//   http://www.apache.org/licenses/LICENSE-2.0
10
//
11
// Unless required by applicable law or agreed to in writing,
12
// software distributed under the License is distributed on an
13
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14
// KIND, either express or implied.  See the License for the
15
// specific language governing permissions and limitations
16
// under the License.
17
18
#include "exprs/vectorized_agg_fn.h"
19
20
#include <fmt/format.h>
21
#include <fmt/ranges.h> // IWYU pragma: keep
22
#include <gen_cpp/Exprs_types.h>
23
#include <gen_cpp/PlanNodes_types.h>
24
#include <glog/logging.h>
25
26
#include <memory>
27
#include <ostream>
28
#include <string_view>
29
30
#include "common/config.h"
31
#include "common/object_pool.h"
32
#include "core/block/block.h"
33
#include "core/block/column_with_type_and_name.h"
34
#include "core/block/materialize_block.h"
35
#include "core/data_type/data_type_agg_state.h"
36
#include "core/data_type/data_type_factory.hpp"
37
#include "exec/common/util.hpp"
38
#include "exprs/aggregate/aggregate_function_ai_agg.h"
39
#include "exprs/aggregate/aggregate_function_java_udaf.h"
40
#include "exprs/aggregate/aggregate_function_python_udaf.h"
41
#include "exprs/aggregate/aggregate_function_rpc.h"
42
#include "exprs/aggregate/aggregate_function_simple_factory.h"
43
#include "exprs/aggregate/aggregate_function_sort.h"
44
#include "exprs/aggregate/aggregate_function_state_merge.h"
45
#include "exprs/aggregate/aggregate_function_state_union.h"
46
#include "exprs/vexpr.h"
47
#include "exprs/vexpr_context.h"
48
49
static constexpr int64_t BE_VERSION_THAT_SUPPORT_NULLABLE_CHECK = 8;
50
51
namespace doris {
52
class RowDescriptor;
53
class Arena;
54
class BufferWritable;
55
class IColumn;
56
} // namespace doris
57
58
namespace doris {
59
60
template <class FunctionType>
61
AggregateFunctionPtr get_agg_state_function(const DataTypes& argument_types,
62
499
                                            DataTypePtr return_type) {
63
499
    return FunctionType::create(
64
499
            assert_cast<const DataTypeAggState*>(argument_types[0].get())->get_nested_function(),
65
499
            argument_types, return_type);
66
499
}
_ZN5doris22get_agg_state_functionINS_19AggregateStateUnionEEESt10shared_ptrINS_18IAggregateFunctionEERKSt6vectorIS2_IKNS_9IDataTypeEESaIS8_EES8_
Line
Count
Source
62
149
                                            DataTypePtr return_type) {
63
149
    return FunctionType::create(
64
149
            assert_cast<const DataTypeAggState*>(argument_types[0].get())->get_nested_function(),
65
149
            argument_types, return_type);
66
149
}
_ZN5doris22get_agg_state_functionINS_19AggregateStateMergeEEESt10shared_ptrINS_18IAggregateFunctionEERKSt6vectorIS2_IKNS_9IDataTypeEESaIS8_EES8_
Line
Count
Source
62
350
                                            DataTypePtr return_type) {
63
350
    return FunctionType::create(
64
350
            assert_cast<const DataTypeAggState*>(argument_types[0].get())->get_nested_function(),
65
350
            argument_types, return_type);
66
350
}
67
68
AggFnEvaluator::AggFnEvaluator(const TExprNode& desc, const bool without_key,
69
                               const bool is_window_function)
70
186k
        : _fn(desc.fn),
71
186k
          _is_merge(desc.agg_expr.is_merge_agg),
72
186k
          _without_key(without_key),
73
186k
          _is_window_function(is_window_function),
74
186k
          _data_type(DataTypeFactory::instance().create_data_type(
75
18.4E
                  desc.fn.ret_type, desc.__isset.is_nullable ? desc.is_nullable : true)) {
76
186k
    if (desc.agg_expr.__isset.param_types) {
77
186k
        const auto& param_types = desc.agg_expr.param_types;
78
186k
        for (const auto& param_type : param_types) {
79
149k
            _argument_types_with_sort.push_back(
80
149k
                    DataTypeFactory::instance().create_data_type(param_type));
81
149k
        }
82
186k
    }
83
186k
}
84
85
Status AggFnEvaluator::create(ObjectPool* pool, const TExpr& desc, const TSortInfo& sort_info,
86
                              const bool without_key, const bool is_window_function,
87
186k
                              AggFnEvaluator** result) {
88
186k
    *result =
89
186k
            pool->add(AggFnEvaluator::create_unique(desc.nodes[0], without_key, is_window_function)
90
186k
                              .release());
91
186k
    auto& agg_fn_evaluator = *result;
92
186k
    int node_idx = 0;
93
356k
    for (int i = 0; i < desc.nodes[0].num_children; ++i) {
94
169k
        ++node_idx;
95
169k
        VExprSPtr expr;
96
169k
        VExprContextSPtr ctx;
97
169k
        RETURN_IF_ERROR(VExpr::create_tree_from_thrift(desc.nodes, &node_idx, expr, ctx));
98
169k
        agg_fn_evaluator->_input_exprs_ctxs.push_back(ctx);
99
169k
    }
100
101
186k
    auto sort_size = sort_info.ordering_exprs.size();
102
186k
    auto real_arguments_size = agg_fn_evaluator->_argument_types_with_sort.size() - sort_size;
103
    // Child arguments contains [real arguments, order by arguments], we pass the arguments
104
    // to the order by functions
105
186k
    for (int i = 0; i < sort_size; ++i) {
106
91
        agg_fn_evaluator->_sort_description.emplace_back(real_arguments_size + i,
107
91
                                                         sort_info.is_asc_order[i] ? 1 : -1,
108
91
                                                         sort_info.nulls_first[i] ? -1 : 1);
109
91
    }
110
111
    // Pass the real arguments to get functions
112
336k
    for (int i = 0; i < real_arguments_size; ++i) {
113
149k
        agg_fn_evaluator->_real_argument_types.emplace_back(
114
149k
                agg_fn_evaluator->_argument_types_with_sort[i]);
115
149k
    }
116
186k
    return Status::OK();
117
186k
}
118
119
Status AggFnEvaluator::prepare(RuntimeState* state, const RowDescriptor& desc,
120
                               const SlotDescriptor* intermediate_slot_desc,
121
186k
                               const SlotDescriptor* output_slot_desc) {
122
186k
    DCHECK(intermediate_slot_desc != nullptr);
123
186k
    DCHECK(_intermediate_slot_desc == nullptr);
124
186k
    _output_slot_desc = output_slot_desc;
125
186k
    _intermediate_slot_desc = intermediate_slot_desc;
126
127
186k
    Status status = VExpr::prepare(_input_exprs_ctxs, state, desc);
128
186k
    RETURN_IF_ERROR(status);
129
130
186k
    DataTypes tmp_argument_types;
131
186k
    tmp_argument_types.reserve(_input_exprs_ctxs.size());
132
133
186k
    std::vector<std::string_view> child_expr_name;
134
135
    // prepare for argument
136
186k
    for (auto& _input_exprs_ctx : _input_exprs_ctxs) {
137
169k
        auto data_type = _input_exprs_ctx->root()->data_type();
138
169k
        tmp_argument_types.emplace_back(data_type);
139
169k
        child_expr_name.emplace_back(_input_exprs_ctx->root()->expr_name());
140
169k
    }
141
142
186k
    std::vector<std::string> column_names;
143
186k
    for (const auto& expr_ctx : _input_exprs_ctxs) {
144
169k
        const auto& root = expr_ctx->root();
145
169k
        if (!root->expr_name().empty() && !root->is_constant()) {
146
73.2k
            column_names.emplace_back(root->expr_name());
147
73.2k
        }
148
169k
    }
149
150
186k
    const DataTypes& argument_types =
151
186k
            _real_argument_types.empty() ? tmp_argument_types : _real_argument_types;
152
153
186k
    if (_fn.binary_type == TFunctionBinaryType::JAVA_UDF) {
154
115
        if (config::enable_java_support) {
155
115
            _function = AggregateJavaUdaf::create(_fn, argument_types, _data_type);
156
115
            RETURN_IF_ERROR(static_cast<AggregateJavaUdaf*>(_function.get())->check_udaf(_fn));
157
115
        } else {
158
0
            return Status::InternalError(
159
0
                    "Java UDAF is not enabled, you can change be config enable_java_support to "
160
0
                    "true and restart be.");
161
0
        }
162
186k
    } else if (_fn.binary_type == TFunctionBinaryType::PYTHON_UDF) {
163
673
        if (config::enable_python_udf_support) {
164
673
            _function = AggregatePythonUDAF::create(_fn, argument_types, _data_type);
165
673
            RETURN_IF_ERROR(static_cast<AggregatePythonUDAF*>(_function.get())->open());
166
673
            LOG(INFO) << fmt::format(
167
673
                    "Created Python UDAF: {}, runtime_version: {}, function_code: {}",
168
673
                    _fn.name.function_name, _fn.runtime_version, _fn.function_code);
169
673
        } else {
170
0
            return Status::InternalError(
171
0
                    "Python UDAF is not enabled, you can change be config "
172
0
                    "enable_python_udf_support to true and restart be.");
173
0
        }
174
186k
    } else if (_fn.binary_type == TFunctionBinaryType::RPC) {
175
0
        _function = AggregateRpcUdaf::create(_fn, argument_types, _data_type);
176
186k
    } else if (_fn.binary_type == TFunctionBinaryType::AGG_STATE) {
177
500
        if (argument_types.size() != 1) {
178
0
            return Status::InternalError("Agg state Function must input 1 argument but get {}",
179
0
                                         argument_types.size());
180
0
        }
181
500
        if (argument_types[0]->is_nullable()) {
182
0
            return Status::InternalError("Agg state function input type must be not nullable");
183
0
        }
184
500
        if (argument_types[0]->get_primitive_type() != PrimitiveType::TYPE_AGG_STATE) {
185
0
            return Status::InternalError(
186
0
                    "Agg state function input type must be agg_state but get {}",
187
0
                    argument_types[0]->get_family_name());
188
0
        }
189
190
500
        std::string type_function_name =
191
500
                assert_cast<const DataTypeAggState*>(argument_types[0].get())->get_function_name();
192
500
        if (type_function_name + AGG_UNION_SUFFIX == _fn.name.function_name) {
193
149
            if (_data_type->is_nullable()) {
194
0
                return Status::InternalError(
195
0
                        "Union function return type must be not nullable, real={}",
196
0
                        _data_type->get_name());
197
0
            }
198
149
            if (_data_type->get_primitive_type() != PrimitiveType::TYPE_AGG_STATE) {
199
0
                return Status::InternalError(
200
0
                        "Union function return type must be AGG_STATE, real={}",
201
0
                        _data_type->get_name());
202
0
            }
203
149
            _function = get_agg_state_function<AggregateStateUnion>(argument_types, _data_type);
204
351
        } else if (type_function_name + AGG_MERGE_SUFFIX == _fn.name.function_name) {
205
350
            auto type = assert_cast<const DataTypeAggState*>(argument_types[0].get())
206
350
                                ->get_nested_function()
207
350
                                ->get_return_type();
208
350
            if (!type->equals(*_data_type)) {
209
0
                return Status::InternalError("{}'s expect return type is {}, but input {}",
210
0
                                             argument_types[0]->get_name(), type->get_name(),
211
0
                                             _data_type->get_name());
212
0
            }
213
350
            _function = get_agg_state_function<AggregateStateMerge>(argument_types, _data_type);
214
350
        } else {
215
1
            return Status::InternalError("{} not match function {}", argument_types[0]->get_name(),
216
1
                                         _fn.name.function_name);
217
1
        }
218
185k
    } else {
219
185k
        const bool is_foreach =
220
185k
                AggregateFunctionSimpleFactory::is_foreach(_fn.name.function_name) ||
221
185k
                AggregateFunctionSimpleFactory::is_foreachv2(_fn.name.function_name);
222
        // Here, only foreachv1 needs special treatment, and v2 can follow the normal code logic.
223
185k
        if (AggregateFunctionSimpleFactory::is_foreach(_fn.name.function_name)) {
224
0
            _function = AggregateFunctionSimpleFactory::instance().get(
225
0
                    _fn.name.function_name, argument_types, _data_type,
226
0
                    AggregateFunctionSimpleFactory::result_nullable_by_foreach(_data_type),
227
0
                    state->be_exec_version(),
228
0
                    {.is_window_function = _is_window_function,
229
0
                     .is_foreach = is_foreach,
230
0
                     .enable_aggregate_function_null_v2 =
231
0
                             state->enable_aggregate_function_null_v2(),
232
0
                     .new_version_percentile =
233
0
                             state->query_options().__isset.new_version_percentile &&
234
0
                             state->query_options().new_version_percentile,
235
0
                     .column_names = std::move(column_names)});
236
185k
        } else {
237
185k
            _function = AggregateFunctionSimpleFactory::instance().get(
238
185k
                    _fn.name.function_name, argument_types, _data_type, _data_type->is_nullable(),
239
185k
                    state->be_exec_version(),
240
185k
                    {.is_window_function = _is_window_function,
241
185k
                     .is_foreach = is_foreach,
242
185k
                     .enable_aggregate_function_null_v2 =
243
185k
                             state->enable_aggregate_function_null_v2(),
244
185k
                     .new_version_percentile =
245
185k
                             state->query_options().__isset.new_version_percentile &&
246
185k
                             state->query_options().new_version_percentile,
247
185k
                     .column_names = std::move(column_names)});
248
185k
        }
249
185k
    }
250
186k
    if (_function == nullptr) {
251
0
        return Status::InternalError("Agg Function {} is not implemented", _fn.signature);
252
0
    }
253
254
186k
    if (!_sort_description.empty()) {
255
79
        _function = transform_to_sort_agg_function(_function, _argument_types_with_sort,
256
79
                                                   _sort_description, state);
257
79
    }
258
259
186k
    if (_fn.name.function_name == "ai_agg") {
260
0
        _function->set_query_context(state->get_query_ctx());
261
0
    }
262
263
    // Foreachv2, like foreachv1, does not check the return type,
264
    // because its return type is related to the internal agg.
265
186k
    if (!AggregateFunctionSimpleFactory::is_foreach(_fn.name.function_name) &&
266
186k
        !AggregateFunctionSimpleFactory::is_foreachv2(_fn.name.function_name)) {
267
186k
        if (state->be_exec_version() >= BE_VERSION_THAT_SUPPORT_NULLABLE_CHECK) {
268
186k
            RETURN_IF_ERROR(
269
186k
                    _function->verify_result_type(_without_key, argument_types, _data_type));
270
186k
        }
271
186k
    }
272
186k
    _expr_name = fmt::format("{}({})", _fn.name.function_name, child_expr_name);
273
186k
    return Status::OK();
274
186k
}
275
276
186k
Status AggFnEvaluator::open(RuntimeState* state) {
277
186k
    return VExpr::open(_input_exprs_ctxs, state);
278
186k
}
279
280
3.51M
void AggFnEvaluator::create(AggregateDataPtr place) {
281
3.51M
    _function->create(place);
282
3.51M
}
283
284
8.22k
void AggFnEvaluator::destroy(AggregateDataPtr place) {
285
8.22k
    _function->destroy(place);
286
8.22k
}
287
288
244k
Status AggFnEvaluator::execute_single_add(Block* block, AggregateDataPtr place, Arena& arena) {
289
244k
    RETURN_IF_ERROR(_calc_argument_columns(block));
290
244k
    _function->add_batch_single_place(block->rows(), place, _agg_columns.data(), arena);
291
244k
    return Status::OK();
292
244k
}
293
294
Status AggFnEvaluator::execute_batch_add(Block* block, size_t offset, AggregateDataPtr* places,
295
63.9k
                                         Arena& arena, bool agg_many) {
296
63.9k
    RETURN_IF_ERROR(_calc_argument_columns(block));
297
63.9k
    _function->add_batch(block->rows(), places, offset, _agg_columns.data(), arena, agg_many);
298
63.9k
    return Status::OK();
299
63.9k
}
300
301
Status AggFnEvaluator::execute_batch_add_selected(Block* block, size_t offset,
302
2
                                                  AggregateDataPtr* places, Arena& arena) {
303
2
    RETURN_IF_ERROR(_calc_argument_columns(block));
304
2
    _function->add_batch_selected(block->rows(), places, offset, _agg_columns.data(), arena);
305
2
    return Status::OK();
306
2
}
307
308
Status AggFnEvaluator::streaming_agg_serialize_to_column(Block* block, MutableColumnPtr& dst,
309
26
                                                         const size_t num_rows, Arena& arena) {
310
26
    RETURN_IF_ERROR(_calc_argument_columns(block));
311
26
    _function->streaming_agg_serialize_to_column(_agg_columns.data(), dst, num_rows, arena);
312
26
    return Status::OK();
313
26
}
314
315
86.9k
void AggFnEvaluator::insert_result_info(AggregateDataPtr place, IColumn* column) {
316
86.9k
    _function->insert_result_into(place, *column);
317
86.9k
}
318
319
void AggFnEvaluator::insert_result_info_vec(const std::vector<AggregateDataPtr>& places,
320
57.3k
                                            size_t offset, IColumn* column, const size_t num_rows) {
321
57.3k
    _function->insert_result_into_vec(places, offset, *column, num_rows);
322
57.3k
}
323
324
15.2k
void AggFnEvaluator::reset(AggregateDataPtr place) {
325
15.2k
    _function->reset(place);
326
15.2k
}
327
328
0
std::string AggFnEvaluator::debug_string(const std::vector<AggFnEvaluator*>& exprs) {
329
0
    std::stringstream out;
330
0
    out << "[";
331
332
0
    for (int i = 0; i < exprs.size(); ++i) {
333
0
        out << (i == 0 ? "" : " ") << exprs[i]->debug_string();
334
0
    }
335
336
0
    out << "]";
337
0
    return out.str();
338
0
}
339
340
0
std::string AggFnEvaluator::debug_string() const {
341
0
    std::stringstream out;
342
0
    out << "AggFnEvaluator(";
343
0
    out << _fn.signature;
344
0
    out << ")";
345
0
    return out.str();
346
0
}
347
348
308k
Status AggFnEvaluator::_calc_argument_columns(Block* block) {
349
308k
    SCOPED_TIMER(_expr_timer);
350
308k
    _agg_columns.resize(_input_exprs_ctxs.size());
351
308k
    std::vector<int> column_ids(_input_exprs_ctxs.size());
352
548k
    for (int i = 0; i < _input_exprs_ctxs.size(); ++i) {
353
239k
        int column_id = -1;
354
239k
        RETURN_IF_ERROR(_input_exprs_ctxs[i]->execute(block, &column_id));
355
239k
        column_ids[i] = column_id;
356
239k
    }
357
308k
    materialize_block_inplace(*block, column_ids.data(),
358
308k
                              column_ids.data() + _input_exprs_ctxs.size());
359
548k
    for (int i = 0; i < _input_exprs_ctxs.size(); ++i) {
360
240k
        _agg_columns[i] = block->get_by_position(column_ids[i]).column.get();
361
240k
    }
362
308k
    return Status::OK();
363
308k
}
364
365
325k
AggFnEvaluator* AggFnEvaluator::clone(RuntimeState* state, ObjectPool* pool) {
366
325k
    return pool->add(AggFnEvaluator::create_unique(*this, state).release());
367
325k
}
368
369
AggFnEvaluator::AggFnEvaluator(AggFnEvaluator& evaluator, RuntimeState* state)
370
325k
        : _fn(evaluator._fn),
371
325k
          _is_merge(evaluator._is_merge),
372
325k
          _without_key(evaluator._without_key),
373
325k
          _is_window_function(evaluator._is_window_function),
374
325k
          _argument_types_with_sort(evaluator._argument_types_with_sort),
375
325k
          _real_argument_types(evaluator._real_argument_types),
376
325k
          _intermediate_slot_desc(evaluator._intermediate_slot_desc),
377
325k
          _output_slot_desc(evaluator._output_slot_desc),
378
325k
          _sort_description(evaluator._sort_description),
379
325k
          _data_type(evaluator._data_type),
380
325k
          _function(evaluator._function),
381
325k
          _expr_name(evaluator._expr_name),
382
325k
          _agg_columns(evaluator._agg_columns) {
383
325k
    if (evaluator._fn.binary_type == TFunctionBinaryType::JAVA_UDF) {
384
671
        DataTypes tmp_argument_types;
385
671
        tmp_argument_types.reserve(evaluator._input_exprs_ctxs.size());
386
        // prepare for argument
387
733
        for (auto& _input_exprs_ctx : evaluator._input_exprs_ctxs) {
388
733
            auto data_type = _input_exprs_ctx->root()->data_type();
389
733
            tmp_argument_types.emplace_back(data_type);
390
733
        }
391
671
        const DataTypes& argument_types =
392
671
                _real_argument_types.empty() ? tmp_argument_types : _real_argument_types;
393
671
        _function = AggregateJavaUdaf::create(evaluator._fn, argument_types, evaluator._data_type);
394
671
        THROW_IF_ERROR(static_cast<AggregateJavaUdaf*>(_function.get())->check_udaf(evaluator._fn));
395
671
    }
396
325k
    DCHECK(_function != nullptr);
397
398
325k
    _input_exprs_ctxs.resize(evaluator._input_exprs_ctxs.size());
399
610k
    for (size_t i = 0; i < _input_exprs_ctxs.size(); i++) {
400
285k
        WARN_IF_ERROR(evaluator._input_exprs_ctxs[i]->clone(state, _input_exprs_ctxs[i]), "");
401
285k
    }
402
325k
}
403
404
Status AggFnEvaluator::check_agg_fn_output(uint32_t key_size,
405
                                           const std::vector<AggFnEvaluator*>& agg_fn,
406
35.2k
                                           const RowDescriptor& output_row_desc) {
407
35.2k
    auto name_and_types = VectorizedUtils::create_name_and_data_types(output_row_desc);
408
131k
    for (uint32_t i = key_size, j = 0; i < name_and_types.size(); i++, j++) {
409
96.1k
        auto&& [name, column_type] = name_and_types[i];
410
96.1k
        auto agg_return_type = agg_fn[j]->function()->get_return_type();
411
96.1k
        if (!column_type->equals(*agg_return_type)) {
412
14.9k
            if (!column_type->is_nullable() || agg_return_type->is_nullable() ||
413
14.9k
                !remove_nullable(column_type)->equals(*agg_return_type)) {
414
0
                return Status::InternalError(
415
0
                        "column_type not match data_types in agg node, column_type={}, "
416
0
                        "data_types={},column name={}",
417
0
                        column_type->get_name(), agg_return_type->get_name(), name);
418
0
            }
419
14.9k
        }
420
96.1k
    }
421
35.2k
    return Status::OK();
422
35.2k
}
423
424
1.48M
bool AggFnEvaluator::is_blockable() const {
425
1.48M
    return _function->is_blockable() ||
426
1.48M
           std::any_of(_input_exprs_ctxs.begin(), _input_exprs_ctxs.end(),
427
1.48M
                       [](VExprContextSPtr ctx) { return ctx->root()->is_blockable(); });
428
1.48M
}
429
430
} // namespace doris