Coverage Report

Created: 2026-10-09 09:59

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
29.4k
            : IAggregateFunctionDataHelper(argument_types_) {}
59
60
5.15k
    bool is_simple_count() const override { return true; }
61
392
    String get_name() const override { return "count"; }
62
63
108k
    DataTypePtr get_return_type() const override { return std::make_shared<DataTypeInt64>(); }
64
65
0
    bool is_trivial() const override { return true; }
66
67
57
    WindowSpillStrategy window_spill_strategy() const override {
68
57
        return WindowSpillStrategy::PARTITION_REDUCE;
69
57
    }
70
71
9.84M
    void add(AggregateDataPtr __restrict place, const IColumn**, ssize_t, Arena&) const override {
72
9.84M
        ++data(place).count;
73
9.84M
    }
74
75
    void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn**,
76
54.2k
                                Arena&) const override {
77
54.2k
        data(place).count += batch_size;
78
54.2k
    }
79
80
309
    void reset(AggregateDataPtr place) const override {
81
309
        AggregateFunctionCount::data(place).count = 0;
82
309
    }
83
84
    void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs,
85
108k
               Arena&) const override {
86
108k
        data(place).count += data(rhs).count;
87
108k
    }
88
89
73
    void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override {
90
73
        buf.write_var_uint(data(place).count);
91
73
    }
92
93
    void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf,
94
73
                     Arena&) const override {
95
73
        buf.read_var_uint(data(place).count);
96
73
    }
97
98
155k
    void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override {
99
155k
        assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(to).get_data().push_back(
100
155k
                data(place).count);
101
155k
    }
102
103
    void serialize_to_column(const std::vector<AggregateDataPtr>& places, size_t offset,
104
5.00k
                             MutableColumnPtr& dst, const size_t num_rows) const override {
105
5.00k
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
106
5.00k
        DCHECK(col.item_size() == sizeof(Data))
107
10
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
108
5.00k
        col.resize(num_rows);
109
5.00k
        auto* data = col.get_data().data();
110
217k
        for (size_t i = 0; i != num_rows; ++i) {
111
212k
            *reinterpret_cast<Data*>(&data[sizeof(Data) * i]) =
112
212k
                    *reinterpret_cast<Data*>(places[i] + offset);
113
212k
        }
114
5.00k
    }
115
116
    void streaming_agg_serialize_to_column(const IColumn** columns, MutableColumnPtr& dst,
117
67
                                           const size_t num_rows, Arena&) const override {
118
67
        auto& dst_col = assert_cast<ColumnFixedLengthObject&>(*dst);
119
67
        DCHECK(dst_col.item_size() == sizeof(Data))
120
0
                << "size is not equal: " << dst_col.item_size() << " " << sizeof(Data);
121
67
        dst_col.resize(num_rows);
122
67
        auto* data = dst_col.get_data().data();
123
342
        for (size_t i = 0; i != num_rows; ++i) {
124
275
            auto& state = *reinterpret_cast<Data*>(&data[sizeof(Data) * i]);
125
275
            state.count = 1;
126
275
        }
127
67
    }
128
129
    void deserialize_and_merge_from_column_range(AggregateDataPtr __restrict place,
130
                                                 const IColumn& column, size_t begin, size_t end,
131
6.93k
                                                 Arena&) const override {
132
6.93k
        DCHECK(end <= column.size() && begin <= end)
133
0
                << ", begin:" << begin << ", end:" << end << ", column.size():" << column.size();
134
6.93k
        auto& col = assert_cast<const ColumnFixedLengthObject&>(column);
135
6.93k
        auto* data = reinterpret_cast<const Data*>(col.get_data().data());
136
116k
        for (size_t i = begin; i <= end; ++i) {
137
109k
            doris::AggregateFunctionCount::data(place).count += data[i].count;
138
109k
        }
139
6.93k
    }
