1use 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#[derive(Clone, Default)]
33pub struct CurrentThreadRuntime {
34 executor: Arc<smol::Executor<'static>>,
35}
36
37impl CurrentThreadRuntime {
38 pub fn new() -> Self {
40 Self::default()
41 }
42
43 pub fn new_pool(&self) -> CurrentThreadWorkerPool {
51 CurrentThreadWorkerPool::new(Arc::clone(&self.executor))
52 }
53
54 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 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 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
121pub 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
135pub struct ThreadSafeIterator<T> {
137 executor: Arc<smol::Executor<'static>>,
138 results: kanal::AsyncReceiver<T>,
139 driver: Arc<Mutex<Option<smol::Task<()>>>>,
143}
144
145impl<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 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 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)] #[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 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 assert_eq!(value.load(Ordering::SeqCst), 0);
220
221 let pool = runtime.new_pool();
223 assert_eq!(value.load(Ordering::SeqCst), 0);
224
225 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 #[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 #[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 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 #[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 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 assert_eq!(panics.load(Ordering::SeqCst), 1);
415 }
416
417 #[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}