Skip to main content

conc_util/future/
join.rs

1use core::{
2    mem::ManuallyDrop,
3    pin::Pin,
4    task::{Context, Poll},
5};
6
7pub struct Join<A: Future, B: Future> {
8    done: bool,
9    a: Result<ManuallyDrop<A::Output>, A>,
10    b: Result<ManuallyDrop<B::Output>, B>,
11}
12
13impl<A: Future, B: Future> Join<A, B> {
14    pub fn new(a: A, b: B) -> Self {
15        Self {
16            done: false,
17            a: Err(a),
18            b: Err(b),
19        }
20    }
21}
22
23impl<A: Future, B: Future> Future for Join<A, B> {
24    type Output = (A::Output, B::Output);
25
26    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
27        if self.done {
28            panic!("Join polled after returning value")
29        }
30
31        // SAFETY: pin projecting fields, which will never move
32        let (mut a, mut b, mut done) = unsafe {
33            let ptr = self.get_unchecked_mut() as *mut Self;
34            (
35                Pin::new_unchecked(&mut (*ptr).a),
36                Pin::new_unchecked(&mut (*ptr).b),
37                Pin::new_unchecked(&mut (*ptr).done),
38            )
39        };
40
41        let a = match pin_result(a.as_mut()) {
42            Ok(result) => Some(result),
43            Err(future) => match future.poll(cx) {
44                Poll::Ready(done) => {
45                    a.set(Ok(ManuallyDrop::new(done)));
46                    // SAFETY: just set it to Ok()
47                    Some(unsafe { pin_result(a).unwrap_unchecked() })
48                }
49                Poll::Pending => None,
50            },
51        };
52        let b = match pin_result(b.as_mut()) {
53            Ok(result) => Some(result),
54            Err(future) => match future.poll(cx) {
55                Poll::Ready(done) => {
56                    b.set(Ok(ManuallyDrop::new(done)));
57                    // SAFETY: just set it to Ok()
58                    Some(unsafe { pin_result(b).unwrap_unchecked() })
59                }
60                Poll::Pending => None,
61            },
62        };
63
64        match (a, b) {
65            (Some(a), Some(b)) => {
66                done.set(true);
67                // SAFETY: we move out of pin but `a` and `b`
68                // are never exposed as pinned so it is fine
69                Poll::Ready(unsafe {
70                    (
71                        ManuallyDrop::take(a.get_unchecked_mut()),
72                        ManuallyDrop::take(b.get_unchecked_mut()),
73                    )
74                })
75            }
76            _ => Poll::Pending,
77        }
78    }
79}
80
81fn pin_result<T, E>(value: Pin<&mut Result<T, E>>) -> Result<Pin<&mut T>, Pin<&mut E>> {
82    // SAFETY: pin projecting arms, which will never move
83    unsafe {
84        match value.get_unchecked_mut() {
85            Ok(ok) => Ok(Pin::new_unchecked(ok)),
86            Err(err) => Err(Pin::new_unchecked(err)),
87        }
88    }
89}
90
91impl<A: Future, B: Future> Drop for Join<A, B> {
92    fn drop(&mut self) {
93        if self.done {
94            return;
95        }
96        if let Ok(a) = &mut self.a {
97            // SAFETY: this is only taken from by `poll` but it also sets `done`
98            unsafe { ManuallyDrop::drop(a) };
99        } else if let Ok(b) = &mut self.b {
100            // SAFETY: this is only taken from by `poll` but it also sets `done`
101            unsafe { ManuallyDrop::drop(b) };
102        }
103    }
104}