Coverage Report

Created: 2026-09-29 14:44

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
be/src/exprs/aggregate/aggregate_function_count.h
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
// This file is copied from
18
// https://github.com/ClickHouse/ClickHouse/blob/master/src/AggregateFunctions/AggregateFunctionCount.h
19
// and modified by Doris
20
21
#pragma once
22
23
#include <stddef.h>
24
25
#include <algorithm>
26
#include <boost/iterator/iterator_facade.hpp>
27
#include <memory>
28
#include <vector>
29
30
#include "core/assert_cast.h"
31
#include "core/column/column.h"
32
#include "core/column/column_fixed_length_object.h"
33
#include "core/column/column_nullable.h"
34
#include "core/column/column_vector.h"
35
#include "core/data_type/data_type.h"
36
#include "core/data_type/data_type_fixed_length_object.h"
37
#include "core/data_type/data_type_number.h"
38
#include "core/types.h"
39
#include "exprs/aggregate/aggregate_function.h"
40
#include "util/simd/bits.h"
41
42
namespace doris {
43
class Arena;
44
class BufferReadable;
45
class BufferWritable;
46
47
struct AggregateFunctionCountData {
48
    UInt64 count = 0;
49
};
50
51
/// Simply count number of calls.
52
class AggregateFunctionCount final
53
        : public IAggregateFunctionDataHelper<AggregateFunctionCountData, AggregateFunctionCount>,
54
          VarargsExpression,
55
          NotNullableAggregateFunction {
56
public:
57
    AggregateFunctionCount(const DataTypes& argument_types_)
58
28.7k
            : IAggregateFunctionDataHelper(argument_types_) {}
59
60
3.91k
    bool is_simple_count() const override { return true; }
61
238
    String get_name() const override { return "count"; }
62
63
106k
    DataTypePtr get_return_type() const override { return std::make_shared<DataTypeInt64>(); }
64
65
0
    bool is_trivial() const override { return true; }
66
67
9.64M
    void add(AggregateDataPtr __restrict place, const IColumn**, ssize_t, Arena&) const override {
68
9.64M
        ++data(place).count;
69
9.64M
    }
70
71
    void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn**,
72
54.6k
                                Arena&) const override {
73
54.6k
        data(place).count += batch_size;
74
54.6k
    }
75
76
156
    void reset(AggregateDataPtr place) const override {
77
156
        AggregateFunctionCount::data(place).count = 0;
78
156
    }
79
80
    void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs,
81
107k
               Arena&) const override {
82
107k
        data(place).count += data(rhs).count;
83
107k
    }
84
85
81
    void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override {
86
81
        buf.write_var_uint(data(place).count);
87
81
    }
88
89
    void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf,
90
81
                     Arena&) const override {
91
81
        buf.read_var_uint(data(place).count);
92
81
    }
93
94
116k
    void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override {
95
116k
        assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(to).get_data().push_back(
96
116k
                data(place).count);
97
116k
    }
98
99
    void serialize_to_column(const std::vector<AggregateDataPtr>& places, size_t offset,
100
3.61k
                             MutableColumnPtr& dst, const size_t num_rows) const override {
101
3.61k
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
102
3.61k
        DCHECK(col.item_size() == sizeof(Data))
103
2
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
104
3.61k
        col.resize(num_rows);
105
3.61k
        auto* data = col.get_data().data();
106
214k
        for (size_t i = 0; i != num_rows; ++i) {
107
210k
            *reinterpret_cast<Data*>(&data[sizeof(Data) * i]) =
108
210k
                    *reinterpret_cast<Data*>(places[i] + offset);
109
210k
        }
110
3.61k
    }
111
112
    void streaming_agg_serialize_to_column(const IColumn** columns, MutableColumnPtr& dst,
113
64
                                           const size_t num_rows, Arena&) const override {
114
64
        auto& dst_col = assert_cast<ColumnFixedLengthObject&>(*dst);
115
64
        DCHECK(dst_col.item_size() == sizeof(Data))
116
0
                << "size is not equal: " << dst_col.item_size() << " " << sizeof(Data);
117
64
        dst_col.resize(num_rows);
118
64
        auto* data = dst_col.get_data().data();
119
339
        for (size_t i = 0; i != num_rows; ++i) {
120
275
            auto& state = *reinterpret_cast<Data*>(&data[sizeof(Data) * i]);
121
275
            state.count = 1;
122
275
        }
123
64
    }
