1use crate::{
4 Runtime,
5 sys::AsSysFd,
6 traits::{Executor, Reactor, RuntimeKit},
7 util::Task,
8};
9use async_compat::{Compat, CompatExt};
10use futures_core::Stream;
11use futures_io::{AsyncRead, AsyncWrite};
12use std::{
13 future::Future,
14 io::{self, Read, Write},
15 net::SocketAddr,
16 pin::Pin,
17 sync::Arc,
18 task::{Context, Poll},
19 time::{Duration, Instant},
20};
21use tokio::{
22 net::TcpStream,
23 runtime::{EnterGuard, Handle, Runtime as TokioRT},
24 time::Sleep,
25};
26use tokio_stream::{StreamExt, wrappers::IntervalStream};
27
28use task::TTask;
29
30pub type TokioRuntime = Runtime<Tokio>;
32
33impl TokioRuntime {
34 pub fn tokio() -> io::Result<Self> {
36 Ok(Self::tokio_with_runtime(TokioRT::new()?))
37 }
38
39 #[must_use]
41 pub fn tokio_current() -> Self {
42 Self::new(Tokio::current())
43 }
44
45 #[must_use]
47 pub fn tokio_with_handle(handle: Handle) -> Self {
48 Self::new(Tokio::default().with_handle(handle))
49 }
50
51 #[must_use]
53 pub fn tokio_with_runtime(runtime: TokioRT) -> Self {
54 Self::new(Tokio::default().with_runtime(runtime))
55 }
56
57 pub fn shutdown_blocking(self) -> Result<(), Self> {
73 self.kit.shutdown_blocking().map_err(Self::new)
74 }
75}
76
77const NO_RUNTIME: &str = "no tokio runtime: use Runtime::tokio() or Runtime::tokio_with_handle()";
83
84#[derive(Default, Clone, Debug)]
86pub struct Tokio {
87 handle: Option<Handle>,
88 runtime: Option<Arc<OwnedRuntime>>,
89}
90
91#[derive(Debug)]
98struct OwnedRuntime(Option<TokioRT>);
99
100impl OwnedRuntime {
101 fn get(&self) -> &TokioRT {
102 self.0
103 .as_ref()
104 .expect("owned runtime is available until drop")
105 }
106}
107
108impl Drop for OwnedRuntime {
109 fn drop(&mut self) {
110 if let Some(runtime) = self.0.take() {
111 if Handle::try_current().is_ok() {
112 runtime.shutdown_background();
113 } else {
114 drop(runtime);
115 }
116 }
117 }
118}
119
120impl Tokio {
121 pub(crate) fn shutdown_blocking(mut self) -> Result<(), Self> {
122 if self.runtime.is_some() && tokio::task::try_id().is_some() {
125 return Err(self);
126 }
127 let Some(runtime) = self.runtime.take() else {
128 return Ok(());
129 };
130 match Arc::try_unwrap(runtime) {
131 Ok(mut runtime) => {
132 drop(runtime.0.take());
133 Ok(())
134 }
135 Err(runtime) => {
136 self.runtime = Some(runtime);
137 Err(self)
138 }
139 }
140 }
141
142 #[must_use]
149 pub fn with_handle(mut self, handle: Handle) -> Self {
150 self.handle = Some(handle);
151 self
152 }
153
154 #[must_use]
156 pub fn with_runtime(mut self, runtime: TokioRT) -> Self {
157 let handle = runtime.handle().clone();
158 self.runtime = Some(Arc::new(OwnedRuntime(Some(runtime))));
159 self.with_handle(handle)
160 }
161
162 #[must_use]
164 pub fn current() -> Self {
165 Self::default().with_handle(Handle::current())
166 }
167
168 fn bound_handle(&self) -> Option<&Handle> {
174 self.runtime
175 .as_ref()
176 .map(|r| r.get().handle())
177 .or(self.handle.as_ref())
178 }
179
180 fn handle(&self) -> Option<Handle> {
181 self.bound_handle()
182 .cloned()
183 .or_else(|| Handle::try_current().ok())
184 }
185
186 fn enter(&self) -> Option<EnterGuard<'_>> {
192 self.bound_handle().map(Handle::enter)
193 }
194
195 fn has_runtime(&self) -> bool {
197 self.bound_handle().is_some() || Handle::try_current().is_ok()
198 }
199
200 fn require_enter(&self) -> Option<EnterGuard<'_>> {
205 assert!(self.has_runtime(), "{NO_RUNTIME}");
206 self.enter()
207 }
208
209 fn require_handle(&self) -> Handle {
210 self.handle().expect(NO_RUNTIME)
211 }
212}
213
214impl RuntimeKit for Tokio {}
215
216impl Executor for Tokio {
217 type Task<T: Send + 'static> = TTask<T>;
218
219 fn block_on<T, F: Future<Output = T>>(&self, f: F) -> T {
220 if let Some(runtime) = self.runtime.as_ref() {
221 runtime.get().block_on(f)
222 } else {
223 self.require_handle().block_on(f)
226 }
227 }
228
229 fn spawn<T: Send + 'static, F: Future<Output = T> + Send + 'static>(
230 &self,
231 f: F,
232 ) -> Task<Self::Task<T>> {
233 TTask(Some(self.require_handle().spawn(f))).into()
234 }
235
236 fn spawn_blocking<T: Send + 'static, F: FnOnce() -> T + Send + 'static>(
237 &self,
238 f: F,
239 ) -> Task<Self::Task<T>> {
240 TTask(Some(self.require_handle().spawn_blocking(f))).into()
241 }
242}
243
244impl Reactor for Tokio {
245 type TcpStream = Compat<TcpStream>;
246 type Sleep = Sleep;
247
248 fn register<H: Read + Write + AsSysFd + Send + 'static>(
249 &self,
250 socket: H,
251 ) -> io::Result<impl AsyncRead + AsyncWrite + Send + Unpin + 'static> {
252 if !self.has_runtime() {
255 return Err(io::Error::other(NO_RUNTIME));
256 }
257 let _enter = self.enter();
258 #[cfg(unix)]
259 {
260 Ok(unix::AsyncFdWrapper(tokio::io::unix::AsyncFd::new(socket)?))
261 }
262 #[cfg(not(unix))]
263 {
264 let _ = socket;
265 Err::<crate::util::DummyIO, _>(io::Error::other(
266 "Registering FD on tokio reactor is only supported on unix",
267 ))
268 }
269 }
270
271 fn sleep(&self, dur: Duration) -> Self::Sleep {
272 let _enter = self.require_enter();
273 tokio::time::sleep(dur)
274 }
275
276 fn interval(&self, dur: Duration) -> impl Stream<Item = Instant> + Send + 'static {
277 let _enter = self.require_enter();
278 IntervalStream::new(tokio::time::interval(dur)).map(tokio::time::Instant::into_std)
279 }
280
281 fn tcp_connect_addr(
282 &self,
283 addr: SocketAddr,
284 ) -> impl Future<Output = io::Result<Self::TcpStream>> + Send + 'static {
285 InTokioContext::new(self.bound_handle().cloned(), async move {
295 if !crate::util::inside_tokio() {
300 return Err(io::Error::other(NO_RUNTIME));
301 }
302 let stream = TcpStream::connect(addr).await?;
303 stream.set_nodelay(true)?;
304 Ok(stream.compat())
305 })
306 }
307}
308
309struct InTokioContext<F: Future> {
318 handle: Option<Handle>,
319 fut: Pin<Box<F>>,
321}
322
323impl<F: Future> InTokioContext<F> {
324 fn new(handle: Option<Handle>, fut: F) -> Self {
325 Self {
326 handle,
327 fut: Box::pin(fut),
328 }
329 }
330}
331
332impl<F: Future> Future for InTokioContext<F> {
333 type Output = F::Output;
334
335 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
336 let this = self.get_mut();
337 let _enter = this.handle.as_ref().map(Handle::enter);
338 this.fut.as_mut().poll(cx)
339 }
340}
341
342mod task {
343 use crate::util::TaskImpl;
344 use async_trait::async_trait;
345 use std::{
346 future::Future,
347 panic,
348 pin::Pin,
349 task::{Context, Poll},
350 };
351
352 #[derive(Debug)]
354 pub struct TTask<T: Send + 'static>(pub(super) Option<tokio::task::JoinHandle<T>>);
355
356 #[async_trait]
357 impl<T: Send + 'static> TaskImpl for TTask<T> {
358 async fn cancel(&mut self) -> Option<T> {
359 let task = self.0.take()?;
360 task.abort();
361 task.await.ok()
362 }
363 }
364
365 impl<T: Send + 'static> Future for TTask<T> {
366 type Output = T;
367
368 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
369 let task = self
370 .0
371 .as_mut()
372 .expect("Task polled after it was canceled or completed");
373 let res = match Pin::new(task).poll(cx) {
374 Poll::Pending => return Poll::Pending,
375 Poll::Ready(res) => res,
376 };
377
378 self.0 = None;
381
382 match res {
383 Ok(res) => Poll::Ready(res),
384 Err(err) if err.is_panic() => panic::resume_unwind(err.into_panic()),
388 Err(err) => panic!("Task did not complete: {err}"),
389 }
390 }
391 }
392}
393
394#[cfg(unix)]
395mod unix {
396 use super::*;
397 use futures_io::{AsyncRead, AsyncWrite};
398 use std::{
399 io::{IoSlice, IoSliceMut},
400 pin::Pin,
401 task::{Context, Poll},
402 };
403 use tokio::io::unix::AsyncFd;
404
405 pub(super) struct AsyncFdWrapper<H: Read + Write + AsSysFd>(pub(super) AsyncFd<H>);
406
407 impl<H: Read + Write + AsSysFd> AsyncFdWrapper<H> {
408 fn read<F: FnOnce(&mut AsyncFd<H>) -> io::Result<usize>>(
409 mut self: Pin<&mut Self>,
410 cx: &mut Context<'_>,
411 f: F,
412 ) -> Option<Poll<io::Result<usize>>> {
413 Some(match self.0.poll_read_ready_mut(cx) {
414 Poll::Pending => Poll::Pending,
415 Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
416 Poll::Ready(Ok(mut guard)) => match guard.try_io(f) {
417 Ok(res) => Poll::Ready(res),
418 Err(_) => return None,
419 },
420 })
421 }
422
423 fn write<R, F: FnOnce(&mut AsyncFd<H>) -> io::Result<R>>(
424 mut self: Pin<&mut Self>,
425 cx: &mut Context<'_>,
426 f: F,
427 ) -> Option<Poll<io::Result<R>>> {
428 Some(match self.0.poll_write_ready_mut(cx) {
429 Poll::Pending => Poll::Pending,
430 Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
431 Poll::Ready(Ok(mut guard)) => match guard.try_io(f) {
432 Ok(res) => Poll::Ready(res),
433 Err(_) => return None,
434 },
435 })
436 }
437 }
438
439 impl<H: Read + Write + AsSysFd> Unpin for AsyncFdWrapper<H> {}
440
441 impl<H: Read + Write + AsSysFd> AsyncRead for AsyncFdWrapper<H> {
442 fn poll_read(
443 mut self: Pin<&mut Self>,
444 cx: &mut Context<'_>,
445 buf: &mut [u8],
446 ) -> Poll<io::Result<usize>> {
447 loop {
448 if let Some(res) = self.as_mut().read(cx, |socket| socket.get_mut().read(buf)) {
449 return res;
450 }
451 }
452 }
453
454 fn poll_read_vectored(
455 mut self: Pin<&mut Self>,
456 cx: &mut Context<'_>,
457 bufs: &mut [IoSliceMut<'_>],
458 ) -> Poll<io::Result<usize>> {
459 loop {
460 if let Some(res) = self
461 .as_mut()
462 .read(cx, |socket| socket.get_mut().read_vectored(bufs))
463 {
464 return res;
465 }
466 }
467 }
468 }
469
470 impl<H: Read + Write + AsSysFd> AsyncWrite for AsyncFdWrapper<H> {
471 fn poll_write(
472 mut self: Pin<&mut Self>,
473 cx: &mut Context<'_>,
474 buf: &[u8],
475 ) -> Poll<io::Result<usize>> {
476 loop {
477 if let Some(res) = self
478 .as_mut()
479 .write(cx, |socket| socket.get_mut().write(buf))
480 {
481 return res;
482 }
483 }
484 }
485
486 fn poll_write_vectored(
487 mut self: Pin<&mut Self>,
488 cx: &mut Context<'_>,
489 bufs: &[IoSlice<'_>],
490 ) -> Poll<io::Result<usize>> {
491 loop {
492 if let Some(res) = self
493 .as_mut()
494 .write(cx, |socket| socket.get_mut().write_vectored(bufs))
495 {
496 return res;
497 }
498 }
499 }
500
501 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
502 loop {
503 if let Some(res) = self.as_mut().write(cx, |socket| socket.get_mut().flush()) {
504 return res;
505 }
506 }
507 }
508
509 fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<futures_io::Result<()>> {
510 self.poll_flush(cx)
511 }
512 }
513}
514
515#[cfg(test)]
516mod tests {
517 use super::*;
518
519 #[test]
520 fn auto_traits() {
521 use crate::util::test::*;
522 let runtime = Runtime::tokio().unwrap();
523 assert_send(&runtime);
524 assert_sync(&runtime);
525 assert_clone(&runtime);
526 }
527
528 #[test]
531 fn panicking_task_does_not_hang() {
532 let res = crate::util::test::with_timeout(|| {
533 let runtime = Runtime::tokio().unwrap();
534 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
535 runtime.block_on(runtime.spawn(async { panic!("boom") }))
536 }))
537 });
538 assert_eq!(
541 res.expect_err("task panic").downcast_ref::<&str>(),
542 Some(&"boom")
543 );
544 }
545
546 #[test]
547 fn last_owned_runtime_can_be_dropped_from_its_worker() {
548 let (release_tx, release_rx) = std::sync::mpsc::channel();
549 let (done_tx, done_rx) = std::sync::mpsc::channel();
550 let runtime = Runtime::tokio().unwrap();
551 let last_owner = runtime.clone();
552
553 drop(runtime.spawn(async move {
556 release_rx.recv().unwrap();
557 let dropped = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
558 drop(last_owner);
559 }));
560 done_tx.send(dropped.is_ok()).unwrap();
561 }));
562 drop(runtime);
563 release_tx.send(()).unwrap();
564 assert!(done_rx.recv_timeout(Duration::from_secs(10)).unwrap());
565 }
566
567 #[test]
568 fn explicit_shutdown_waits_with_a_foreign_handle_entered() {
569 let runtime = Runtime::tokio().unwrap();
570 let (started_tx, started_rx) = std::sync::mpsc::channel();
571 let (release_tx, release_rx) = std::sync::mpsc::channel();
572 let (done_tx, done_rx) = std::sync::mpsc::channel();
573
574 drop(runtime.spawn_blocking(move || {
575 started_tx.send(()).unwrap();
576 release_rx.recv().unwrap();
577 done_tx.send(()).unwrap();
578 }));
579 started_rx.recv_timeout(Duration::from_secs(5)).unwrap();
580
581 let (shutdown_started_tx, shutdown_started_rx) = std::sync::mpsc::channel();
582 let (shutdown_done_tx, shutdown_done_rx) = std::sync::mpsc::channel();
583 let shutdown_thread = std::thread::spawn(move || {
584 let ambient = TokioRT::new().unwrap();
585 let _enter = ambient.enter();
586 shutdown_started_tx.send(()).unwrap();
587 runtime.shutdown_blocking().unwrap();
588 shutdown_done_tx.send(()).unwrap();
589 });
590 shutdown_started_rx
591 .recv_timeout(Duration::from_secs(5))
592 .unwrap();
593 let premature = shutdown_done_rx.recv_timeout(Duration::from_millis(100));
594 release_tx.send(()).unwrap();
595 assert!(matches!(
596 premature,
597 Err(std::sync::mpsc::RecvTimeoutError::Timeout)
598 ));
599 shutdown_done_rx
600 .recv_timeout(Duration::from_secs(5))
601 .unwrap();
602 shutdown_thread.join().unwrap();
603 assert!(done_rx.try_recv().is_ok());
604 }
605
606 #[test]
607 fn explicit_shutdown_refuses_its_own_blocking_task() {
608 let runtime = Runtime::tokio().unwrap();
609 let last_owner = runtime.clone();
610 let (release_tx, release_rx) = std::sync::mpsc::channel();
611 let (result_tx, result_rx) = std::sync::mpsc::channel();
612
613 drop(runtime.spawn_blocking(move || {
614 release_rx.recv().unwrap();
615 result_tx.send(last_owner.shutdown_blocking()).unwrap();
616 }));
617 drop(runtime);
618 release_tx.send(()).unwrap();
619 let runtime = result_rx
620 .recv_timeout(Duration::from_secs(5))
621 .unwrap()
622 .unwrap_err();
623 runtime.shutdown_blocking().unwrap();
624 }
625
626 #[test]
627 fn explicit_shutdown_requires_last_owner() {
628 let runtime = Runtime::tokio().unwrap();
629 let other_owner = runtime.clone();
630 let runtime = runtime.shutdown_blocking().unwrap_err();
631 drop(other_owner);
632 runtime.shutdown_blocking().unwrap();
633 }
634
635 #[test]
638 fn tcp_connect_addr_polled_off_runtime() {
639 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
640 let addr = listener.local_addr().unwrap();
641
642 let (_runtime, mut stream) = crate::util::test::with_timeout(move || {
644 let runtime = Runtime::tokio().unwrap();
645 let connect = runtime.tcp_connect_addr(addr);
646 let stream = crate::util::simple_block_on(connect).expect("connect");
647 (runtime, stream)
648 });
649
650 let (mut socket, _) = listener.accept().expect("accept");
654 Write::write_all(&mut socket, b"hello").expect("write");
655
656 let read = crate::util::test::with_timeout(move || {
660 let mut buf = [0_u8; 5];
661 let mut read = 0;
662 crate::util::simple_block_on(std::future::poll_fn(|cx| {
663 while read < buf.len() {
664 match Pin::new(&mut stream).poll_read(cx, &mut buf[read..]) {
665 Poll::Ready(Ok(0)) => break,
666 Poll::Ready(Ok(n)) => read += n,
667 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
668 Poll::Pending => return Poll::Pending,
669 }
670 }
671 Poll::Ready(Ok(buf))
672 }))
673 .expect("read")
674 });
675 assert_eq!(&read, b"hello");
676 }
677
678 #[test]
681 fn one_kit_binds_everything_to_the_same_runtime() {
682 let other = TokioRT::new().unwrap();
683 let runtime = Runtime::new(
684 Tokio::default()
685 .with_runtime(TokioRT::new().unwrap())
686 .with_handle(other.handle().clone()),
687 );
688 drop(other);
690
691 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
692 let addr = listener.local_addr().unwrap();
693 let accepted = std::thread::spawn(move || listener.accept().map(|_| ()));
694 runtime.block_on(async { runtime.tcp_connect_addr(addr).await.expect("connect") });
695 accepted.join().expect("accept thread").expect("accept");
696 }
697
698 #[test]
701 fn tcp_connect_addr_without_a_runtime_reports_an_error() {
702 let runtime = Runtime::new(Tokio::default());
703 let addr = "127.0.0.1:1".parse().unwrap();
704 let Err(err) = crate::util::simple_block_on(runtime.tcp_connect_addr(addr)) else {
705 panic!("connect succeeded without a runtime");
706 };
707 assert!(err.to_string().contains("no tokio runtime"), "{err}");
708 }
709
710 #[test]
715 fn tcp_connect_addr_built_off_runtime_uses_the_one_polling_it() {
716 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
717 let addr = listener.local_addr().unwrap();
718 let accepted = std::thread::spawn(move || listener.accept().map(|_| ()));
719
720 let connect = Runtime::new(Tokio::default()).tcp_connect_addr(addr);
722 TokioRT::new()
723 .unwrap()
724 .block_on(connect)
725 .expect("connect polled inside a runtime");
726 accepted.join().expect("accept thread").expect("accept");
727 }
728
729 #[test]
732 #[cfg(unix)]
733 fn register_without_a_runtime_reports_an_error() {
734 let runtime = Runtime::new(Tokio::default());
735 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
736 let socket = std::net::TcpStream::connect(listener.local_addr().unwrap()).unwrap();
737 let Err(err) = runtime.register(socket) else {
738 panic!("register succeeded without a runtime");
739 };
740 assert!(err.to_string().contains("no tokio runtime"), "{err}");
741 }
742
743 #[test]
744 fn panicking_blocking_task_does_not_hang() {
745 let res = crate::util::test::with_timeout(|| {
746 let runtime = Runtime::tokio().unwrap();
747 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
748 runtime.block_on(runtime.spawn_blocking(|| -> u32 { panic!("boom") }))
749 }))
750 });
751 assert_eq!(
752 res.expect_err("task panic").downcast_ref::<&str>(),
753 Some(&"boom")
754 );
755 }
756}