140
141
    void deserialize_and_merge_vec(const AggregateDataPtr* places, size_t offset,
142
                                   AggregateDataPtr rhs, const IColumn* column, Arena& arena,
143
3.79k
                                   const size_t num_rows) const override {
144
3.79k
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
145
3.79k
        const auto* data = col.get_data().data();
146
3.79k
        this->merge_vec(places, offset, AggregateDataPtr(data), arena, num_rows);
147
3.79k
    }
148
149
    void deserialize_and_merge_vec_selected(const AggregateDataPtr* places, size_t offset,
150
                                            AggregateDataPtr rhs, const IColumn* column,
151
1
                                            Arena& arena, const size_t num_rows) const override {
152
1
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
153
1
        const auto* data = col.get_data().data();
154
1
        this->merge_vec_selected(places, offset, AggregateDataPtr(data), arena, num_rows);
155
1
    }
156
157
    void serialize_without_key_to_column(ConstAggregateDataPtr __restrict place,
158
6.32k
                                         IColumn& to) const override {
159
6.32k
        auto& col = assert_cast<ColumnFixedLengthObject&>(to);
160
6.32k
        DCHECK(col.item_size() == sizeof(Data))
161
1
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
162
6.32k
        size_t old_size = col.size();
163
6.32k
        col.resize(old_size + 1);
164
6.32k
        (reinterpret_cast<Data*>(col.get_data().data()) + old_size)->count =
165
6.32k
                AggregateFunctionCount::data(place).count;
166
6.32k
    }
167
168
12.5k
    MutableColumnPtr create_serialize_column() const override {
169
12.5k
        return ColumnFixedLengthObject::create(sizeof(Data));
170
12.5k
    }
171
172
12.6k
    DataTypePtr get_serialized_type() const override {
173
12.6k
        return std::make_shared<DataTypeFixedLengthObject>();
174
12.6k
    }
175
176
    void add_range_single_place(int64_t partition_start, int64_t partition_end, int64_t frame_start,
177
                                int64_t frame_end, AggregateDataPtr place, const IColumn** columns,
178
                                Arena& arena, UInt8* use_null_result,
179
1.23k
                                UInt8* could_use_previous_result) const override {
180
1.23k
        frame_start = std::max<int64_t>(frame_start, partition_start);
181
1.23k
        frame_end = std::min<int64_t>(frame_end, partition_end);
182
1.23k
        if (frame_start >= frame_end) {
183
19
            if (!*could_use_previous_result) {
184
2
                *use_null_result = true;
185
2
            }
186
1.21k
        } else {
187
1.21k
            AggregateFunctionCount::data(place).count += frame_end - frame_start;
188
1.21k
            *use_null_result = false;
189
1.21k
            *could_use_previous_result = true;
190
1.21k
        }
191
1.23k
    }
192
};
193
194
// Used for unary count(nullable_expr). SQL count(expr) counts non-NULL values.
195
class AggregateFunctionCountNotNullUnary final
196
        : public IAggregateFunctionDataHelper<AggregateFunctionCountData,
197
                                              AggregateFunctionCountNotNullUnary> {
198
public:
199
    AggregateFunctionCountNotNullUnary(const DataTypes& argument_types_)
200
13.9k
            : IAggregateFunctionDataHelper(argument_types_) {}
201
202
425
    String get_name() const override { return "count"; }
203
204
106k
    DataTypePtr get_return_type() const override { return std::make_shared<DataTypeInt64>(); }
205
206
0
    bool is_trivial() const override { return true; }
207
208
56
    WindowSpillStrategy window_spill_strategy() const override {
209
56
        return WindowSpillStrategy::PARTITION_REDUCE;
210
56
    }
211
212
    void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num,
213
1.41M
             Arena&) const override {
214
1.41M
        data(place).count +=
215
1.41M
                !assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0])
216
1.41M
                         .is_null_at(row_num);
217
1.41M
    }
218
219
    void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn** columns,
220
16.6k
                                Arena&) const override {
221
16.6k
        const auto& nullable_column =
222
16.6k
                assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0]);