124
125
    void deserialize_and_merge_from_column_range(AggregateDataPtr __restrict place,
126
                                                 const IColumn& column, size_t begin, size_t end,
127
5.53k
                                                 Arena&) const override {
128
5.53k
        DCHECK(end <= column.size() && begin <= end)
129
0
                << ", begin:" << begin << ", end:" << end << ", column.size():" << column.size();
130
5.53k
        auto& col = assert_cast<const ColumnFixedLengthObject&>(column);
131
5.53k
        auto* data = reinterpret_cast<const Data*>(col.get_data().data());
132
113k
        for (size_t i = begin; i <= end; ++i) {
133
108k
            doris::AggregateFunctionCount::data(place).count += data[i].count;
134
108k
        }
135
5.53k
    }
136
137
    void deserialize_and_merge_vec(const AggregateDataPtr* places, size_t offset,
138
                                   AggregateDataPtr rhs, const IColumn* column, Arena& arena,
139
2.99k
                                   const size_t num_rows) const override {
140
2.99k
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
141
2.99k
        const auto* data = col.get_data().data();
142
2.99k
        this->merge_vec(places, offset, AggregateDataPtr(data), arena, num_rows);
143
2.99k
    }
144
145
    void deserialize_and_merge_vec_selected(const AggregateDataPtr* places, size_t offset,
146
                                            AggregateDataPtr rhs, const IColumn* column,
147
1
                                            Arena& arena, const size_t num_rows) const override {
148
1
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
149
1
        const auto* data = col.get_data().data();
150
1
        this->merge_vec_selected(places, offset, AggregateDataPtr(data), arena, num_rows);
151
1
    }
152
153
    void serialize_without_key_to_column(ConstAggregateDataPtr __restrict place,
154
5.08k
                                         IColumn& to) const override {
155
5.08k
        auto& col = assert_cast<ColumnFixedLengthObject&>(to);
156
5.08k
        DCHECK(col.item_size() == sizeof(Data))
157
18
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
158
5.08k
        size_t old_size = col.size();
159
5.08k
        col.resize(old_size + 1);
160
5.08k
        (reinterpret_cast<Data*>(col.get_data().data()) + old_size)->count =
161
5.08k
                AggregateFunctionCount::data(place).count;
162
5.08k
    }
163
164
9.65k
    MutableColumnPtr create_serialize_column() const override {
165
9.65k
        return ColumnFixedLengthObject::create(sizeof(Data));
166
9.65k
    }
167
168
9.79k
    DataTypePtr get_serialized_type() const override {
169
9.79k
        return std::make_shared<DataTypeFixedLengthObject>();
170
9.79k
    }
171
172
    void add_range_single_place(int64_t partition_start, int64_t partition_end, int64_t frame_start,
173
                                int64_t frame_end, AggregateDataPtr place, const IColumn** columns,
174
                                Arena& arena, UInt8* use_null_result,
175
264
                                UInt8* could_use_previous_result) const override {
176
264
        frame_start = std::max<int64_t>(frame_start, partition_start);
177
264
        frame_end = std::min<int64_t>(frame_end, partition_end);
178
264
        if (frame_start >= frame_end) {
179
19
            if (!*could_use_previous_result) {
180
2
                *use_null_result = true;
181
2
            }
182
245
        } else {
183
245
            AggregateFunctionCount::data(place).count += frame_end - frame_start;
184
245
            *use_null_result = false;
185
245
            *could_use_previous_result = true;
186
245
        }
187
264
    }
188
};
189
190
// Used for unary count(nullable_expr). SQL count(expr) counts non-NULL values.
191
class AggregateFunctionCountNotNullUnary final
192
        : public IAggregateFunctionDataHelper<AggregateFunctionCountData,
