1use std::{fmt, pin::Pin, task::Context, task::Poll, task::ready};
2
3use async_task::Task;
4
5#[derive(Debug)]
6pub struct JoinHandle<T> {
8 task: Option<Task<T>>,
9}
10
11impl<T> JoinHandle<T> {
12 pub(crate) fn new(task: Task<T>) -> Self {
13 JoinHandle { task: Some(task) }
14 }
15
16 pub fn cancel(mut self) {
18 if let Some(t) = self.task.take() {
19 drop(t.cancel());
20 }
21 }
22
23 pub fn detach(mut self) {
25 if let Some(t) = self.task.take() {
26 t.detach();
27 }
28 }
29
30 pub fn is_finished(&self) -> bool {
32 match &self.task {
33 Some(fut) => fut.is_finished(),
34 None => true,
35 }
36 }
37}
38
39impl<T> Drop for JoinHandle<T> {
40 fn drop(&mut self) {
41 if let Some(fut) = self.task.take() {
42 fut.detach();
43 }
44 }
45}
46
47impl<T> Future for JoinHandle<T> {
48 type Output = Result<T, JoinError>;
49
50 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
51 Poll::Ready(match self.task.as_mut() {
52 Some(fut) => Ok(ready!(Pin::new(fut).poll(cx))),
53 None => Err(JoinError),
54 })
55 }
56}
57
58#[derive(Debug, Copy, Clone)]
60pub struct JoinError;
61
62impl fmt::Display for JoinError {
63 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64 write!(f, "JoinError")
65 }
66}
67
68impl std::error::Error for JoinError {}
69
70#[cfg(all(test, not(feature = "tokio"), not(feature = "compio")))]
71mod tests {
72 use std::future::{pending, poll_fn};
73
74 use crate::{Handle, System, testing::TestRunner};
75
76 #[test]
77 fn join_handle() {
78 System::new("test", TestRunner).block_on(async {
79 let hnd = crate::spawn(pending::<()>());
80 assert!(!hnd.is_finished());
81 hnd.cancel();
82
83 assert_eq!(crate::spawn(async { 1 }).await.unwrap(), 1);
84
85 let hnd = crate::spawn(async {});
86 poll_fn(|cx| {
87 if hnd.is_finished() {
88 std::task::Poll::Ready(())
89 } else {
90 cx.waker().wake_by_ref();
91 std::task::Poll::Pending
92 }
93 })
94 .await;
95
96 let hnd = Handle::current().clone();
97 hnd.notify().unwrap();
98 });
99 assert_eq!(super::JoinError.to_string(), "JoinError");
100 assert_eq!(crate::rt_default::JoinError.to_string(), "JoinError");
101 }
102}