223
16.6k
        const auto& null_map = nullable_column.get_null_map_data();
224
16.6k
        DCHECK_LE(batch_size, null_map.size());
225
16.6k
        if (!nullable_column.has_null(0, batch_size)) {
226
13.2k
            data(place).count += batch_size;
227
13.2k
            return;
228
13.2k
        }
229
3.31k
        data(place).count +=
230
3.31k
                simd::count_zero_num(reinterpret_cast<const int8_t*>(null_map.data()), batch_size);
231
3.31k
    }
232
233
59.9k
    void reset(AggregateDataPtr place) const override { data(place).count = 0; }
234
235
    void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs,
236
103k
               Arena&) const override {
237
103k
        data(place).count += data(rhs).count;
238
103k
    }
239
240
321
    void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override {
241
321
        buf.write_var_uint(data(place).count);
242
321
    }
243
244
    void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf,
245
323
                     Arena&) const override {
246
323
        buf.read_var_uint(data(place).count);
247
323
    }
248
249
196k
    void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override {
250
196k
        if (is_column_nullable(to)) {
251
1
            auto& null_column = assert_cast<ColumnNullable&, TypeCheckOnRelease::DISABLE>(to);
252
1
            null_column.get_null_map_data().push_back(0);
253
1
            assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(null_column.get_nested_column())
254
1
                    .get_data()
255
1
                    .push_back(data(place).count);
256
196k
        } else {
257
196k
            assert_cast<ColumnInt64&, TypeCheckOnRelease::DISABLE>(to).get_data().push_back(
258
196k
                    data(place).count);
259
196k
        }
260
196k
    }
261
262
73.4k
    void check_result_column_type(const IColumn& to) const override {
263
73.4k
        if (const auto* null_column = check_and_get_column<ColumnNullable>(to)) {
264
1
            IAggregateFunction::check_result_column_type(null_column->get_nested_column());
265
1
            return;
266
1
        }
267
73.4k
        IAggregateFunction::check_result_column_type(to);
268
73.4k
    }
269
270
    void serialize_to_column(const std::vector<AggregateDataPtr>& places, size_t offset,
271
4.06k
                             MutableColumnPtr& dst, const size_t num_rows) const override {
272
4.06k
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
273
18.4E
        DCHECK(col.item_size() == sizeof(Data))
274
18.4E
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
275
4.06k
        col.resize(num_rows);
276
4.06k
        auto* data = col.get_data().data();
277
184k
        for (size_t i = 0; i != num_rows; ++i) {
278
180k
            *reinterpret_cast<Data*>(&data[sizeof(Data) * i]) =
279
180k
                    *reinterpret_cast<Data*>(places[i] + offset);
280
180k
        }
281
4.06k
    }
282
283
    void streaming_agg_serialize_to_column(const IColumn** columns, MutableColumnPtr& dst,
284
70
                                           const size_t num_rows, Arena&) const override {
285
70
        auto& col = assert_cast<ColumnFixedLengthObject&>(*dst);
286
70
        DCHECK(col.item_size() == sizeof(Data))
287
0
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
288
70
        col.resize(num_rows);
289
70
        auto& data = col.get_data();
290
70
        const ColumnNullable& input_col = assert_cast<const ColumnNullable&>(*columns[0]);
291
211
        for (size_t i = 0; i < num_rows; i++) {
292
141
            auto& state = *reinterpret_cast<Data*>(&data[sizeof(Data) * i]);
293
141
            state.count = !input_col.is_null_at(i);
294
141
        }
295
70
    }