193
                                              AggregateFunctionCountNotNullUnary> {
194
public:
195
    AggregateFunctionCountNotNullUnary(const DataTypes& argument_types_)
196
12.1k
            : IAggregateFunctionDataHelper(argument_types_) {}
197
198
260
    String get_name() const override { return "count"; }
199
200
44.4k
    DataTypePtr get_return_type() const override { return std::make_shared<DataTypeInt64>(); }
201
202
0
    bool is_trivial() const override { return true; }
203
204
    void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num,
205
1.23M
             Arena&) const override {
206
1.23M
        data(place).count +=
207
1.23M
                !assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0])
208
1.23M
                         .is_null_at(row_num);
209
1.23M
    }
210
211
    void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn** columns,
212
13.8k
                                Arena&) const override {
213
13.8k
        const auto& nullable_column =
214
13.8k
                assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0]);
215
13.8k
        const auto& null_map = nullable_column.get_null_map_data();
216
13.8k
        DCHECK_LE(batch_size, null_map.size());
217
13.8k
        if (!nullable_column.has_null(0, batch_size)) {
218
11.2k
            data(place).count += batch_size;
219
11.2k
            return;
220
11.2k
        }
221
2.55k
        data(place).count +=
222
2.55k
                simd::count_zero_num(reinterpret_cast<const int8_t*>(null_map.data()), batch_size);
223
2.55k
    }
224
225
253
    void reset(AggregateDataPtr place) const override { data(place).count = 0; }
226
227
    void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs,
228
76.1k
               Arena&) const override {
229
76.1k
        data(place).count += data(rhs).count;
230
76.1k
    }
231
232
11
    void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override {
233
11
        buf.write_var_uint(data(place).count);
234
11
    }
235
236
    void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf,
237
11
                     Arena&) const override {
238
11
        buf.read_var_uint(data(place).count);
239
11
    }
240
241
86.8k
    void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override {
242
86.8k
        if (is_column_nullable(to)) {
243
1
            auto& null_column = assert_cast<ColumnNullable&, TypeCheckOnRelease::DISABLE>(to);
244
1
            null_column.get_null_map_data().push_back(0);
245
1
            assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(null_column.get_nested_column())
246
1
                    .get_data()
247
1
                    .push_back(data(place).count);
248
86.8k
        } else {
249
86.8k
            assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(to).get_data().push_back(
250
86.8k
                    data(place).count);
251
86.8k
        }
252
86.8k
    }
253
254
13.4k
    void check_result_column_type(const IColumn& to) const override {
255
13.4k
        if (const auto* null_column = check_and_get_column<ColumnNullable>(to)) {
256
1
            IAggregateFunction::check_result_column_type(null_column->get_nested_column());
257
1
            return;
258
1
        }
259
13.4k
        IAggregateFunction::check_result_column_type(to);
260
13.4k
    }
261
262
    void serialize_to_column(const std::vector<AggregateDataPtr>& places, size_t offset,
263
2.73k
                             MutableColumnPtr& dst, const size_t num_rows) const override {
264
2.73k
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
265
2.73k
        DCHECK(col.item_size() == sizeof(Data))
266
3
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
267
2.73k
        col.resize(num_rows);
268
2.73k
        auto* data = col.get_data().data();
269
163k
        for (size_t i = 0; i != num_rows; ++i) {
270
160k
            *reinterpret_cast<Data*>(&data[sizeof(Data) * i]) =
271
160k
                    *reinterpret_cast<Data*>(places[i] + offset);
272
160k
        }
273
2.73k
    }
274
275
    void streaming_agg_serialize_to_column(const IColumn** columns, MutableColumnPtr& dst,
276
4
                                           const size_t num_rows, Arena&) const override {
277
4
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
278
4
        DCHECK(col.item_size() == sizeof(Data))
279
0
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
280
4
        col.resize(num_rows);
281
4
        auto& data = col.get_data();
282
4
        const ColumnNullable& input_col = assert_cast<const ColumnNullable&>(*columns[0]);
283
25
        for (size_t i = 0; i < num_rows; i++) {
284
21
            auto& state = *reinterpret_cast<Data*>(&data[sizeof(Data) * i]);
285
21
            state.count = !input_col.is_null_at(i);
286
21
        }
287
4
    }
