Skip to main content

marigold_impl/
keep_first_n.rs

1use async_trait::async_trait;
2use binary_heap_plus::BinaryHeap;
3use futures::stream::Stream;
4use futures::stream::StreamExt;
5use std::cmp::Ordering;
6use tracing::instrument;
7
8// Batch size for ready_chunks: amortizes spawn overhead (~8,405 ns) over many items (~28 ns each).
9// Measured crossover is ~300 items; 256 is a power-of-2 heuristic just below that. Workloads with
10// heavier comparators would benefit from a smaller value; trivially cheap ones from a larger one.
11#[cfg(any(feature = "tokio", feature = "async-std"))]
12const READY_CHUNK_SIZE: usize = 256;
13
14#[async_trait]
15pub trait KeepFirstN<T, F>
16where
17    F: Fn(&T, &T) -> Ordering,
18{
19    /// Takes the largest N values according to the sorted function, returned in descending order
20    /// (max first). Exhausts the stream.
21    async fn keep_first_n(
22        self,
23        n: usize,
24        sorted_by: F,
25    ) -> futures::stream::Iter<std::vec::IntoIter<T>>;
26}
27
28#[cfg(any(feature = "tokio", feature = "async-std"))]
29#[async_trait]
30impl<SInput, T, F> KeepFirstN<T, F> for SInput
31where
32    SInput: Stream<Item = T> + Send + Unpin + std::marker::Sync + 'static,
33    T: Clone + Send + std::marker::Sync + std::fmt::Debug + 'static,
34    F: Fn(&T, &T) -> Ordering + std::marker::Send + std::marker::Sync + std::marker::Copy + 'static,
35{
36    #[instrument(skip(self, sorted_by))]
37    async fn keep_first_n(
38        mut self,
39        n: usize,
40        sorted_by: F,
41    ) -> futures::stream::Iter<std::vec::IntoIter<T>> {
42        // use the reverse ordering so that the smallest value is always the first to pop.
43        let first_n = BinaryHeap::with_capacity_by(n, move |a, b| sorted_by(a, b).reverse());
44        impl_keep_first_n(self, first_n, n, sorted_by).await
45    }
46}
47
48/// Internal logic for keep_first_n. This is in a separate function so that we can get the full
49/// type of the binary heap, which includes a lambda for reversing the ordering fromt the passed
50/// sort_by function. By declaring a new function, we can use generics to describe its type, and
51/// then can use that type while unsafely casting pointers.
52///
53/// This implementation wraps items with their stream index to provide deterministic tie-breaking
54/// when the user's comparison function returns Equal. Lower indices (earlier in stream) are
55/// preferred to ensure consistent results even with parallel processing.
56#[cfg(any(feature = "tokio", feature = "async-std"))]
57async fn impl_keep_first_n<SInput, T, F, FReversed>(
58    sinput: SInput,
59    _first_n: BinaryHeap<T, binary_heap_plus::FnComparator<FReversed>>,
60    n: usize,
61    sorted_by: F,
62) -> futures::stream::Iter<std::vec::IntoIter<T>>
63where
64    SInput: Stream<Item = T> + Send + Unpin + std::marker::Sync + 'static,
65    T: Clone + Send + std::marker::Sync + std::fmt::Debug + 'static,
66    F: Fn(&T, &T) -> Ordering + std::marker::Send + std::marker::Sync + std::marker::Copy + 'static,
67    FReversed: Fn(&T, &T) -> std::cmp::Ordering + Clone + Send + 'static,
68{
69    // n=0 means keep nothing; return an empty stream immediately without touching the heap
70    // (the heap is empty, so peek().unwrap() would panic below).
71    if n == 0 {
72        return futures::stream::iter(vec![]);
73    }
74
75    // Add indices to items for deterministic tie-breaking
76    let mut indexed_stream = sinput.enumerate();
77
78    // Create a heap that stores (index, item) tuples with tie-breaking comparator
79    let indexed_comparator = move |a: &(usize, T), b: &(usize, T)| {
80        match sorted_by(&a.1, &b.1) {
81            Ordering::Less => Ordering::Less,
82            Ordering::Greater => Ordering::Greater,
83            // When equal, prefer lower index (earlier in stream)
84            Ordering::Equal => a.0.cmp(&b.0),
85        }
86    };
87    let mut first_n =
88        BinaryHeap::with_capacity_by(n, move |a, b| indexed_comparator(a, b).reverse());
89
90    // Iterate through values in a single thread until we have seen n values.
91    while first_n.len() < n {
92        if let Some(indexed_item) = indexed_stream.next().await {
93            first_n.push(indexed_item);
94        } else {
95            break;
96        }
97    }
98
99    // If we have exhausted the stream before reaching n values, we can exit early.
100    if first_n.len() < n {
101        return futures::stream::iter(
102            first_n
103                .into_sorted_vec()
104                .into_iter()
105                .map(|(_idx, item)| item) // Unwrap indices
106                .collect::<Vec<_>>(),
107        );
108    }
109
110    // Otherwise, we can check each remaining value in the stream against the smallest
111    // kept value, updating the kept values only when a keepable value is found. This
112    // is done by spawning tasks, which can be parallelized by multithreaded runtimes.
113    //
114    // A double-check pattern is used: the RwLock provides a fast-path filter
115    // (most items are rejected without touching the mutex), and after acquiring
116    // the mutex the condition is re-checked against the current heap state to
117    // eliminate the TOCTOU race where two tasks could both pass the fast-path
118    // check but only one should actually replace the smallest kept value.
119    let first_n_mutex = std::sync::Arc::new(parking_lot::Mutex::new(first_n));
120    let smallest_kept = std::sync::Arc::new(parking_lot::RwLock::new(
121        first_n_mutex.lock().peek().unwrap().to_owned(),
122    ));
123    let first_n_arc = first_n_mutex.clone();
124    let smallest_kept_arc = smallest_kept.clone();
125    let parallel_work = async move {
126        let mut ongoing_tasks = indexed_stream
127            .ready_chunks(READY_CHUNK_SIZE)
128            .map(move |chunk: Vec<(usize, T)>| {
129                let first_n_arc = first_n_arc.clone();
130                let smallest_kept_arc = smallest_kept_arc.clone();
131                crate::async_runtime::spawn(async move {
132                    #[cfg(feature = "bench-instrumentation")]
133                    let _worker_span = tracing::info_span!("keep_first_n_worker_task").entered();
134                    for indexed_item in chunk {
135                        let smallest = smallest_kept_arc.read();
136                        let should_keep = match sorted_by(&smallest.1, &indexed_item.1) {
137                            Ordering::Less => true,
138                            Ordering::Greater => false,
139                            Ordering::Equal => indexed_item.0 < smallest.0,
140                        };
141                        drop(smallest);
142
143                        if should_keep {
144                            let mut update_first_n = first_n_arc.lock();
145                            let current_smallest = update_first_n.peek().unwrap();
146                            let still_should_keep =
147                                match sorted_by(&current_smallest.1, &indexed_item.1) {
148                                    Ordering::Less => true,
149                                    Ordering::Greater => false,
150                                    Ordering::Equal => indexed_item.0 < current_smallest.0,
151                                };
152                            if still_should_keep {
153                                update_first_n.pop();
154                                update_first_n.push(indexed_item);
155                                let mut update_smallest_kept = smallest_kept_arc.write();
156                                *update_smallest_kept = update_first_n.peek().unwrap().to_owned();
157                            }
158                        }
159                    }
160                })
161            })
162            .buffer_unordered(num_cpus::get() * 4);
163        while let Some(_task) = ongoing_tasks.next().await {}
164    };
165    #[cfg(feature = "bench-instrumentation")]
166    {
167        use tracing::Instrument;
168        parallel_work
169            .instrument(tracing::info_span!("keep_first_n_parallel_section"))
170            .await;
171    }
172    #[cfg(not(feature = "bench-instrumentation"))]
173    parallel_work.await;
174    futures::stream::iter(
175        std::sync::Arc::try_unwrap(first_n_mutex)
176            .expect("Dangling references to mutex")
177            .into_inner()
178            .into_sorted_vec()
179            .into_iter()
180            .map(|(_idx, item)| item) // Unwrap indices
181            .collect::<Vec<_>>(),
182    )
183}
184
185#[async_trait]
186#[cfg(not(any(feature = "tokio", feature = "async-std")))]
187impl<SInput, T, F> KeepFirstN<T, F> for SInput
188where
189    SInput: Stream<Item = T> + Send + Unpin,
190    T: Clone + Send + std::marker::Sync,
191    F: Fn(&T, &T) -> Ordering + std::marker::Send + std::marker::Sync + 'static,
192{
193    #[instrument(skip(self, sorted_by))]
194    async fn keep_first_n(
195        mut self,
196        n: usize,
197        sorted_by: F,
198    ) -> futures::stream::Iter<std::vec::IntoIter<T>> {
199        // n=0 means keep nothing; return an empty stream immediately without touching the heap
200        // (the heap is empty, so peek().unwrap() would panic below).
201        if n == 0 {
202            return futures::stream::iter(vec![].into_iter());
203        }
204
205        // use the reverse ordering so that the smallest value is always the first to pop.
206        let mut first_n = BinaryHeap::with_capacity_by(n, |a, b| match sorted_by(a, b) {
207            Ordering::Less => Ordering::Greater,
208            Ordering::Equal => Ordering::Equal,
209            Ordering::Greater => Ordering::Less,
210        });
211
212        while first_n.len() < n {
213            if let Some(item) = self.next().await {
214                first_n.push(item);
215            } else {
216                break;
217            }
218        }
219
220        // If we have exhausted the stream before reaching n values, we can exit early.
221        if first_n.len() < n {
222            return futures::stream::iter(first_n.into_sorted_vec().into_iter());
223        }
224
225        // Otherwise, we can check each remaining value in the stream against the smallest
226        // kept value, updating the kept values only when a keepable value is found.
227        let first_n_mutex = parking_lot::Mutex::new(first_n);
228        let smallest_kept =
229            parking_lot::RwLock::new(first_n_mutex.lock().peek().unwrap().to_owned());
230
231        self.for_each_concurrent(
232            /* arbitrarily set concurrency limit */ 256,
233            |item| async {
234                if sorted_by(&*smallest_kept.read(), &item) == Ordering::Less {
235                    let mut first_n_mut = first_n_mutex.lock();
236                    first_n_mut.pop();
237                    first_n_mut.push(item);
238                    let mut update_smallest_kept = smallest_kept.write();
239                    *update_smallest_kept = first_n_mut.peek().unwrap().to_owned();
240                }
241            },
242        )
243        .await;
244
245        futures::stream::iter(first_n_mutex.into_inner().into_sorted_vec().into_iter())
246    }
247}
248
249#[cfg(test)]
250mod tests {
251    use super::KeepFirstN;
252    use futures::stream::StreamExt;
253
254    #[tokio::test]
255    async fn keep_first_n() {
256        assert_eq!(
257            futures::stream::iter(1..10)
258                .keep_first_n(5, |a, b| (a % 2).cmp(&(b % 2))) // keep odd numbers
259                .await
260                .keep_first_n(2, |a, b| a.cmp(b)) // keep largest odd 2 numbers
261                .await
262                .collect::<Vec<_>>()
263                .await,
264            vec![9, 7]
265        );
266    }
267
268    #[tokio::test]
269    async fn large_stream_correctness() {
270        let items: Vec<u64> = (0..10_000).map(|i| (i * 7 + 3) % 10_000).collect();
271        let mut expected: Vec<u64> = items.clone();
272        expected.sort_by(|a, b| b.cmp(a));
273        expected.truncate(10);
274
275        let result = futures::stream::iter(items)
276            .keep_first_n(10, |a, b| a.cmp(b))
277            .await
278            .collect::<Vec<_>>()
279            .await;
280
281        assert_eq!(result, expected);
282    }
283
284    #[tokio::test]
285    async fn chunk_boundary_exact() {
286        let items: Vec<u32> = (0..256).collect();
287        let result = futures::stream::iter(items)
288            .keep_first_n(5, |a, b| a.cmp(b))
289            .await
290            .collect::<Vec<_>>()
291            .await;
292        assert_eq!(result, vec![255, 254, 253, 252, 251]);
293    }
294
295    #[tokio::test]
296    async fn chunk_boundary_plus_one() {
297        let items: Vec<u32> = (0..257).collect();
298        let result = futures::stream::iter(items)
299            .keep_first_n(5, |a, b| a.cmp(b))
300            .await
301            .collect::<Vec<_>>()
302            .await;
303        assert_eq!(result, vec![256, 255, 254, 253, 252]);
304    }
305
306    #[tokio::test]
307    async fn chunk_boundary_less_than_chunk() {
308        let items: Vec<u32> = (0..100).collect();
309        let result = futures::stream::iter(items)
310            .keep_first_n(5, |a, b| a.cmp(b))
311            .await
312            .collect::<Vec<_>>()
313            .await;
314        assert_eq!(result, vec![99, 98, 97, 96, 95]);
315    }
316
317    #[tokio::test]
318    async fn chunk_boundary_keep_all() {
319        let items: Vec<u32> = (0..10).collect();
320        let result = futures::stream::iter(items)
321            .keep_first_n(10, |a, b| a.cmp(b))
322            .await
323            .collect::<Vec<_>>()
324            .await;
325        assert_eq!(result, vec![9, 8, 7, 6, 5, 4, 3, 2, 1, 0]);
326    }
327
328    #[tokio::test]
329    async fn tie_breaking_determinism() {
330        let items: Vec<(u32, usize)> = (0..100).map(|i| (42u32, i)).collect();
331        let result = futures::stream::iter(items)
332            .keep_first_n(5, |a, b| a.0.cmp(&b.0))
333            .await
334            .collect::<Vec<_>>()
335            .await;
336
337        assert_eq!(result.len(), 5);
338        let mut indices: Vec<usize> = result.iter().map(|(_, i)| *i).collect();
339        indices.sort_unstable();
340        assert_eq!(indices, vec![0, 1, 2, 3, 4]);
341    }
342
343    #[tokio::test]
344    async fn stream_shorter_than_n() {
345        let result = futures::stream::iter(vec![3u32, 1, 2])
346            .keep_first_n(10, |a, b| a.cmp(b))
347            .await
348            .collect::<Vec<_>>()
349            .await;
350        assert_eq!(result, vec![3, 2, 1]);
351    }
352
353    #[tokio::test]
354    async fn cross_chunk_concurrent_correctness() {
355        let n = 10_000usize;
356        let items: Vec<u64> = (0..n as u64).collect();
357        let mut expected: Vec<u64> = items.clone();
358        expected.sort_by(|a, b| b.cmp(a));
359        expected.truncate(50);
360
361        let result = futures::stream::iter(items)
362            .keep_first_n(50, |a, b| a.cmp(b))
363            .await
364            .collect::<Vec<_>>()
365            .await;
366
367        assert_eq!(result, expected);
368    }
369
370    #[tokio::test]
371    async fn test_keep_first_n_single_element() {
372        // Keep only 1 element (the largest).
373        let result = futures::stream::iter(vec![5, 3, 8, 1])
374            .keep_first_n(1, |a, b| a.cmp(b))
375            .await
376            .collect::<Vec<_>>()
377            .await;
378        assert_eq!(result, vec![8]);
379    }
380
381    #[tokio::test]
382    async fn test_keep_first_n_empty_stream() {
383        // Empty stream should return empty results regardless of n.
384        let result = futures::stream::iter(Vec::<i32>::new())
385            .keep_first_n(5, |a, b| a.cmp(b))
386            .await
387            .collect::<Vec<_>>()
388            .await;
389        assert!(result.is_empty());
390    }
391}