use std::ffi::{c_char, c_void, CStr};
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex, PoisonError};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
use crate::panic_safe::catch_user_panic;
struct SyncCompletionState<T> {
result: Option<Result<T, String>>,
}
struct SyncCompletionInner<T> {
consumed: AtomicBool,
state: Mutex<SyncCompletionState<T>>,
cvar: Condvar,
}
pub struct SyncCompletion<T> {
inner: Arc<SyncCompletionInner<T>>,
}
pub type SyncCompletionPtr = *mut c_void;
impl<T> SyncCompletion<T> {
#[must_use]
pub fn new() -> (Self, SyncCompletionPtr) {
let inner = Arc::new(SyncCompletionInner {
consumed: AtomicBool::new(false),
state: Mutex::new(SyncCompletionState { result: None }),
cvar: Condvar::new(),
});
let raw = Arc::into_raw(Arc::clone(&inner));
(Self { inner }, raw as SyncCompletionPtr)
}
pub fn wait(self) -> Result<T, String> {
let mut state = self
.inner
.state
.lock()
.unwrap_or_else(PoisonError::into_inner);
loop {
if let Some(result) = state.result.take() {
return result;
}
state = self
.inner
.cvar
.wait(state)
.unwrap_or_else(PoisonError::into_inner);
}
}
#[must_use]
pub fn wait_timeout(self, timeout: Duration) -> Option<Result<T, String>> {
let start = Instant::now();
let mut state = self
.inner
.state
.lock()
.unwrap_or_else(PoisonError::into_inner);
loop {
if let Some(result) = state.result.take() {
return Some(result);
}
let remaining = timeout.checked_sub(start.elapsed())?;
state = self
.inner
.cvar
.wait_timeout(state, remaining)
.unwrap_or_else(PoisonError::into_inner)
.0;
}
}
#[must_use]
pub fn context_ptr(&self) -> SyncCompletionPtr {
Arc::as_ptr(&self.inner).cast_mut().cast()
}
pub unsafe fn complete_ok(context: SyncCompletionPtr, value: T) {
Self::complete_with_result(context, Ok(value));
}
pub unsafe fn complete_err(context: SyncCompletionPtr, error: String) {
Self::complete_with_result(context, Err(error));
}
pub unsafe fn complete_with_result(context: SyncCompletionPtr, result: Result<T, String>) {
if context.is_null() {
return;
}
let inner_ref = unsafe { &*context.cast::<SyncCompletionInner<T>>() };
if inner_ref.consumed.swap(true, Ordering::AcqRel) {
eprintln!(
"doom-fish-utils: SyncCompletion callback fired more than once; \
ignoring duplicate to avoid double-free"
);
return;
}
let inner = unsafe { Arc::from_raw(context.cast::<SyncCompletionInner<T>>()) };
{
let mut state = inner
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.result = Some(result);
}
inner.cvar.notify_one();
}
}
impl<T> Default for SyncCompletion<T> {
fn default() -> Self {
Self::new().0
}
}
struct AsyncCompletionState<T> {
result: Option<Result<T, String>>,
waker: Option<Waker>,
}
struct AsyncCompletionInner<T> {
consumed: AtomicBool,
state: Mutex<AsyncCompletionState<T>>,
}
pub struct AsyncCompletion<T> {
_marker: std::marker::PhantomData<T>,
}
pub struct AsyncCompletionFuture<T> {
inner: Arc<AsyncCompletionInner<T>>,
}
impl<T> AsyncCompletion<T> {
#[must_use]
pub fn create() -> (AsyncCompletionFuture<T>, SyncCompletionPtr) {
let inner = Arc::new(AsyncCompletionInner {
consumed: AtomicBool::new(false),
state: Mutex::new(AsyncCompletionState {
result: None,
waker: None,
}),
});
let raw = Arc::into_raw(Arc::clone(&inner));
(AsyncCompletionFuture { inner }, raw as SyncCompletionPtr)
}
pub unsafe fn complete_ok(context: SyncCompletionPtr, value: T) {
Self::complete_with_result(context, Ok(value));
}
pub unsafe fn complete_err(context: SyncCompletionPtr, error: String) {
Self::complete_with_result(context, Err(error));
}
pub unsafe fn complete_with_result(context: SyncCompletionPtr, result: Result<T, String>) {
if context.is_null() {
return;
}
let inner_ref = unsafe { &*context.cast::<AsyncCompletionInner<T>>() };
if inner_ref.consumed.swap(true, Ordering::AcqRel) {
eprintln!(
"doom-fish-utils: AsyncCompletion callback fired more than once; \
ignoring duplicate to avoid double-free"
);
return;
}
let inner = unsafe { Arc::from_raw(context.cast::<AsyncCompletionInner<T>>()) };
let waker = {
let mut state = inner
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.result = Some(result);
state.waker.take()
};
if let Some(w) = waker {
w.wake();
}
}
}
impl<T> Future for AsyncCompletionFuture<T> {
type Output = Result<T, String>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut state = self
.inner
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.result.take().map_or_else(
|| {
let waker = cx.waker();
match state.waker {
Some(ref existing) if existing.will_wake(waker) => {}
_ => state.waker = Some(waker.clone()),
}
Poll::Pending
},
Poll::Ready,
)
}
}
#[must_use]
pub unsafe fn error_from_cstr(msg: *const c_char) -> String {
if msg.is_null() {
"Unknown error".to_string()
} else {
CStr::from_ptr(msg)
.to_str()
.map_or_else(|_| "Unknown error".to_string(), String::from)
}
}
pub type UnitCompletion = SyncCompletion<()>;
impl UnitCompletion {
pub unsafe extern "C" fn callback(context: *mut c_void, success: bool, msg: *const c_char) {
catch_user_panic("UnitCompletion::callback", || {
if success {
unsafe { Self::complete_ok(context, ()) };
} else {
let error = unsafe { error_from_cstr(msg) };
unsafe { Self::complete_err(context, error) };
}
});
}
}
#[cfg(test)]
mod tests {
use std::future::Future;
use std::pin::Pin;
use std::ptr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll, Wake, Waker};
use std::thread;
use std::time::{Duration, Instant};
use super::{AsyncCompletion, SyncCompletion, UnitCompletion};
#[test]
fn unit_completion_callback_matches_shared_alias() {
let callback: crate::ffi_callbacks::UnitCompletionCallback = UnitCompletion::callback;
let _ = callback;
}
#[test]
fn unit_completion_callback_completes_successfully() {
let (completion, context) = UnitCompletion::new();
unsafe { UnitCompletion::callback(context, true, ptr::null()) };
assert_eq!(completion.wait(), Ok(()));
}
#[test]
fn unit_completion_callback_reports_errors() {
let (completion, context) = UnitCompletion::new();
unsafe { UnitCompletion::callback(context, false, c"denied".as_ptr()) };
assert_eq!(completion.wait(), Err("denied".to_string()));
}
struct DropCounter(Arc<AtomicUsize>);
impl Drop for DropCounter {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn sync_completion_ignores_a_duplicate_callback() {
let (completion, context) = SyncCompletion::<u32>::new();
unsafe { SyncCompletion::<u32>::complete_ok(context, 1) };
unsafe { SyncCompletion::<u32>::complete_ok(context, 2) };
assert_eq!(completion.wait(), Ok(1));
}
#[test]
fn sync_completion_waits_for_another_thread() {
let (completion, context) = SyncCompletion::<String>::new();
let context = context as usize;
let callback = thread::spawn(move || {
thread::sleep(Duration::from_millis(20));
unsafe { SyncCompletion::<String>::complete_ok(context as *mut _, "done".to_string()) };
});
assert_eq!(completion.wait(), Ok("done".to_string()));
callback.join().unwrap();
}
#[test]
fn wait_timeout_returns_none_when_no_callback_arrives() {
let (completion, _context) = SyncCompletion::<u32>::new();
let timeout = Duration::from_millis(30);
let start = Instant::now();
assert_eq!(completion.wait_timeout(timeout), None);
assert!(start.elapsed() >= timeout);
}
#[test]
fn wait_timeout_returns_the_result() {
let (completion, context) = SyncCompletion::<u32>::new();
unsafe { SyncCompletion::<u32>::complete_err(context, "failed".to_string()) };
assert_eq!(
completion.wait_timeout(Duration::from_secs(5)),
Some(Err("failed".to_string()))
);
let (completion, context) = SyncCompletion::<u32>::new();
let context = context as usize;
let callback = thread::spawn(move || {
thread::sleep(Duration::from_millis(20));
unsafe { SyncCompletion::<u32>::complete_ok(context as *mut _, 9) };
});
assert_eq!(completion.wait_timeout(Duration::MAX), Some(Ok(9)));
callback.join().unwrap();
}
#[test]
fn late_callback_after_timeout_releases_the_value() {
let drops = Arc::new(AtomicUsize::new(0));
let (completion, context) = SyncCompletion::<DropCounter>::new();
assert!(completion.wait_timeout(Duration::from_millis(1)).is_none());
unsafe {
SyncCompletion::<DropCounter>::complete_ok(context, DropCounter(Arc::clone(&drops)));
};
assert_eq!(drops.load(Ordering::SeqCst), 1);
}
#[test]
fn default_sync_completion_can_be_completed() {
let completion = SyncCompletion::<u32>::default();
let context = completion.context_ptr();
unsafe { SyncCompletion::<u32>::complete_ok(context, 5) };
assert_eq!(completion.wait(), Ok(5));
let (completion, context) = SyncCompletion::<u32>::new();
assert_eq!(completion.context_ptr(), context);
unsafe { SyncCompletion::<u32>::complete_ok(context, 6) };
assert_eq!(completion.wait(), Ok(6));
}
#[test]
fn wait_recovers_from_a_poisoned_lock() {
let (completion, context) = SyncCompletion::<u32>::new();
let inner = Arc::clone(&completion.inner);
assert!(thread::spawn(move || {
let _state = inner.state.lock().unwrap();
panic!("poison completion state");
})
.join()
.is_err());
unsafe { SyncCompletion::<u32>::complete_ok(context, 3) };
assert_eq!(completion.wait(), Ok(3));
let (completion, _context) = SyncCompletion::<u32>::new();
let inner = Arc::clone(&completion.inner);
assert!(thread::spawn(move || {
let _state = inner.state.lock().unwrap();
panic!("poison completion state");
})
.join()
.is_err());
assert_eq!(completion.wait_timeout(Duration::from_millis(1)), None);
}
#[derive(Default)]
struct CountingWake(AtomicUsize);
impl Wake for CountingWake {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn async_completion_resolves_with_the_value() {
let (future, context) = AsyncCompletion::<u32>::create();
unsafe { AsyncCompletion::<u32>::complete_ok(context, 42) };
assert_eq!(pollster::block_on(future), Ok(42));
}
#[test]
fn async_completion_wakes_a_pending_future() {
let (mut future, context) = AsyncCompletion::<u32>::create();
let probe = Arc::new(CountingWake::default());
let waker = Waker::from(Arc::clone(&probe));
let mut cx = Context::from_waker(&waker);
assert_eq!(Pin::new(&mut future).poll(&mut cx), Poll::Pending);
let context = context as usize;
thread::spawn(move || unsafe { AsyncCompletion::<u32>::complete_ok(context as *mut _, 8) })
.join()
.unwrap();
assert_eq!(probe.0.load(Ordering::SeqCst), 1);
assert_eq!(Pin::new(&mut future).poll(&mut cx), Poll::Ready(Ok(8)));
}
#[test]
fn async_completion_resolves_with_the_error() {
let (future, context) = AsyncCompletion::<u32>::create();
unsafe { AsyncCompletion::<u32>::complete_err(context, "denied".to_string()) };
assert_eq!(pollster::block_on(future), Err("denied".to_string()));
}
#[test]
fn async_completion_ignores_a_duplicate_callback() {
let drops = Arc::new(AtomicUsize::new(0));
let (future, context) = AsyncCompletion::<DropCounter>::create();
unsafe {
AsyncCompletion::<DropCounter>::complete_ok(context, DropCounter(Arc::clone(&drops)));
};
unsafe { AsyncCompletion::<DropCounter>::complete_err(context, "duplicate".to_string()) };
let value = pollster::block_on(future).unwrap();
assert_eq!(drops.load(Ordering::SeqCst), 0);
drop(value);
assert_eq!(drops.load(Ordering::SeqCst), 1);
}
#[test]
fn dropped_async_future_releases_a_late_value() {
let drops = Arc::new(AtomicUsize::new(0));
let (future, context) = AsyncCompletion::<DropCounter>::create();
drop(future);
unsafe {
AsyncCompletion::<DropCounter>::complete_ok(context, DropCounter(Arc::clone(&drops)));
};
assert_eq!(drops.load(Ordering::SeqCst), 1);
}
}