288
289
    void deserialize_and_merge_from_column_range(AggregateDataPtr __restrict place,
290
                                                 const IColumn& column, size_t begin, size_t end,
291
3.89k
                                                 Arena&) const override {
292
3.89k
        DCHECK(end <= column.size() && begin <= end)
293
0
                << ", begin:" << begin << ", end:" << end << ", column.size():" << column.size();
294
3.89k
        auto& col = assert_cast<const ColumnFixedLengthObject&>(column);
295
3.89k
        auto* data = reinterpret_cast<const Data*>(col.get_data().data());
296
81.4k
        for (size_t i = begin; i <= end; ++i) {
297
77.5k
            doris::AggregateFunctionCountNotNullUnary::data(place).count += data[i].count;
298
77.5k
        }
299
3.89k
    }
300
301
    void deserialize_and_merge_vec(const AggregateDataPtr* places, size_t offset,
302
                                   AggregateDataPtr rhs, const IColumn* column, Arena& arena,
303
1.85k
                                   const size_t num_rows) const override {
304
1.85k
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
305
1.85k
        const auto* data = col.get_data().data();
306
1.85k
        this->merge_vec(places, offset, AggregateDataPtr(data), arena, num_rows);
307
1.85k
    }
308
309
    void deserialize_and_merge_vec_selected(const AggregateDataPtr* places, size_t offset,
310
                                            AggregateDataPtr rhs, const IColumn* column,
311
2
                                            Arena& arena, const size_t num_rows) const override {
312
2
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
313
2
        const auto* data = col.get_data().data();
314
2
        this->merge_vec_selected(places, offset, AggregateDataPtr(data), arena, num_rows);
315
2
    }
316
317
    void serialize_without_key_to_column(ConstAggregateDataPtr __restrict place,
318
3.75k
                                         IColumn& to) const override {
319
3.75k
        auto& col = assert_cast<ColumnFixedLengthObject&>(to);
320
3.75k
        DCHECK(col.item_size() == sizeof(Data))
321
0
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
322
3.75k
        size_t old_size = col.size();
323
3.75k
        col.resize(old_size + 1);
324
3.75k
        (reinterpret_cast<Data*>(col.get_data().data()) + old_size)->count =
325
3.75k
                AggregateFunctionCountNotNullUnary::data(place).count;
326
3.75k
    }
327
328
6.46k
    MutableColumnPtr create_serialize_column() const override {
329
6.46k
        return ColumnFixedLengthObject::create(sizeof(Data));
330
6.46k
    }
331
332
6.47k
    DataTypePtr get_serialized_type() const override {
333
6.47k
        return std::make_shared<DataTypeFixedLengthObject>();
334
6.47k
    }
335
336
    void add_range_single_place(int64_t partition_start, int64_t partition_end, int64_t frame_start,
337
                                int64_t frame_end, AggregateDataPtr place, const IColumn** columns,
338
                                Arena& arena, UInt8* use_null_result,
339
337
                                UInt8* could_use_previous_result) const override {
340
337
        frame_start = std::max<int64_t>(frame_start, partition_start);
341
337
        frame_end = std::min<int64_t>(frame_end, partition_end);
342
337
        if (frame_start >= frame_end) {
343
12
            if (!*could_use_previous_result) {
344
0
                *use_null_result = true;
345
0
            }
346
325
        } else {
347
325
            const auto& nullable_column =
348
325
                    assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0]);
349
325
            size_t count = 0;
350
325
            if (nullable_column.has_null()) {
351
175
                for (int64_t i = frame_start; i < frame_end; ++i) {
352
121
                    if (!nullable_column.is_null_at(i)) {
353
54
                        ++count;
354
54
                    }
355
121
                }
356
271
            } else {
357
271
                count = frame_end - frame_start;
358
271
            }
359
325
            *use_null_result = false;
360
325
            *could_use_previous_result = true;
361
325
            AggregateFunctionCountNotNullUnary::data(place).count += count;
362
325
        }
363
337
    }
364
};
365
366
} // namespace doris