296
297
    void deserialize_and_merge_from_column_range(AggregateDataPtr __restrict place,
298
                                                 const IColumn& column, size_t begin, size_t end,
299
7.02k
                                                 Arena&) const override {
300
7.02k
        DCHECK(end <= column.size() && begin <= end)
301
1
                << ", begin:" << begin << ", end:" << end << ", column.size():" << column.size();
302
7.02k
        auto& col = assert_cast<const ColumnFixedLengthObject&>(column);
303
7.02k
        auto* data = reinterpret_cast<const Data*>(col.get_data().data());
304
88.5k
        for (size_t i = begin; i <= end; ++i) {
305
81.5k
            doris::AggregateFunctionCountNotNullUnary::data(place).count += data[i].count;
306
81.5k
        }
307
7.02k
    }
308
309
    void deserialize_and_merge_vec(const AggregateDataPtr* places, size_t offset,
310
                                   AggregateDataPtr rhs, const IColumn* column, Arena& arena,
311
2.34k
                                   const size_t num_rows) const override {
312
2.34k
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
313
2.34k
        const auto* data = col.get_data().data();
314
2.34k
        this->merge_vec(places, offset, AggregateDataPtr(data), arena, num_rows);
315
2.34k
    }
316
317
    void deserialize_and_merge_vec_selected(const AggregateDataPtr* places, size_t offset,
318
                                            AggregateDataPtr rhs, const IColumn* column,
319
2
                                            Arena& arena, const size_t num_rows) const override {
320
2
        const auto& col = assert_cast<const ColumnFixedLengthObject&>(*column);
321
2
        const auto* data = col.get_data().data();
322
2
        this->merge_vec_selected(places, offset, AggregateDataPtr(data), arena, num_rows);
323
2
    }
324
325
    void serialize_without_key_to_column(ConstAggregateDataPtr __restrict place,
326
6.05k
                                         IColumn& to) const override {
327
6.05k
        auto& col = assert_cast<ColumnFixedLengthObject&>(to);
328
6.05k
        DCHECK(col.item_size() == sizeof(Data))
329
2
                << "size is not equal: " << col.item_size() << " " << sizeof(Data);
330
6.05k
        size_t old_size = col.size();
331
6.05k
        col.resize(old_size + 1);
332
6.05k
        (reinterpret_cast<Data*>(col.get_data().data()) + old_size)->count =
333
6.05k
                AggregateFunctionCountNotNullUnary::data(place).count;
334
6.05k
    }
335
336
10.1k
    MutableColumnPtr create_serialize_column() const override {
337
10.1k
        return ColumnFixedLengthObject::create(sizeof(Data));
338
10.1k
    }
339
340
10.4k
    DataTypePtr get_serialized_type() const override {
341
10.4k
        return std::make_shared<DataTypeFixedLengthObject>();
342
10.4k
    }
343
344
    void add_range_single_place(int64_t partition_start, int64_t partition_end, int64_t frame_start,
345
                                int64_t frame_end, AggregateDataPtr place, const IColumn** columns,
346
                                Arena& arena, UInt8* use_null_result,
347
60.3k
                                UInt8* could_use_previous_result) const override {
348
60.3k
        frame_start = std::max<int64_t>(frame_start, partition_start);
349
60.3k
        frame_end = std::min<int64_t>(frame_end, partition_end);
350
60.3k
        if (frame_start >= frame_end) {
351
12
            if (!*could_use_previous_result) {
352
0
                *use_null_result = true;
353
0
            }
354
60.2k
        } else {
355
60.2k
            const auto& nullable_column =
356
60.2k
                    assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*columns[0]);
357
60.2k
            size_t count = 0;
358
60.2k
            if (nullable_column.has_null()) {
359
119k
                for (int64_t i = frame_start; i < frame_end; ++i) {
360
63.6k
                    if (!nullable_column.is_null_at(i)) {
361
60.8k
                        ++count;
362
60.8k
                    }
363
63.6k
                }
364
55.8k
            } else {
365
4.46k
                count = frame_end - frame_start;
366
4.46k
            }
367
60.2k
            *use_null_result = false;
368
60.2k
            *could_use_previous_result = true;
369
60.2k
            AggregateFunctionCountNotNullUnary::data(place).count += count;
370
60.2k
        }
371
60.3k
    }
372
};
373
374
} // namespace doris