1use 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
71pub 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 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 match heap_elem.batch.next() {
141 Some(mut item) => {
144 std::mem::swap(&mut item, &mut heap_elem.item);
145 Some(item)
146 }
147 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 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 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 fn sorted_nested_vec_strategy() -> impl Strategy<Value = SortedNestedVec> {
271 prop::collection::vec(any::<u64>(), 0..=100).prop_flat_map(|mut flat_vec| {
273 flat_vec.sort_unstable();
274
275 let total_len = flat_vec.len();
277 if total_len == 0 {
278 return Just(SortedNestedVec(vec![vec![]])).boxed();
279 }
280
281 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 #[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 #[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 for [a, b] in merged_data.array_windows() {
340 prop_assert!(a <= b, "Merged data should be sorted");
341 }
342 }
343
344 #[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 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 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}