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#[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 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 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#[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 if n == 0 {
72 return futures::stream::iter(vec![]);
73 }
74
75 let mut indexed_stream = sinput.enumerate();
77
78 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 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 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 first_n.len() < n {
101 return futures::stream::iter(
102 first_n
103 .into_sorted_vec()
104 .into_iter()
105 .map(|(_idx, item)| item) .collect::<Vec<_>>(),
107 );
108 }
109
110 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(¤t_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) .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 if n == 0 {
202 return futures::stream::iter(vec![].into_iter());
203 }
204
205 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 first_n.len() < n {
222 return futures::stream::iter(first_n.into_sorted_vec().into_iter());
223 }
224
225 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 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))) .await
260 .keep_first_n(2, |a, b| a.cmp(b)) .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 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 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}