use super::{
FutureLowerReturn, RustFutureCallback, RustFutureContinuationBoundCallback, RustFuturePoll,
Scheduler, UniffiCompatibleFuture,
};
use crate::{try_rust_call, FfiDefault, LiftArgsError, RustCallResult, RustCallStatus};
use std::{
future, panic,
pin::{pin, Pin},
sync::{Arc, Mutex},
task::{Context, Poll, RawWaker, RawWakerVTable, Waker},
};
struct WrappedFuture<FfiType> {
future: Option<Pin<Box<dyn UniffiCompatibleFuture<RustCallResult<FfiType>>>>>,
result: Option<Result<FfiType, RustCallStatus>>,
}
impl<FfiType> WrappedFuture<FfiType> {
fn new<F, T, UT>(future: F) -> Self
where
F: UniffiCompatibleFuture<Result<T, LiftArgsError>> + 'static,
T: FutureLowerReturn<UT, ReturnType = FfiType>,
{
let wrapped_future = async {
let mut future = pin!(future);
future::poll_fn(move |cx| {
let call_result = try_rust_call(
panic::AssertUnwindSafe(|| match future.as_mut().poll(cx) {
Poll::Pending => Ok(Poll::Pending),
Poll::Ready(Ok(v)) => T::lower_return(v).map(Poll::Ready),
Poll::Ready(Err(e)) => T::handle_failed_lift(e).map(Poll::Ready),
}),
);
match call_result {
Ok(Poll::Pending) => Poll::Pending,
Ok(Poll::Ready(v)) => Poll::Ready(Ok(v)),
Err(call_status) => Poll::Ready(Err(call_status)),
}
})
.await
};
Self {
future: Some(Box::pin(wrapped_future)),
result: None,
}
}
fn poll(&mut self, context: &mut Context<'_>) -> bool {
if self.result.is_some() {
true
} else if let Some(future) = &mut self.future {
match future.as_mut().poll(context) {
Poll::Pending => false,
Poll::Ready(result) => {
self.future = None;
self.result = Some(result);
true
}
}
} else {
trace!("poll with neither future nor result set");
true
}
}
fn complete(&mut self, out_status: &mut RustCallStatus) -> FfiType
where
FfiType: FfiDefault,
{
let mut return_value = FfiType::ffi_default();
match self.result.take() {
Some(Ok(v)) => return_value = v,
Some(Err(call_status)) => *out_status = call_status,
None => *out_status = RustCallStatus::cancelled(),
}
self.free();
return_value
}
fn free(&mut self) {
self.future = None;
self.result = None;
}
}
pub struct RustFuture<FfiType, Callback = RustFutureContinuationBoundCallback> {
future: Mutex<WrappedFuture<FfiType>>,
scheduler: Scheduler<Callback>,
}
impl<FfiType, Callback> RustFuture<FfiType, Callback>
where
Callback: RustFutureCallback,
{
pub fn new<F, T, UT>(future: F, _tag: UT) -> Self
where
F: UniffiCompatibleFuture<Result<T, LiftArgsError>> + 'static,
T: FutureLowerReturn<UT, ReturnType = FfiType>,
{
Self {
future: Mutex::new(WrappedFuture::new(future)),
scheduler: Scheduler::new(),
}
}
pub fn poll(self: Arc<Self>, callback: Callback) {
let cancelled = self.is_cancelled();
let ready = cancelled || {
let mut locked = self.future.lock().unwrap();
let waker = Arc::clone(&self).into_waker();
locked.poll(&mut Context::from_waker(&waker))
};
if ready {
trace!("RustFuture::poll is ready (cancelled: {cancelled})");
callback.invoke(RustFuturePoll::Ready)
} else {
self.scheduler.store(callback);
}
}
pub fn is_cancelled(&self) -> bool {
self.scheduler.is_cancelled()
}
pub fn cancel(&self) {
self.scheduler.cancel();
}
pub fn complete(&self, call_status: &mut RustCallStatus) -> FfiType
where
FfiType: FfiDefault,
{
self.future.lock().unwrap().complete(call_status)
}
pub fn free(&self) {
self.scheduler.cancel();
self.future.lock().unwrap().free();
}
}
impl<FfiType, Callback> RustFuture<FfiType, Callback>
where
Scheduler: Send + Sync,
Callback: RustFutureCallback,
{
unsafe fn waker_clone(ptr: *const ()) -> RawWaker {
trace!("RustFuture::waker_clone called ({ptr:?})");
Arc::<Self>::increment_strong_count(ptr.cast::<Self>());
RawWaker::new(
ptr,
&RawWakerVTable::new(
Self::waker_clone,
Self::waker_wake,
Self::waker_wake_by_ref,
Self::waker_drop,
),
)
}
unsafe fn waker_wake(ptr: *const ()) {
trace!("RustFuture::waker_wake called ({ptr:?})");
Self::recreate_arc(ptr).scheduler.wake();
}
unsafe fn waker_wake_by_ref(ptr: *const ()) {
trace!("RustFuture::waker_wake_by_ref called ({ptr:?})");
let ptr = ptr.cast::<Self>();
(*ptr).scheduler.wake();
}
unsafe fn waker_drop(ptr: *const ()) {
trace!("RustFuture::waker_drop called ({ptr:?})");
drop(Self::recreate_arc(ptr));
}
unsafe fn recreate_arc(ptr: *const ()) -> Arc<Self> {
let ptr = ptr.cast::<Self>();
Arc::<Self>::from_raw(ptr)
}
fn into_waker(self: Arc<Self>) -> Waker {
trace!("RustFuture::creating waker ({:?})", Arc::as_ptr(&self));
let raw_waker = RawWaker::new(
Arc::into_raw(self).cast::<()>(),
&RawWakerVTable::new(
Self::waker_clone,
Self::waker_wake,
Self::waker_wake_by_ref,
Self::waker_drop,
),
);
unsafe { Waker::from_raw(raw_waker) }
}
}