Skip to main content

nautilus_persistence/backend/
kmerge_batch.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
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// -------------------------------------------------------------------------------------------------
15
16use std::vec::IntoIter;
17
18use futures::{Stream, StreamExt};
19use tokio::{
20    sync::mpsc::{self, Receiver},
21    task::JoinHandle,
22};
23
24use super::{
25    binary_heap::{BinaryHeap, PeekMut},
26    compare::Compare,
27};
28
29pub struct EagerStream<T> {
30    rx: Receiver<T>,
31    task: JoinHandle<()>,
32    runtime: tokio::runtime::Handle,
33}
34
35impl<T> EagerStream<T> {
36    pub fn from_stream_with_runtime<S>(stream: S, runtime: tokio::runtime::Handle) -> Self
37    where
38        S: Stream<Item = T> + Send + 'static,
39        T: Send + 'static,
40    {
41        let (tx, rx) = mpsc::channel(1);
42
43        let task = runtime.spawn(async move {
44            futures::pin_mut!(stream);
45            while let Some(item) = stream.next().await {
46                if tx.send(item).await.is_err() {
47                    break;
48                }
49            }
50        });
51
52        Self { rx, task, runtime }
53    }
54}
55
56impl<T> Iterator for EagerStream<T> {
57    type Item = T;
58
59    fn next(&mut self) -> Option<Self::Item> {
60        super::block_on(&self.runtime, self.rx.recv())
61    }
62}
63
64impl<T> Drop for EagerStream<T> {
65    fn drop(&mut self) {
66        self.rx.close();
67        self.task.abort();
68    }
69}
70
71// TODO: Investigate implementing Iterator for ElementBatchIter
72// to reduce next element duplication. May be difficult to make it peekable.
73pub struct ElementBatchIter<I, T>
74where
75    I: Iterator<Item = IntoIter<T>>,
76{
77    pub item: T,
78    batch: I::Item,
79    iter: I,
80}
81
82impl<I, T> ElementBatchIter<I, T>
83where
84    I: Iterator<Item = IntoIter<T>>,
85{
86    fn new_from_iter(mut iter: I) -> Option<Self> {
87        loop {
88            let Some(mut batch) = iter.next() else {
89                break None;
90            };
91
92            if let Some(item) = batch.next() {
93                break Some(Self { item, batch, iter });
94            }
95        }
96    }
97}
98
99pub struct KMerge<I, T, C>
100where
101    I: Iterator<Item = IntoIter<T>>,
102{
103    heap: BinaryHeap<ElementBatchIter<I, T>, C>,
104}
105
106impl<I, T, C> KMerge<I, T, C>
107where
108    I: Iterator<Item = IntoIter<T>>,
109    C: Compare<ElementBatchIter<I, T>>,
110{
111    /// Creates a new [`KMerge`] instance.
112    pub fn new(cmp: C) -> Self {
113        Self {
114            heap: BinaryHeap::from_vec_cmp(Vec::new(), cmp),
115        }
116    }
117
118    pub fn push_iter(&mut self, s: I) {
119        if let Some(heap_elem) = ElementBatchIter::new_from_iter(s) {
120            self.heap.push(heap_elem);
121        }
122    }
123
124    pub fn clear(&mut self) {
125        self.heap.clear();
126    }
127}
128
129impl<I, T, C> Iterator for KMerge<I, T, C>
130where
131    I: Iterator<Item = IntoIter<T>>,
132    C: Compare<ElementBatchIter<I, T>>,
133{
134    type Item = T;
135
136    fn next(&mut self) -> Option<Self::Item> {
137        match self.heap.peek_mut() {
138            Some(mut heap_elem) => {
139                // Get next element from batch
140                match heap_elem.batch.next() {
141                    // Swap current heap element with new element
142                    // return the old element
143                    Some(mut item) => {
144                        std::mem::swap(&mut item, &mut heap_elem.item);
145                        Some(item)
146                    }
147                    // Otherwise get the next batch and the element from it
148                    // Unless the underlying iterator is exhausted
149                    None => loop {
150                        let Some(mut batch) = heap_elem.iter.next() else {
151                            let ElementBatchIter {
152                                item,
153                                batch: _,
154                                iter: _,
155                            } = PeekMut::pop(heap_elem);
156                            break Some(item);
157                        };
158
159                        if let Some(mut item) = batch.next() {
160                            heap_elem.batch = batch;
161                            std::mem::swap(&mut item, &mut heap_elem.item);
162                            break Some(item);
163                        }
164                    },
165                }
166            }
167            None => None,
168        }
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use proptest::prelude::*;
175    use rstest::rstest;
176
177    use super::*;
178
179    struct OrdComparator;
180    impl<S> Compare<ElementBatchIter<S, i32>> for OrdComparator
181    where
182        S: Iterator<Item = IntoIter<i32>>,
183    {
184        fn compare(
185            &self,
186            l: &ElementBatchIter<S, i32>,
187            r: &ElementBatchIter<S, i32>,
188        ) -> std::cmp::Ordering {
189            // Max heap ordering must be reversed
190            l.item.cmp(&r.item).reverse()
191        }
192    }
193
194    impl<S> Compare<ElementBatchIter<S, u64>> for OrdComparator
195    where
196        S: Iterator<Item = IntoIter<u64>>,
197    {
198        fn compare(
199            &self,
200            l: &ElementBatchIter<S, u64>,
201            r: &ElementBatchIter<S, u64>,
202        ) -> std::cmp::Ordering {
203            // Max heap ordering must be reversed
204            l.item.cmp(&r.item).reverse()
205        }
206    }
207
208    #[rstest]
209    fn test1() {
210        let iter_a = vec![vec![1, 2, 3].into_iter(), vec![7, 8, 9].into_iter()].into_iter();
211        let iter_b = vec![vec![4, 5, 6].into_iter()].into_iter();
212        let mut kmerge: KMerge<_, i32, _> = KMerge::new(OrdComparator);
213        kmerge.push_iter(iter_a);
214        kmerge.push_iter(iter_b);
215
216        let values: Vec<i32> = kmerge.collect();
217        assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]);
218    }
219
220    #[rstest]
221    fn test2() {
222        let iter_a = vec![vec![1, 2, 6].into_iter(), vec![7, 8, 9].into_iter()].into_iter();
223        let iter_b = vec![vec![3, 4, 5, 6].into_iter()].into_iter();
224        let mut kmerge: KMerge<_, i32, _> = KMerge::new(OrdComparator);
225        kmerge.push_iter(iter_a);
226        kmerge.push_iter(iter_b);
227
228        let values: Vec<i32> = kmerge.collect();
229        assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 6, 7, 8, 9]);
230    }
231
232    #[rstest]
233    fn test3() {
234        let iter_a = vec![vec![1, 4, 7].into_iter(), vec![24, 35, 56].into_iter()].into_iter();
235        let iter_b = vec![vec![2, 4, 8].into_iter()].into_iter();
236        let iter_c = vec![vec![3, 5, 9].into_iter(), vec![12, 12, 90].into_iter()].into_iter();
237        let mut kmerge: KMerge<_, i32, _> = KMerge::new(OrdComparator);
238        kmerge.push_iter(iter_a);
239        kmerge.push_iter(iter_b);
240        kmerge.push_iter(iter_c);
241
242        let values: Vec<i32> = kmerge.collect();
243        assert_eq!(
244            values,
245            vec![1, 2, 3, 4, 4, 5, 7, 8, 9, 12, 12, 24, 35, 56, 90]
246        );
247    }
248
249    #[rstest]
250    fn test5() {
251        let iter_a = vec![
252            vec![1, 3, 5].into_iter(),
253            vec![].into_iter(),
254            vec![7, 9, 11].into_iter(),
255        ]
256        .into_iter();
257        let iter_b = vec![vec![2, 4, 6].into_iter()].into_iter();
258        let mut kmerge: KMerge<_, i32, _> = KMerge::new(OrdComparator);
259        kmerge.push_iter(iter_a);
260        kmerge.push_iter(iter_b);
261
262        let values: Vec<i32> = kmerge.collect();
263        assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 9, 11]);
264    }
265
266    #[derive(Debug, Clone)]
267    struct SortedNestedVec(Vec<Vec<u64>>);
268
269    /// Strategy to generate nested vectors where each inner vector is sorted.
270    fn sorted_nested_vec_strategy() -> impl Strategy<Value = SortedNestedVec> {
271        // Generate a vector of u64 values, then split into sorted chunks
272        prop::collection::vec(any::<u64>(), 0..=100).prop_flat_map(|mut flat_vec| {
273            flat_vec.sort_unstable();
274
275            // Generate chunk sizes that will split the sorted vector
276            let total_len = flat_vec.len();
277            if total_len == 0 {
278                return Just(SortedNestedVec(vec![vec![]])).boxed();
279            }
280
281            // Generate random chunk boundaries
282            prop::collection::vec(0..=total_len, 0..=10)
283                .prop_map(move |mut boundaries| {
284                    boundaries.push(0);
285                    boundaries.push(total_len);
286                    boundaries.sort_unstable();
287                    boundaries.dedup();
288
289                    let mut nested_vec = Vec::new();
290                    for [start, end] in boundaries.array_windows() {
291                        nested_vec.push(flat_vec[*start..*end].to_vec());
292                    }
293
294                    SortedNestedVec(nested_vec)
295                })
296                .boxed()
297        })
298    }
299
300    proptest! {
301        /// Property: K-way merge should produce the same result as sorting all data together
302        #[rstest]
303        fn prop_kmerge_equivalent_to_sort(
304            all_data in prop::collection::vec(sorted_nested_vec_strategy(), 0..=10)
305        ) {
306            let mut kmerge: KMerge<_, u64, _> = KMerge::new(OrdComparator);
307
308            let copy_data = all_data.clone();
309            for stream in copy_data {
310                let input = stream.0.into_iter().map(std::iter::IntoIterator::into_iter);
311                kmerge.push_iter(input);
312            }
313            let merged_data: Vec<u64> = kmerge.collect();
314
315            let mut sorted_data: Vec<u64> = all_data
316                .into_iter()
317                .flat_map(|stream| stream.0.into_iter().flatten())
318                .collect();
319            sorted_data.sort_unstable();
320
321            prop_assert_eq!(merged_data.len(), sorted_data.len(), "Lengths should be equal");
322            prop_assert_eq!(merged_data, sorted_data, "Merged data should equal sorted data");
323        }
324
325        /// Property: K-way merge should preserve sortedness when inputs are sorted
326        #[rstest]
327        fn prop_kmerge_preserves_sort_order(
328            all_data in prop::collection::vec(sorted_nested_vec_strategy(), 1..=5)
329        ) {
330            let mut kmerge: KMerge<_, u64, _> = KMerge::new(OrdComparator);
331
332            for stream in all_data {
333                let input = stream.0.into_iter().map(std::iter::IntoIterator::into_iter);
334                kmerge.push_iter(input);
335            }
336            let merged_data: Vec<u64> = kmerge.collect();
337
338            // Check that the merged data is sorted
339            for [a, b] in merged_data.array_windows() {
340                prop_assert!(a <= b, "Merged data should be sorted");
341            }
342        }
343
344        /// Property: Empty iterators should not affect the merge result
345        #[rstest]
346        fn prop_kmerge_handles_empty_iterators(
347            data in sorted_nested_vec_strategy(),
348            empty_count in 0usize..=5
349        ) {
350            let mut kmerge_with_empty: KMerge<_, u64, _> = KMerge::new(OrdComparator);
351            let mut kmerge_without_empty: KMerge<_, u64, _> = KMerge::new(OrdComparator);
352
353            // Add the actual data to both merges
354            let input_with_empty = data.0.clone().into_iter().map(std::iter::IntoIterator::into_iter);
355            let input_without_empty = data.0.into_iter().map(std::iter::IntoIterator::into_iter);
356
357            kmerge_with_empty.push_iter(input_with_empty);
358            kmerge_without_empty.push_iter(input_without_empty);
359
360            // Add empty iterators to the first merge
361            for _ in 0..empty_count {
362                let empty_vec: Vec<Vec<u64>> = vec![];
363                let empty_input = empty_vec.into_iter().map(std::iter::IntoIterator::into_iter);
364                kmerge_with_empty.push_iter(empty_input);
365            }
366
367            let result_with_empty: Vec<u64> = kmerge_with_empty.collect();
368            let result_without_empty: Vec<u64> = kmerge_without_empty.collect();
369
370            prop_assert_eq!(result_with_empty, result_without_empty, "Empty iterators should not affect result");
371        }
372    }
373}