Skip to main content

vortex_io/runtime/
current.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use std::future::Future;
5use std::sync::Arc;
6
7use futures::Stream;
8use futures::StreamExt;
9use futures::stream::BoxStream;
10use parking_lot::Mutex;
11use smol::block_on;
12use vortex_utils::parallelism::get_available_parallelism;
13
14use crate::runtime::BlockingRuntime;
15use crate::runtime::Executor;
16use crate::runtime::Handle;
17pub use crate::runtime::pool::CurrentThreadWorkerPool;
18
19/// A current thread runtime allows callers to much more explicitly drive Vortex futures than with
20/// a Tokio runtime.
21///
22/// The current thread runtime will do no work unless `block_on` is called. In other words, the
23/// default behavior is single-threaded with code running on the thread that called `block_on`.
24///
25/// It's also possible to clone the runtime onto other threads, each of which can call `block_on`
26/// to drive work on that thread. Each thread shares the same underlying executor with the same
27/// set of tasks, allowing work to be driven in parallel.
28///
29/// For automatic driving of work, a [`CurrentThreadWorkerPool`] can be created from the runtime
30/// by calling [`new_pool`](CurrentThreadRuntime::new_pool). The returned pool can be configured
31/// with the desired number of worker threads that will drive work on behalf of the runtime.
32#[derive(Clone, Default)]
33pub struct CurrentThreadRuntime {
34    executor: Arc<smol::Executor<'static>>,
35}
36
37impl CurrentThreadRuntime {
38    /// Create a new current thread runtime.
39    pub fn new() -> Self {
40        Self::default()
41    }
42
43    /// Create a new worker pool for driving the runtime in the background.
44    ///
45    /// This pool can be used to offload work from the current thread to a set of worker threads
46    /// that will drive the runtime's executor.
47    ///
48    /// By default, the pool has no worker threads; the caller must set the desired number of
49    /// worker threads using the `set_workers` method on the returned pool.
50    pub fn new_pool(&self) -> CurrentThreadWorkerPool {
51        CurrentThreadWorkerPool::new(Arc::clone(&self.executor))
52    }
53
54    /// Returns an iterator wrapper around a stream, blocking the current thread for each item.
55    ///
56    /// ## Multi-threaded Usage
57    ///
58    /// To drive the iterator from multiple threads, simply clone it and call `next()` on each
59    /// clone. Results on each thread are ordered with respect to the stream, but there is no
60    /// ordering guarantee between threads.
61    pub fn block_on_stream_thread_safe<F, S, R>(&self, f: F) -> ThreadSafeIterator<R>
62    where
63        F: FnOnce(Handle) -> S,
64        S: Stream<Item = R> + Send + 'static,
65        R: Send + 'static,
66    {
67        let stream = f(self.handle());
68
69        // We create an MPMC result channel and spawn a task to drive the stream and send results.
70        // This allows multiple worker threads to drive the execution while all waiting for results
71        // on the channel. Channel buffers up to one item per core so calling threads can get a
72        // ready item without every one of them having to become an executor.
73        let capacity = get_available_parallelism().unwrap_or(1).max(1);
74        let (result_tx, result_rx) = kanal::bounded_async(capacity);
75        let driver = self.executor.spawn(async move {
76            futures::pin_mut!(stream);
77            while let Some(item) = stream.next().await {
78                // If all receivers are dropped, we stop driving the stream.
79                if let Err(e) = result_tx.send(item).await {
80                    tracing::trace!("all receivers dropped, stopping stream: {}", e);
81                    break;
82                }
83            }
84        });
85
86        ThreadSafeIterator {
87            executor: Arc::clone(&self.executor),
88            results: result_rx,
89            driver: Arc::new(Mutex::new(Some(driver))),
90        }
91    }
92}
93
94impl BlockingRuntime for CurrentThreadRuntime {
95    type BlockingIterator<'a, R: 'a> = CurrentThreadIterator<'a, R>;
96
97    fn handle(&self) -> Handle {
98        let executor: Arc<dyn Executor> = Arc::clone(&self.executor) as Arc<dyn Executor>;
99        Handle::new(Arc::downgrade(&executor))
100    }
101
102    fn block_on<Fut, R>(&self, fut: Fut) -> R
103    where
104        Fut: Future<Output = R>,
105    {
106        block_on(self.executor.run(fut))
107    }
108
109    fn block_on_stream<'a, S, R>(&self, stream: S) -> Self::BlockingIterator<'a, R>
110    where
111        S: Stream<Item = R> + Send + 'a,
112        R: Send + 'a,
113    {
114        CurrentThreadIterator {
115            executor: Arc::clone(&self.executor),
116            stream: stream.boxed(),
117        }
118    }
119}
120
121/// An iterator that wraps up a stream to drive it using the current thread execution.
122pub struct CurrentThreadIterator<'a, T> {
123    executor: Arc<smol::Executor<'static>>,
124    stream: BoxStream<'a, T>,
125}
126
127impl<T> Iterator for CurrentThreadIterator<'_, T> {
128    type Item = T;
129
130    fn next(&mut self) -> Option<Self::Item> {
131        block_on(self.executor.run(self.stream.next()))
132    }
133}
134
135/// An iterator that drives a stream from multiple threads.
136pub struct ThreadSafeIterator<T> {
137    executor: Arc<smol::Executor<'static>>,
138    results: kanal::AsyncReceiver<T>,
139    /// Handle to the task driving the stream. Once the stream ends, the first consumer to
140    /// observe it joins the task so a panic raised while driving the stream is re-raised rather
141    /// than silently ending the iterator.
142    driver: Arc<Mutex<Option<smol::Task<()>>>>,
143}
144
145// Manual clone implementation since `T` does not need to be `Clone`.
146impl<T> Clone for ThreadSafeIterator<T> {
147    fn clone(&self) -> Self {
148        Self {
149            executor: Arc::clone(&self.executor),
150            results: self.results.clone(),
151            driver: Arc::clone(&self.driver),
152        }
153    }
154}
155
156impl<T> Iterator for ThreadSafeIterator<T> {
157    type Item = T;
158
159    fn next(&mut self) -> Option<Self::Item> {
160        // If driver already has an item, take it without driving the executor
161        match self.results.try_recv() {
162            Ok(Some(item)) => return Some(item),
163            Ok(None) => {}
164            Err(_) => return self.get_task_error(),
165        }
166
167        match block_on(self.executor.run(self.results.recv())) {
168            Ok(item) => Some(item),
169            Err(_) => self.get_task_error(),
170        }
171    }
172}
173
174impl<T> ThreadSafeIterator<T> {
175    // Join current task so panics are re-raised
176    fn get_task_error(&self) -> Option<T> {
177        let task = self.driver.lock().take();
178        if let Some(task) = task {
179            block_on(self.executor.run(task));
180        }
181        None
182    }
183}
184
185#[expect(clippy::if_then_some_else_none)] // Clippy is wrong when if/else has await.
186#[cfg(test)]
187mod tests {
188    use std::any::Any;
189    use std::panic::AssertUnwindSafe;
190    use std::sync::Arc;
191    use std::sync::Barrier;
192    use std::sync::atomic::AtomicUsize;
193    use std::sync::atomic::Ordering;
194    use std::task::Poll;
195    use std::thread;
196    use std::time::Duration;
197
198    use futures::StreamExt;
199    use futures::stream;
200    use parking_lot::Mutex;
201
202    use super::*;
203
204    #[test]
205    fn test_worker_thread() {
206        let runtime = CurrentThreadRuntime::new();
207
208        // We spawn a future that sets a value on a separate thread.
209        let value = Arc::new(AtomicUsize::new(0));
210        let value2 = Arc::clone(&value);
211        runtime
212            .handle()
213            .spawn(async move {
214                value2.store(42, Ordering::SeqCst);
215            })
216            .detach();
217
218        // By default, nothing has driven the executor, so the value should still be 0.
219        assert_eq!(value.load(Ordering::SeqCst), 0);
220
221        // An empty pool still does nothing.
222        let pool = runtime.new_pool();
223        assert_eq!(value.load(Ordering::SeqCst), 0);
224
225        // Adding a worker thread should drive the executor.
226        pool.set_workers(1);
227        for _ in 0..10 {
228            if value.load(Ordering::SeqCst) == 42 {
229                break;
230            }
231            thread::sleep(Duration::from_millis(10));
232        }
233        assert_eq!(value.load(Ordering::SeqCst), 42);
234    }
235
236    #[test]
237    fn test_block_on_stream_single_thread() {
238        let mut iter =
239            CurrentThreadRuntime::new().block_on_stream(stream::iter(vec![1, 2, 3, 4, 5]).boxed());
240
241        assert_eq!(iter.next(), Some(1));
242        assert_eq!(iter.next(), Some(2));
243        assert_eq!(iter.next(), Some(3));
244        assert_eq!(iter.next(), Some(4));
245        assert_eq!(iter.next(), Some(5));
246        assert_eq!(iter.next(), None);
247    }
248
249    #[test]
250    fn test_block_on_stream_multiple_threads() {
251        let counter = Arc::new(AtomicUsize::new(0));
252        let num_threads = 4;
253        let items_per_thread = 25;
254        let total_items = 100;
255
256        let iter = CurrentThreadRuntime::new()
257            .block_on_stream_thread_safe(|_h| stream::iter(0..total_items).boxed());
258
259        let barrier = Arc::new(Barrier::new(num_threads));
260        let results = Arc::new(Mutex::new(Vec::new()));
261
262        let threads: Vec<_> = (0..num_threads)
263            .map(|_| {
264                let mut iter = iter.clone();
265                let counter = Arc::clone(&counter);
266                let barrier = Arc::clone(&barrier);
267                let results = Arc::clone(&results);
268
269                thread::spawn(move || {
270                    barrier.wait();
271                    let mut local_results = Vec::new();
272
273                    for _ in 0..items_per_thread {
274                        if let Some(item) = iter.next() {
275                            counter.fetch_add(1, Ordering::SeqCst);
276                            local_results.push(item);
277                        }
278                    }
279
280                    results.lock().push(local_results);
281                })
282            })
283            .collect();
284
285        for thread in threads {
286            thread.join().unwrap();
287        }
288
289        assert_eq!(counter.load(Ordering::SeqCst), total_items);
290
291        let all_results = results.lock();
292        let mut collected: Vec<_> = all_results.iter().flatten().copied().collect();
293        collected.sort();
294        assert_eq!(collected, (0..total_items).collect::<Vec<_>>());
295    }
296
297    #[test]
298    fn test_block_on_stream_thread_safe_propagates_driver_panic() {
299        let runtime = CurrentThreadRuntime::new();
300        let mut iter = runtime.block_on_stream_thread_safe(|_h| {
301            stream::poll_fn(|_| -> Poll<Option<usize>> {
302                panic!("stream driver panic");
303            })
304            .boxed()
305        });
306
307        let panic = std::panic::catch_unwind(AssertUnwindSafe(|| iter.next()))
308            .expect_err("stream panic must propagate through iterator");
309        let message = panic
310            .downcast_ref::<&'static str>()
311            .copied()
312            .or_else(|| panic.downcast_ref::<String>().map(String::as_str))
313            .unwrap_or("<unknown panic>");
314        assert!(message.contains("stream driver panic"));
315    }
316
317    fn panic_message(panic: &(dyn Any + Send)) -> &str {
318        panic
319            .downcast_ref::<&'static str>()
320            .copied()
321            .or_else(|| panic.downcast_ref::<String>().map(String::as_str))
322            .unwrap_or("<unknown panic>")
323    }
324
325    // A driver panic must propagate on *every* run, regardless of executor scheduling. Running
326    // the scenario many times guards against a return to timing-dependent propagation.
327    #[test]
328    fn test_block_on_stream_thread_safe_panic_propagation_is_deterministic() {
329        for i in 0..2000 {
330            let mut iter = CurrentThreadRuntime::new().block_on_stream_thread_safe(|_h| {
331                stream::poll_fn(|_| -> Poll<Option<usize>> {
332                    panic!("deterministic driver panic");
333                })
334                .boxed()
335            });
336
337            let outcome = std::panic::catch_unwind(AssertUnwindSafe(|| iter.next()));
338            assert!(
339                outcome.is_err(),
340                "driver panic was swallowed on iteration {i}: next() returned {:?}",
341                outcome.ok().flatten(),
342            );
343        }
344    }
345
346    // A panic after some items were already produced must still surface, not be seen as a clean
347    // end of stream.
348    #[test]
349    fn test_block_on_stream_thread_safe_panic_after_items() {
350        let mut emitted = 0usize;
351        let iter = CurrentThreadRuntime::new().block_on_stream_thread_safe(move |_h| {
352            stream::poll_fn(move |_| -> Poll<Option<usize>> {
353                if emitted < 3 {
354                    emitted += 1;
355                    Poll::Ready(Some(emitted))
356                } else {
357                    panic!("driver panic after items");
358                }
359            })
360            .boxed()
361        });
362
363        // Drain the iterator. The terminal event must be a propagated panic, never a clean
364        // `None`. This avoids depending on exactly how many buffered items survive channel close.
365        let outcome = std::panic::catch_unwind(AssertUnwindSafe(move || iter.collect::<Vec<_>>()));
366        match outcome {
367            Ok(items) => panic!("driver panic was swallowed; stream ended cleanly with {items:?}"),
368            Err(panic) => assert!(panic_message(&*panic).contains("driver panic after items")),
369        }
370    }
371
372    // With multiple consumers, a driver panic must reach *every* consumer that observes the end
373    // of the stream; it must never be swallowed by all of them. Exactly one consumer joins the
374    // driver and observes the panic; the rest see the stream end.
375    #[test]
376    fn test_block_on_stream_thread_safe_multi_consumer_panic_surfaced() {
377        let iter = CurrentThreadRuntime::new().block_on_stream_thread_safe(|_h| {
378            stream::poll_fn(|_| -> Poll<Option<usize>> {
379                panic!("multi consumer driver panic");
380            })
381            .boxed()
382        });
383
384        let num_threads = 4;
385        let barrier = Arc::new(Barrier::new(num_threads));
386        let panics = Arc::new(AtomicUsize::new(0));
387
388        let handles: Vec<_> = (0..num_threads)
389            .map(|_| {
390                let mut iter = iter.clone();
391                let barrier = Arc::clone(&barrier);
392                let panics = Arc::clone(&panics);
393                thread::spawn(move || {
394                    barrier.wait();
395                    match std::panic::catch_unwind(AssertUnwindSafe(|| iter.next())) {
396                        // The driver panicked before producing anything, so a clean end is the
397                        // only non-panic outcome a consumer may observe.
398                        Ok(None) => {}
399                        Ok(Some(_)) => panic!("no item was produced before the driver panicked"),
400                        Err(panic) => {
401                            assert!(panic_message(&*panic).contains("multi consumer driver panic"));
402                            panics.fetch_add(1, Ordering::SeqCst);
403                        }
404                    }
405                })
406            })
407            .collect();
408
409        for handle in handles {
410            handle.join().expect("consumer thread panicked uncaught");
411        }
412
413        // The panic surfaces to exactly one consumer and is never swallowed by all of them.
414        assert_eq!(panics.load(Ordering::SeqCst), 1);
415    }
416
417    // Clean completion must return `None` with no spurious panic.
418    #[test]
419    fn test_block_on_stream_thread_safe_clean_completion_returns_none() {
420        let mut iter = CurrentThreadRuntime::new()
421            .block_on_stream_thread_safe(|_h| stream::iter(vec![1usize, 2, 3]).boxed());
422
423        assert_eq!(iter.next(), Some(1));
424        assert_eq!(iter.next(), Some(2));
425        assert_eq!(iter.next(), Some(3));
426        assert_eq!(iter.next(), None);
427        assert_eq!(iter.next(), None);
428    }
429
430    #[test]
431    fn test_block_on_stream_concurrent_clone_and_drive() {
432        let num_items = 50;
433        let num_threads = 3;
434
435        let iter = CurrentThreadRuntime::new().block_on_stream_thread_safe(|h| {
436            stream::unfold(0, move |state| {
437                let h = h.clone();
438                async move {
439                    if state < num_items {
440                        h.spawn_cpu(move || {
441                            thread::sleep(Duration::from_micros(10));
442                            state
443                        })
444                        .await;
445                        Some((state, state + 1))
446                    } else {
447                        None
448                    }
449                }
450            })
451        });
452
453        let collected = Arc::new(Mutex::new(Vec::new()));
454        let barrier = Arc::new(Barrier::new(num_threads));
455
456        let threads: Vec<_> = (0..num_threads)
457            .map(|thread_id| {
458                let iter = iter.clone();
459                let collected = Arc::clone(&collected);
460                let barrier = Arc::clone(&barrier);
461
462                thread::spawn(move || {
463                    barrier.wait();
464                    let mut local_items = Vec::new();
465
466                    for item in iter {
467                        local_items.push((thread_id, item));
468                        if local_items.len() >= 5 {
469                            break;
470                        }
471                    }
472
473                    collected.lock().extend(local_items);
474                })
475            })
476            .collect();
477
478        for thread in threads {
479            thread.join().unwrap();
480        }
481
482        let results = collected.lock();
483        let mut values: Vec<_> = results.iter().map(|(_, v)| *v).collect();
484        values.sort();
485        values.dedup();
486
487        assert!(values.len() >= 5);
488        assert!(values.iter().all(|&v| v < num_items));
489    }
490
491    #[test]
492    fn test_block_on_stream_async_work() {
493        let runtime = CurrentThreadRuntime::new();
494        let handle = runtime.handle();
495        let iter = runtime.block_on_stream({
496            stream::unfold((handle, 0), |(h, state)| async move {
497                if state < 10 {
498                    let value = h
499                        .spawn(async move { futures::future::ready(state * 2).await })
500                        .await;
501                    Some((value, (h, state + 1)))
502                } else {
503                    None
504                }
505            })
506        });
507
508        let results: Vec<_> = iter.collect();
509        assert_eq!(results, vec![0, 2, 4, 6, 8, 10, 12, 14, 16, 18]);
510    }
511
512    #[test]
513    fn test_block_on_stream_drop_receivers_early() {
514        let counter = Arc::new(AtomicUsize::new(0));
515        let c = Arc::clone(&counter);
516
517        let mut iter = CurrentThreadRuntime::new().block_on_stream({
518            stream::unfold(0, move |state| {
519                let c = Arc::clone(&c);
520                async move {
521                    (state < 100).then(|| {
522                        c.fetch_add(1, Ordering::SeqCst);
523                        (state, state + 1)
524                    })
525                }
526            })
527            .boxed()
528        });
529
530        assert_eq!(iter.next(), Some(0));
531        assert_eq!(iter.next(), Some(1));
532        assert_eq!(iter.next(), Some(2));
533
534        drop(iter);
535
536        let final_count = counter.load(Ordering::SeqCst);
537        assert!(
538            final_count < 100,
539            "Stream should stop when all receivers are dropped"
540        );
541    }
542
543    #[test]
544    fn test_block_on_stream_interleaved_access() {
545        let barrier = Arc::new(Barrier::new(2));
546        let iter = CurrentThreadRuntime::new()
547            .block_on_stream_thread_safe(|_h| stream::iter(0..20).boxed());
548
549        let iter1 = iter.clone();
550        let iter2 = iter;
551        let barrier1 = Arc::clone(&barrier);
552        let barrier2 = barrier;
553
554        let thread1 = thread::spawn(move || {
555            let mut iter = iter1;
556            let mut results = Vec::new();
557            barrier1.wait();
558
559            for _ in 0..5 {
560                if let Some(val) = iter.next() {
561                    results.push(val);
562                    thread::sleep(Duration::from_micros(50));
563                }
564            }
565            results
566        });
567
568        let thread2 = thread::spawn(move || {
569            let mut iter = iter2;
570            let mut results = Vec::new();
571            barrier2.wait();
572
573            for _ in 0..5 {
574                if let Some(val) = iter.next() {
575                    results.push(val);
576                    thread::sleep(Duration::from_micros(50));
577                }
578            }
579            results
580        });
581
582        let results1 = thread1.join().unwrap();
583        let results2 = thread2.join().unwrap();
584
585        let mut all_results = results1;
586        all_results.extend(results2);
587        all_results.sort();
588
589        assert_eq!(all_results, (0..10).collect::<Vec<_>>());
590
591        for i in 0..10 {
592            assert_eq!(all_results.iter().filter(|&&x| x == i).count(), 1);
593        }
594    }
595
596    #[test]
597    fn test_block_on_stream_stress_test() {
598        let num_threads = 10;
599        let num_items = 1000;
600
601        let iter = CurrentThreadRuntime::new()
602            .block_on_stream_thread_safe(|_h| stream::iter(0..num_items).boxed());
603
604        let received = Arc::new(Mutex::new(Vec::new()));
605        let barrier = Arc::new(Barrier::new(num_threads));
606
607        let threads: Vec<_> = (0..num_threads)
608            .map(|_| {
609                let iter = iter.clone();
610                let received = Arc::clone(&received);
611                let barrier = Arc::clone(&barrier);
612
613                thread::spawn(move || {
614                    barrier.wait();
615                    for val in iter {
616                        received.lock().push(val);
617                    }
618                })
619            })
620            .collect();
621
622        for thread in threads {
623            thread.join().unwrap();
624        }
625
626        let mut results = received.lock().clone();
627        results.sort();
628
629        assert_eq!(results.len(), num_items);
630        assert_eq!(results, (0..num_items).collect::<Vec<_>>());
631    }
632}