Skip to main content

risingwave_batch_executors/executor/
group_top_n.rs

1// Copyright 2024 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::marker::PhantomData;
16use std::mem::swap;
17use std::sync::Arc;
18
19use futures_async_stream::try_stream;
20use hashbrown::HashMap;
21use itertools::Itertools;
22use prometheus::core::Atomic;
23use risingwave_common::array::DataChunk;
24use risingwave_common::bitmap::FilterByBitmap;
25use risingwave_common::catalog::Schema;
26use risingwave_common::hash::{HashKey, HashKeyDispatcher, PrecomputedBuildHasher};
27use risingwave_common::memory::{MemoryContext, MonitoredGlobalAlloc};
28use risingwave_common::metrics::TrAdderAtomic;
29use risingwave_common::types::DataType;
30use risingwave_common::util::chunk_coalesce::DataChunkBuilder;
31use risingwave_common::util::iter_util::ZipEqFast;
32use risingwave_common::util::memcmp_encoding::encode_chunk;
33use risingwave_common::util::sort_util::ColumnOrder;
34use risingwave_pb::batch_plan::plan_node::NodeBody;
35
36use super::top_n::{HeapElem, TopNHeap};
37use crate::error::{BatchError, Result};
38use crate::executor::{
39    BoxedDataChunkStream, BoxedExecutor, BoxedExecutorBuilder, Executor, ExecutorBuilder,
40};
41
42/// Group Top-N Executor
43///
44/// For each group, use a N-heap to store the smallest N rows.
45pub struct GroupTopNExecutor<K: HashKey> {
46    child: BoxedExecutor,
47    column_orders: Vec<ColumnOrder>,
48    offset: usize,
49    limit: usize,
50    group_key: Vec<usize>,
51    with_ties: bool,
52    schema: Schema,
53    identity: String,
54    chunk_size: usize,
55    mem_ctx: MemoryContext,
56    _phantom: PhantomData<K>,
57}
58
59pub struct GroupTopNExecutorBuilder {
60    child: BoxedExecutor,
61    column_orders: Vec<ColumnOrder>,
62    offset: usize,
63    limit: usize,
64    group_key: Vec<usize>,
65    group_key_types: Vec<DataType>,
66    with_ties: bool,
67    identity: String,
68    chunk_size: usize,
69    mem_ctx: MemoryContext,
70}
71
72impl HashKeyDispatcher for GroupTopNExecutorBuilder {
73    type Output = BoxedExecutor;
74
75    fn dispatch_impl<K: HashKey>(self) -> Self::Output {
76        Box::new(GroupTopNExecutor::<K>::new(
77            self.child,
78            self.column_orders,
79            self.offset,
80            self.limit,
81            self.with_ties,
82            self.group_key,
83            self.identity,
84            self.chunk_size,
85            self.mem_ctx,
86        ))
87    }
88
89    fn data_types(&self) -> &[DataType] {
90        &self.group_key_types
91    }
92}
93
94impl BoxedExecutorBuilder for GroupTopNExecutorBuilder {
95    async fn new_boxed_executor(
96        source: &ExecutorBuilder<'_>,
97        inputs: Vec<BoxedExecutor>,
98    ) -> Result<BoxedExecutor> {
99        let [child]: [_; 1] = inputs.try_into().unwrap();
100
101        let top_n_node = try_match_expand!(
102            source.plan_node().get_node_body().unwrap(),
103            NodeBody::GroupTopN
104        )?;
105
106        let column_orders = top_n_node
107            .column_orders
108            .iter()
109            .map(ColumnOrder::from_protobuf)
110            .collect();
111
112        let group_key = top_n_node
113            .group_key
114            .iter()
115            .map(|x| *x as usize)
116            .collect_vec();
117        let child_schema = child.schema();
118        let group_key_types = group_key
119            .iter()
120            .map(|x| child_schema.fields[*x].data_type())
121            .collect();
122
123        let identity = source.plan_node().get_identity().clone();
124
125        let builder = Self {
126            child,
127            column_orders,
128            offset: top_n_node.get_offset() as usize,
129            limit: top_n_node.get_limit() as usize,
130            group_key,
131            group_key_types,
132            with_ties: top_n_node.get_with_ties(),
133            identity: identity.clone(),
134            chunk_size: source.context().get_config().developer.chunk_size,
135            mem_ctx: source.context().create_executor_mem_context(&identity),
136        };
137
138        Ok(builder.dispatch())
139    }
140}
141
142impl<K: HashKey> GroupTopNExecutor<K> {
143    pub fn new(
144        child: BoxedExecutor,
145        column_orders: Vec<ColumnOrder>,
146        offset: usize,
147        limit: usize,
148        with_ties: bool,
149        group_key: Vec<usize>,
150        identity: String,
151        chunk_size: usize,
152        mem_ctx: MemoryContext,
153    ) -> Self {
154        let schema = child.schema().clone();
155        Self {
156            child,
157            column_orders,
158            offset,
159            limit,
160            with_ties,
161            group_key,
162            schema,
163            identity,
164            chunk_size,
165            mem_ctx,
166            _phantom: PhantomData,
167        }
168    }
169}
170
171impl<K: HashKey> Executor for GroupTopNExecutor<K> {
172    fn schema(&self) -> &Schema {
173        &self.schema
174    }
175
176    fn identity(&self) -> &str {
177        &self.identity
178    }
179
180    fn execute(self: Box<Self>) -> BoxedDataChunkStream {
181        self.do_execute()
182    }
183}
184
185impl<K: HashKey> GroupTopNExecutor<K> {
186    #[try_stream(boxed, ok = DataChunk, error = BatchError)]
187    async fn do_execute(self: Box<Self>) {
188        if self.limit == 0 {
189            return Ok(());
190        }
191        // Keep all groups' payload charges until execution exits, including errors and cancellation.
192        let mem_ctx = MemoryContext::new(Some(self.mem_ctx.clone()), TrAdderAtomic::new(0));
193        let mut groups =
194            HashMap::<K, TopNHeap, PrecomputedBuildHasher, MonitoredGlobalAlloc>::with_hasher_in(
195                PrecomputedBuildHasher,
196                mem_ctx.global_allocator(),
197            );
198
199        #[for_await]
200        for chunk in self.child.execute() {
201            let chunk = Arc::new(chunk?);
202            let keys = K::build_many(self.group_key.as_slice(), &chunk);
203
204            for (row_id, (encoded_row, key)) in encode_chunk(&chunk, &self.column_orders)?
205                .into_iter()
206                .zip_eq_fast(keys.into_iter())
207                .enumerate()
208                .filter_by_bitmap(chunk.visibility())
209            {
210                let heap = groups.entry(key).or_insert_with(|| {
211                    TopNHeap::new(self.limit, self.offset, self.with_ties, mem_ctx.clone())
212                });
213                heap.push(HeapElem::new(encoded_row, chunk.row_at(row_id).0));
214            }
215        }
216
217        let mut chunk_builder = DataChunkBuilder::new(self.schema.data_types(), self.chunk_size);
218        for (_, h) in &mut groups {
219            let mut heap = TopNHeap::empty();
220            swap(&mut heap, h);
221            for elem in heap.dump() {
222                let output = chunk_builder.append_one_row(elem.row());
223                drop(elem);
224                if let Some(output) = output {
225                    yield output
226                }
227            }
228        }
229        if let Some(spilled) = chunk_builder.consume_all() {
230            yield spilled
231        }
232    }
233}
234
235#[cfg(test)]
236mod tests {
237    use futures::stream::StreamExt;
238    use risingwave_common::catalog::Field;
239    use risingwave_common::metrics::LabelGuardedIntGauge;
240    use risingwave_common::test_prelude::DataChunkTestExt;
241    use risingwave_common::util::sort_util::OrderType;
242
243    use super::*;
244    use crate::executor::test_utils::MockExecutor;
245
246    const CHUNK_SIZE: usize = 1024;
247
248    #[tokio::test]
249    async fn test_group_top_n_executor() {
250        let parent_mem = MemoryContext::root(LabelGuardedIntGauge::test_int_gauge::<4>(), u64::MAX);
251        {
252            let schema = Schema {
253                fields: vec![
254                    Field::unnamed(DataType::Int32),
255                    Field::unnamed(DataType::Int32),
256                    Field::unnamed(DataType::Int32),
257                ],
258            };
259            let mut mock_executor = MockExecutor::new(schema);
260            mock_executor.add(DataChunk::from_pretty(
261                "i i i
262             1 5 1
263             2 4 1
264             3 3 1
265             4 2 1
266             5 1 1
267             1 6 2
268             2 5 2
269             3 4 2
270             4 3 2
271             5 2 2
272             ",
273            ));
274            let column_orders = vec![
275                ColumnOrder {
276                    column_index: 1,
277                    order_type: OrderType::ascending(),
278                },
279                ColumnOrder {
280                    column_index: 0,
281                    order_type: OrderType::ascending(),
282                },
283            ];
284            let mem_ctx = MemoryContext::new(
285                Some(parent_mem.clone()),
286                LabelGuardedIntGauge::test_int_gauge::<4>(),
287            );
288            let top_n_executor = (GroupTopNExecutorBuilder {
289                child: Box::new(mock_executor),
290                column_orders,
291                offset: 1,
292                limit: 3,
293                with_ties: false,
294                group_key: vec![2],
295                group_key_types: vec![DataType::Int32],
296                identity: "GroupTopNExecutor".to_owned(),
297                chunk_size: CHUNK_SIZE,
298                mem_ctx,
299            })
300            .dispatch();
301
302            let fields = &top_n_executor.schema().fields;
303            assert_eq!(fields[0].data_type, DataType::Int32);
304            assert_eq!(fields[1].data_type, DataType::Int32);
305
306            let mut stream = top_n_executor.execute();
307            let res = stream.next().await;
308
309            assert!(res.is_some());
310            if let Some(res) = res {
311                let res = res.unwrap();
312                assert!(
313                    res == DataChunk::from_pretty(
314                        "
315                    i i i
316                    4 2 1
317                    3 3 1
318                    2 4 1
319                    4 3 2
320                    3 4 2
321                    2 5 2
322                    "
323                    ) || res
324                        == DataChunk::from_pretty(
325                            "
326                    i i i
327                    4 3 2
328                    3 4 2
329                    2 5 2
330                    4 2 1
331                    3 3 1
332                    2 4 1
333                    "
334                        )
335                );
336            }
337
338            let res = stream.next().await;
339            assert!(res.is_none());
340        }
341
342        assert_eq!(0, parent_mem.get_bytes_used());
343    }
344}