use saddle_admission::{AdmissionError, ExactStored, RpcDependencyAudit, RpcRetiredOwner, RpcStagePermit, StorageDemand, StoragePermit};
use std::{
alloc::Layout,
future::Future,
pin::Pin,
sync::{Arc, Mutex},
task::{Context, Poll, Wake, Waker},
};
type Driver = Pin<Box<dyn Future<Output = ()> + Send>>;
struct Backing {
tasks: ExactStored<Box<[Option<Driver>]>>,
audit: RpcDependencyAudit,
}
impl Drop for Backing {
fn drop(&mut self) {
for slot in self.tasks.get_mut().iter_mut() {
if let Some(driver) = slot.take() { self.audit.drop_driver(driver); }
}
}
}
struct State {
backing: Option<ExactStored<Box<Backing>>>,
wake: Option<Waker>,
failure: Option<AdmissionError>,
closed: bool,
}
#[repr(C)]
struct Shared {
state: Mutex<State>,
permit: StoragePermit,
audit: RpcDependencyAudit,
stage: RpcStagePermit,
}
impl RpcRetiredOwner for Shared {
fn destroy_last(self: Arc<Self>) {
if Arc::weak_count(&self) != 0 { std::process::abort(); }
let value = Arc::into_inner(self).unwrap_or_else(|| std::process::abort());
drop(value);
}
}
impl Shared {
fn dispatch_wake(&self) {
self.audit.framework_callback(|| {
let wake = self.state.lock().unwrap_or_else(|p| p.into_inner()).wake.take();
if let Some(wake) = wake { wake.wake(); }
});
}
}
impl Wake for Shared {
fn wake(self: Arc<Self>) { self.dispatch_wake(); }
fn wake_by_ref(self: &Arc<Self>) { self.dispatch_wake(); }
}
pub(crate) struct DriverExecutor(Option<Arc<Shared>>);
impl Clone for DriverExecutor {
fn clone(&self) -> Self {
Self(Some(self.0.as_ref().unwrap().clone()))
}
}
impl Drop for DriverExecutor {
fn drop(&mut self) {
if let Some(shared) = self.0.take().and_then(Arc::into_inner) {
drop(shared);
}
}
}
fn destroy_backing(backing: ExactStored<Box<Backing>>) {
drop(backing);
}
impl DriverExecutor {
fn shared(&self) -> &Shared {
self.0.as_ref().unwrap()
}
fn new(stage: &RpcStagePermit) -> Result<Self, AdmissionError> {
let suffix = Layout::new::<RpcDependencyAudit>()
.extend(Layout::new::<RpcStagePermit>())
.map_err(|_| AdmissionError::SizeOverflow)?.0.pad_to_align();
let demand = StorageDemand::embedded(
Layout::new::<Mutex<State>>(),
suffix,
&[(Layout::new::<usize>(), 2)],
)?;
debug_assert_eq!(demand.bytes(), std::mem::size_of::<Shared>() + 2 * std::mem::size_of::<usize>());
let stage_owner = stage.retain_physical_owner()?;
let permit = stage.try_storage(demand)?;
Ok(Self(Some(Arc::new(Shared {
state: Mutex::new(State {
backing: None,
wake: None,
failure: None,
closed: false,
}),
permit,
audit: stage.dependency_audit(),
stage: stage_owner,
}))))
}
fn close(&self) {
let backing = self.shared().audit.framework_callback(|| {
let mut state = self.shared().state.lock().unwrap_or_else(|p| p.into_inner());
state.closed = true;
let wake = state.wake.take();
let backing = state.backing.take();
drop(state);
drop(wake);
backing
});
if let Some(backing) = backing {
destroy_backing(backing);
}
}
fn retire(&mut self) {
if let Some(shared) = self.0.take() {
let handoff = shared.stage.retirement_handle();
let owner: Arc<dyn RpcRetiredOwner> = shared;
handoff.install(owner);
handoff.reclaim();
}
}
fn failure(&self) -> Option<AdmissionError> {
self.shared()
.state
.lock()
.unwrap_or_else(|p| p.into_inner())
.failure
.take()
}
fn poll_drivers(&self, cx: &mut Context<'_>) {
let incoming = self.shared().audit.framework_callback(|| cx.waker().clone());
let count = self.shared().audit.framework_callback(|| {
let mut state = self.shared().state.lock().unwrap_or_else(|p| p.into_inner());
let previous = state.wake.replace(incoming);
let count = state.backing.as_ref().map_or(0, |b| b.get().tasks.get().len());
drop(state);
drop(previous);
count
});
for index in 0..count {
let task = {
let mut state = self
.shared()
.state
.lock()
.unwrap_or_else(|p| p.into_inner());
state.backing.as_mut().and_then(|b| b.get_mut().tasks.get_mut()[index].take())
};
if let Some(mut task) = task {
let bridge = Waker::from(Arc::clone(self.0.as_ref().unwrap()));
let mut driver_cx = Context::from_waker(&bridge);
if self.shared().audit.poll_driver(task.as_mut(), &mut driver_cx).is_pending() {
let mut state = self
.shared()
.state
.lock()
.unwrap_or_else(|p| p.into_inner());
if !state.closed {
state.backing.as_mut().unwrap().get_mut().tasks.get_mut()[index] = Some(task);
} else {
drop(state);
self.shared().audit.drop_driver(task);
}
} else {
self.shared().audit.drop_driver(task);
}
}
}
}
}
impl hyper::rt::Executor<Driver> for DriverExecutor {
fn execute(&self, future: Driver) {
struct PendingDriver {
future: Option<Driver>,
audit: RpcDependencyAudit,
}
impl Drop for PendingDriver {
fn drop(&mut self) {
if let Some(driver) = self.future.take() { self.audit.drop_driver(driver); }
}
}
self.shared().audit.framework_callback(|| {
let mut future = PendingDriver { future: Some(future), audit: self.shared().audit.clone() };
let mut retired = None;
let wake = {
let mut state = self
.shared()
.state
.lock()
.unwrap_or_else(|p| p.into_inner());
if !state.closed && state.failure.is_none() {
let result = (|| {
let count = state
.backing
.as_ref()
.map_or(0, |b| b.get().tasks.get().len())
.checked_add(1)
.ok_or(AdmissionError::SizeOverflow)?;
let layout = Layout::array::<Option<Driver>>(count)
.map_err(|_| AdmissionError::SizeOverflow)?;
let array_permit = self.shared().permit.try_reserve(
StorageDemand::separate(&[(layout, 1)])?)?;
let box_permit = self.shared().permit.try_reserve(
StorageDemand::separate(&[(Layout::new::<Backing>(), 1)])?)?;
let mut tasks = array_permit.allocate_exact_none_slice::<Driver>(count);
if let Some(old) = state.backing.as_mut() {
for (target, source) in tasks.get_mut().iter_mut()
.zip(old.get_mut().tasks.get_mut().iter_mut()) {
*target = source.take();
}
}
tasks.get_mut()[count - 1] = future.future.take();
let next = box_permit.allocate_exact(Layout::new::<Backing>(),
|| Box::new(Backing { tasks, audit: self.shared().audit.clone() }));
retired = state.backing.replace(next);
Ok::<_, AdmissionError>(())
})();
if let Err(error) = result {
state.failure = Some(error);
}
}
state.wake.take()
};
if let Some(old) = retired { destroy_backing(old); }
if let Some(wake) = wake {
wake.wake();
}
});
}
}
#[repr(C)]
struct CallHeap<F> {
future: Pin<Box<F>>,
permit: StoragePermit,
}
pub(crate) struct RpcCall<F> {
heap: Option<Box<CallHeap<F>>>,
executor: DriverExecutor,
stage: Option<RpcStagePermit>,
}
impl<F: Future> RpcCall<F> {
pub(crate) fn new(
stage: RpcStagePermit,
make: impl FnOnce(DriverExecutor) -> F,
) -> Result<Self, AdmissionError> {
struct Construction(Option<DriverExecutor>);
impl Drop for Construction {
fn drop(&mut self) {
if let Some(executor) = self.0.take() {
executor.close();
}
}
}
let mut construction = Construction(Some(DriverExecutor::new(&stage)?));
let permit = stage.try_storage(StorageDemand::embedded(
Layout::new::<Pin<Box<F>>>(),
Layout::new::<()>(),
&[(Layout::new::<F>(), 1)],
)?)?;
let future = Box::pin(make(construction.0.as_ref().unwrap().clone()));
let executor = construction.0.take().unwrap();
Ok(Self {
heap: Some(Box::new(CallHeap { future, permit })),
executor,
stage: Some(stage),
})
}
}
impl<F> RpcCall<F> {
fn close(&mut self) {
if self.stage.is_none() && self.heap.is_none() { return; }
struct Close<'a>(&'a DriverExecutor);
impl Drop for Close<'_> {
fn drop(&mut self) {
self.0.close();
}
}
let stage = self.stage.take();
let close = Close(&self.executor);
if let Some(heap) = self.heap.take() {
let CallHeap { future, permit } = *heap;
drop(future);
drop(permit);
}
drop(close);
self.executor.retire();
drop(stage);
}
}
impl<F> Drop for RpcCall<F> {
fn drop(&mut self) {
self.close();
}
}
impl<F: Future> Future for RpcCall<F> {
type Output = Result<F::Output, AdmissionError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
this.executor.poll_drivers(cx);
if let Some(error) = this.executor.failure() {
this.close();
return Poll::Ready(Err(error));
}
let result = this
.heap
.as_mut()
.expect("no poll after completion")
.future
.as_mut()
.poll(cx);
if let Some(error) = this.executor.failure() {
this.close();
return Poll::Ready(Err(error));
}
match result {
Poll::Ready(result) => {
this.close();
Poll::Ready(Ok(result))
}
Poll::Pending => Poll::Pending,
}
}
}
#[cfg(test)]
#[path = "rpc_driver_tests.rs"]
mod tests;