use crate::{raw, AtmiCtx, AtmiError, AtmiResult, TypedBuffer};
use std::cell::{Cell, RefCell};
use std::collections::HashMap;
use std::ffi::c_char;
use std::fmt;
use std::future::{poll_fn, Future};
use std::io;
use std::ops::Deref;
use std::pin::pin;
use std::task::{Poll, Waker};
use std::time::Instant;
#[cfg(any(feature = "async-io", feature = "tokio"))]
use std::os::fd::{FromRawFd, OwnedFd, RawFd};
pub trait AsyncReplyDriver: Sized {
type Readiness<'a>
where
Self: 'a;
fn register(reply_fd: i32) -> io::Result<Self>;
fn readable(&self) -> impl Future<Output = io::Result<Self::Readiness<'_>>> + '_;
fn clear_readiness(&self, readiness: &mut Self::Readiness<'_>);
fn sleep_until(&self, deadline: Instant) -> impl Future<Output = ()> + '_;
}
struct ParkedBuf {
ptr: *mut c_char,
len: usize,
}
#[cfg(feature = "ctx-send")]
unsafe impl Send for ParkedBuf {}
impl ParkedBuf {
const EMPTY: Self = Self {
ptr: std::ptr::null_mut(),
len: 0,
};
unsafe fn free(self, ctx: &AtmiCtx) {
if !self.ptr.is_null() {
drop(unsafe { TypedBuffer::from_raw(ctx, self.ptr) });
}
}
}
enum Slot {
Waiting {
waker: Option<Waker>,
claimed: bool,
},
Ready {
outcome: AtmiResult<()>,
buf: ParkedBuf,
claimed: bool,
},
}
struct AnyWaiterGuard<'a> {
demux: &'a ReplyDemux,
id: u64,
}
impl<'a> AnyWaiterGuard<'a> {
fn new(demux: &'a ReplyDemux) -> Self {
Self {
id: demux.register_any_waiter(),
demux,
}
}
fn id(&self) -> u64 {
self.id
}
}
impl Drop for AnyWaiterGuard<'_> {
fn drop(&mut self) {
self.demux.deregister_any_waiter(self.id);
}
}
#[derive(Default)]
struct AnyWaiter {
waker: Option<Waker>,
error: Option<AtmiError>,
}
#[derive(Clone, Copy)]
enum Target {
One(i32),
Any(u64),
}
#[derive(Default)]
struct ReplyDemux {
slots: RefCell<HashMap<i32, Slot>>,
scratch: RefCell<Option<ParkedBuf>>,
orphans: RefCell<Vec<ParkedBuf>>,
any_waiters: RefCell<HashMap<u64, AnyWaiter>>,
next_any_id: Cell<u64>,
deadlines: RefCell<HashMap<i32, Option<Instant>>>,
draining: Cell<bool>,
}
impl fmt::Debug for ReplyDemux {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ReplyDemux")
.field("slots", &self.slots.borrow().len())
.finish_non_exhaustive()
}
}
impl ReplyDemux {
fn register_fresh(&self, cd: i32, claimed: bool) {
let previous = self.slots.borrow_mut().insert(
cd,
Slot::Waiting {
waker: None,
claimed,
},
);
debug_assert!(
previous.is_none(),
"descriptor {cd} was reused while its slot was still occupied"
);
}
fn is_descriptor_busy(&self, cd: i32) -> bool {
self.slots.borrow().contains_key(&cd)
}
fn claim(&self, cd: i32) {
match self.slots.borrow_mut().entry(cd) {
std::collections::hash_map::Entry::Vacant(slot) => {
slot.insert(Slot::Waiting {
waker: None,
claimed: true,
});
}
std::collections::hash_map::Entry::Occupied(mut slot) => {
match slot.get_mut() {
Slot::Waiting { claimed, .. } => *claimed = true,
Slot::Ready { claimed, .. } => *claimed = true,
}
}
}
}
fn record_deadline(&self, cd: i32, deadline: Option<Instant>) {
self.deadlines.borrow_mut().insert(cd, deadline);
}
fn peek_deadline(&self, cd: i32) -> Option<Option<Instant>> {
self.deadlines.borrow().get(&cd).copied()
}
fn earliest_deadline(&self) -> Option<Instant> {
self.deadlines.borrow().values().filter_map(|d| *d).min()
}
fn forget_deadline(&self, cd: i32) {
self.deadlines.borrow_mut().remove(&cd);
}
fn is_ready(&self, cd: i32) -> bool {
matches!(self.slots.borrow().get(&cd), Some(Slot::Ready { .. }))
}
fn any_unclaimed(&self) -> Option<i32> {
self.slots
.borrow()
.iter()
.find_map(|(cd, slot)| match slot {
Slot::Ready { claimed: false, .. } => Some(*cd),
_ => None,
})
}
fn is_target_ready(&self, target: Target) -> bool {
match target {
Target::One(cd) => self.is_ready(cd),
Target::Any(id) => self.any_unclaimed().is_some() || self.has_any_error(id),
}
}
fn has_any_error(&self, id: u64) -> bool {
self.any_waiters
.borrow()
.get(&id)
.is_some_and(|entry| entry.error.is_some())
}
fn park_waker(&self, target: Target, id: u64, waker: &Waker) {
match target {
Target::One(cd) => {
if let Some(Slot::Waiting {
waker: slot_waker, ..
}) = self.slots.borrow_mut().get_mut(&cd)
{
match slot_waker {
Some(existing) if existing.will_wake(waker) => {}
other => *other = Some(waker.clone()),
}
}
}
Target::Any(_) => {
if let Some(entry) = self.any_waiters.borrow_mut().get_mut(&id) {
match &mut entry.waker {
Some(existing) if existing.will_wake(waker) => {}
other => *other = Some(waker.clone()),
}
}
}
}
}
fn wake_any_waiters(&self) {
let wakers: Vec<Waker> = self
.any_waiters
.borrow_mut()
.values_mut()
.filter_map(|entry| entry.waker.take())
.collect();
for waker in wakers {
waker.wake();
}
}
fn register_any_waiter(&self) -> u64 {
let id = self.next_any_id.get().wrapping_add(1);
self.next_any_id.set(id);
self.any_waiters
.borrow_mut()
.insert(id, AnyWaiter::default());
id
}
fn deregister_any_waiter(&self, id: u64) {
self.any_waiters.borrow_mut().remove(&id);
}
fn take_any_error(&self, id: u64) -> Option<AtmiError> {
self.any_waiters
.borrow_mut()
.get_mut(&id)
.and_then(|entry| entry.error.take())
}
fn deregister(&self, cd: i32) -> Option<ParkedBuf> {
match self.slots.borrow_mut().remove(&cd) {
Some(Slot::Ready { buf, .. }) => Some(buf),
_ => None,
}
}
fn take_ready(
&self,
target: Target,
data: &mut TypedBuffer<'_>,
flags: i64,
) -> Option<(i32, AtmiResult<()>)> {
let cd = match target {
Target::One(cd) if self.is_ready(cd) => cd,
Target::One(_) => return None,
Target::Any(_) => self.any_unclaimed()?,
};
let Some(Slot::Ready { outcome, buf, .. }) = self.slots.borrow_mut().remove(&cd) else {
return None;
};
Some((cd, self.hand_over(outcome, buf, data, flags)))
}
fn hand_over(
&self,
outcome: AtmiResult<()>,
buf: ParkedBuf,
data: &mut TypedBuffer<'_>,
flags: i64,
) -> AtmiResult<()> {
if buf.ptr.is_null() {
return outcome;
}
if flags & raw::TPNOCHANGE as i64 != 0 {
let incoming = unsafe { TypedBuffer::borrowed_from_raw(data.ctx, buf.ptr) };
match (incoming.tptypes(), data.tptypes()) {
(Ok(got), Ok(want))
if got.type_name != want.type_name || got.subtype != want.subtype =>
{
let message = format!(
"TPNOCHANGE: receiver expects {}/{} but got {}/{} buffer",
want.type_name, want.subtype, got.type_name, got.subtype
);
self.stash(buf);
return Err(AtmiError::new(raw::TPEOTYPE, message));
}
_ => {}
}
}
let previous = ParkedBuf {
ptr: data.as_ptr(),
len: data.len(),
};
data.replace_ptr(buf.ptr);
data.set_len(buf.len);
self.stash(previous);
outcome
}
fn stash(&self, buf: ParkedBuf) {
if buf.ptr.is_null() {
return;
}
let mut scratch = self.scratch.borrow_mut();
if scratch.is_none() {
*scratch = Some(buf);
} else {
drop(scratch);
self.orphans.borrow_mut().push(buf);
}
}
fn drain(&self, ctx: &AtmiCtx, flags: i64) {
if self.draining.replace(true) {
return;
}
self.drain_inner(ctx, flags);
self.draining.set(false);
}
fn drain_inner(&self, ctx: &AtmiCtx, _flags: i64) {
let get_flags = raw::TPNOBLOCK as i64 | raw::TPGETANY as i64;
loop {
let mut buffer = match self.take_scratch(ctx) {
Ok(buffer) => buffer,
Err(err) => return self.fail_all_waiting(err),
};
let mut cd = 0i32;
let outcome = ctx.tpgetrply(&mut cd, &mut buffer, get_flags);
let len = buffer.len();
let buf = ParkedBuf {
ptr: buffer.into_raw(),
len,
};
let blocked = matches!(&outcome, Err(err) if err.code == raw::TPEBLOCK);
if blocked {
self.stash(buf);
return;
}
if cd > 0 {
self.route(cd, outcome, buf);
continue;
}
self.stash(buf);
match outcome {
Err(err) => return self.fail_all_waiting(err),
Ok(()) => return,
}
}
}
fn take_scratch<'c>(&self, ctx: &'c AtmiCtx) -> AtmiResult<TypedBuffer<'c>> {
if let Some(buf) = self.scratch.borrow_mut().take() {
let mut buffer = unsafe { TypedBuffer::from_raw(ctx, buf.ptr) };
buffer.set_len(buf.len);
return Ok(buffer);
}
ctx.tpalloc("CARRAY", "", 1024)
}
fn route(&self, cd: i32, outcome: AtmiResult<()>, buf: ParkedBuf) {
let waker = {
let mut slots = self.slots.borrow_mut();
match slots.get_mut(&cd) {
Some(slot) => {
let (waker, was_claimed) = match slot {
Slot::Waiting { waker, claimed } => (waker.take(), *claimed),
Slot::Ready { .. } => {
drop(slots);
self.stash(buf);
return;
}
};
*slot = Slot::Ready {
outcome,
buf,
claimed: was_claimed,
};
waker
}
None => {
slots.insert(
cd,
Slot::Ready {
outcome,
buf,
claimed: false,
},
);
None
}
}
};
if let Some(waker) = waker {
waker.wake();
}
self.wake_any_waiters();
}
fn fail_all_waiting(&self, err: AtmiError) {
{
let mut any = self.any_waiters.borrow_mut();
for entry in any.values_mut() {
entry.error = Some(err.clone());
}
}
let mut wakers = Vec::new();
{
let mut slots = self.slots.borrow_mut();
for slot in slots.values_mut() {
if let Slot::Waiting { waker, .. } = slot {
wakers.extend(waker.take());
*slot = Slot::Ready {
outcome: Err(err.clone()),
buf: ParkedBuf::EMPTY,
claimed: true,
};
}
}
}
for waker in wakers {
waker.wake();
}
self.wake_any_waiters();
}
unsafe fn release(&self, ctx: &AtmiCtx) {
let slots: Vec<Slot> = self.slots.borrow_mut().drain().map(|(_, v)| v).collect();
for slot in slots {
if let Slot::Ready { buf, .. } = slot {
unsafe { buf.free(ctx) };
}
}
if let Some(buf) = self.scratch.borrow_mut().take() {
unsafe { buf.free(ctx) };
}
for buf in self.orphans.borrow_mut().drain(..) {
unsafe { buf.free(ctx) };
}
}
}
#[derive(Debug)]
pub struct AsyncAtmiCtx<D> {
driver: D,
demux: ReplyDemux,
context: AtmiCtx,
}
impl<D: AsyncReplyDriver> AsyncAtmiCtx<D> {
pub fn new(context: AtmiCtx) -> AtmiResult<Self> {
let reply_fd = context.reply_queue_fd()?;
let driver = D::register(reply_fd).map_err(|err| {
AtmiError::new(
raw::TPEOS,
format!("failed to register Enduro/X reply queue fd: {err}"),
)
})?;
Ok(Self {
driver,
demux: ReplyDemux::default(),
context,
})
}
pub const SUPPORTED: bool = cfg!(endurox_pollable);
pub fn context(&self) -> &AtmiCtx {
&self.context
}
pub fn into_inner(self) -> AtmiCtx {
let this = std::mem::ManuallyDrop::new(self);
unsafe {
let driver = std::ptr::read(&this.driver);
let demux = std::ptr::read(&this.demux);
let context = std::ptr::read(&this.context);
drop(driver);
demux.release(&context);
context
}
}
pub async fn tpcall(
&self,
svc: &str,
idata: &TypedBuffer<'_>,
odata: &mut TypedBuffer<'_>,
flags: i64,
) -> AtmiResult<()> {
Self::check_supported_flags(flags)?;
Self::reject_tpnoreply(flags)?;
let deadline = self.deadline_for(flags)?;
let cd = self.submit(svc, idata, flags, deadline, true)?;
let mut pending = AsyncPendingCall::new(self, cd);
let result = match self
.await_reply(Target::One(pending.cd), odata, flags, deadline, false)
.await
{
Ok((_, outcome)) => outcome,
Err(err) => Err(err),
};
if result.is_ok() {
pending.complete();
}
result
}
pub fn tpacall(&self, svc: &str, idata: &TypedBuffer<'_>, flags: i64) -> AtmiResult<i32> {
Self::check_supported_flags(flags)?;
let deadline = self.deadline_for(flags)?;
self.submit(svc, idata, flags, deadline, false)
}
fn submit(
&self,
svc: &str,
idata: &TypedBuffer<'_>,
flags: i64,
deadline: Option<Instant>,
claimed: bool,
) -> AtmiResult<i32> {
let cd = self.context.tpacall(svc, idata, flags)?;
if cd > 0 && self.demux.is_descriptor_busy(cd) {
let _ = self.context.tpcancel(cd);
self.demux.drain(&self.context, 0);
self.demux.wake_any_waiters();
return Err(AtmiError::new(
raw::TPELIMIT,
"call descriptor is still tracked by an earlier call; collect or \
cancel that one before issuing another request",
));
}
if cd > 0 {
self.demux.register_fresh(cd, claimed);
self.demux.record_deadline(cd, deadline);
}
Ok(cd)
}
pub async fn tpcall_async(
&self,
svc: &str,
idata: &TypedBuffer<'_>,
odata: &mut TypedBuffer<'_>,
flags: i64,
) -> AtmiResult<()> {
self.tpcall(svc, idata, odata, flags).await
}
pub async fn tpgetrply(
&self,
cd: &mut i32,
data: &mut TypedBuffer<'_>,
flags: i64,
) -> AtmiResult<()> {
Self::check_supported_flags(flags)?;
let any = flags & raw::TPGETANY as i64 != 0;
let deadline = if any {
match self.demux.earliest_deadline() {
Some(deadline) => Some(deadline),
None => self.deadline_for(flags)?,
}
} else {
match self.demux.peek_deadline(*cd) {
Some(recorded) => recorded,
None => self.deadline_for(flags)?,
}
};
let any_waiter = if any {
Some(AnyWaiterGuard::new(&self.demux))
} else {
None
};
let target = if let Some(guard) = &any_waiter {
Target::Any(guard.id())
} else {
self.demux.claim(*cd);
Target::One(*cd)
};
let (replied_cd, outcome) = self
.await_reply(target, data, flags, deadline, true)
.await?;
self.demux.forget_deadline(replied_cd);
*cd = replied_cd;
outcome
}
fn is_nonblocking(flags: i64) -> bool {
flags & raw::TPNOBLOCK as i64 != 0
}
pub async fn tpgetrply_async(
&self,
cd: &mut i32,
data: &mut TypedBuffer<'_>,
flags: i64,
) -> AtmiResult<()> {
self.tpgetrply(cd, data, flags).await
}
pub fn tpcancel(&self, cd: i32) -> AtmiResult<()> {
self.release_slot(cd);
let result = self.context.tpcancel(cd);
self.demux.drain(&self.context, 0);
self.demux.wake_any_waiters();
result
}
pub fn tpterm(self) -> AtmiResult<()> {
self.into_inner().tpterm()
}
fn check_supported_flags(flags: i64) -> AtmiResult<()> {
const UNSUPPORTED: &[(i64, &str)] = &[
(raw::TPNOABORT as i64, "TPNOABORT"),
(raw::TPTRANSUSPEND as i64, "TPTRANSUSPEND"),
];
for (bit, name) in UNSUPPORTED {
if flags & *bit != 0 {
return Err(AtmiError::new(
raw::TPEINVAL,
format!(
"{name} cannot be honoured per call by the async reply demux, because one drain collects replies for several descriptors; use the blocking API for this call"
),
));
}
}
Ok(())
}
fn reject_tpnoreply(flags: i64) -> AtmiResult<()> {
if flags & raw::TPNOREPLY as i64 != 0 {
return Err(AtmiError::new(
raw::TPEINVAL,
"TPNOREPLY cannot be used with tpcall()",
));
}
Ok(())
}
fn deadline_for(&self, flags: i64) -> AtmiResult<Option<Instant>> {
if flags & raw::TPNOTIME as i64 != 0 {
return Ok(None);
}
self.context.reply_deadline()
}
fn release_slot(&self, cd: i32) {
self.demux.forget_deadline(cd);
if let Some(buf) = self.demux.deregister(cd) {
unsafe { buf.free(&self.context) };
}
}
async fn await_reply(
&self,
target: Target,
data: &mut TypedBuffer<'_>,
flags: i64,
deadline: Option<Instant>,
cancel_on_timeout: bool,
) -> AtmiResult<(i32, AtmiResult<()>)> {
loop {
if let Some(ready) = self.demux.take_ready(target, data, flags) {
return Ok(ready);
}
self.demux.drain(&self.context, flags);
if let Some(ready) = self.demux.take_ready(target, data, flags) {
return Ok(ready);
}
if let Target::Any(id) = target {
if let Some(err) = self.demux.take_any_error(id) {
return Err(err);
}
}
if Self::is_nonblocking(flags) {
return Err(AtmiError::new(
raw::TPEBLOCK,
"TPNOBLOCK was specified and no reply is available",
));
}
let (wake, readiness) = self.wait_for_wake(target, deadline).await?;
match wake {
Wake::Progress => {
self.demux.drain(&self.context, flags);
if let Some(ready) = self.demux.take_ready(target, data, flags) {
return Ok(ready);
}
if let Some(mut token) = readiness {
self.driver.clear_readiness(&mut token);
}
continue;
}
Wake::Timeout => {
self.demux.drain(&self.context, flags);
if let Some(ready) = self.demux.take_ready(target, data, flags) {
return Ok(ready);
}
if let Target::One(cd) = target {
if cancel_on_timeout {
let _ = self.tpcancel(cd);
} else {
self.release_slot(cd);
}
}
return Err(AtmiError::new(raw::TPETIME, "async reply wait timed out"));
}
}
}
}
#[allow(clippy::type_complexity)]
async fn wait_for_wake(
&self,
target: Target,
deadline: Option<Instant>,
) -> AtmiResult<(Wake, Option<D::Readiness<'_>>)> {
let readable = self.driver.readable();
let mut readable = pin!(readable);
let mut readiness = None;
let wake = match deadline {
Some(deadline) => {
let timer = self.driver.sleep_until(deadline);
let mut timer = pin!(timer);
poll_fn(|cx| {
match target {
Target::One(cd) => self.demux.park_waker(target, cd as u64, cx.waker()),
Target::Any(id) => self.demux.park_waker(target, id, cx.waker()),
}
if self.demux.is_target_ready(target) {
return Poll::Ready(Ok(Wake::Progress));
}
if let Poll::Ready(result) = readable.as_mut().poll(cx) {
return Poll::Ready(match result {
Ok(token) => {
readiness = Some(token);
Ok(Wake::Progress)
}
Err(err) => Err(driver_error(err)),
});
}
if timer.as_mut().poll(cx).is_ready() {
return Poll::Ready(Ok(Wake::Timeout));
}
Poll::Pending
})
.await
}
None => {
poll_fn(|cx| {
match target {
Target::One(cd) => self.demux.park_waker(target, cd as u64, cx.waker()),
Target::Any(id) => self.demux.park_waker(target, id, cx.waker()),
}
if self.demux.is_target_ready(target) {
return Poll::Ready(Ok(Wake::Progress));
}
if let Poll::Ready(result) = readable.as_mut().poll(cx) {
return Poll::Ready(match result {
Ok(token) => {
readiness = Some(token);
Ok(Wake::Progress)
}
Err(err) => Err(driver_error(err)),
});
}
Poll::Pending
})
.await
}
}?;
Ok((wake, readiness))
}
}
enum Wake {
Progress,
Timeout,
}
impl<D> Deref for AsyncAtmiCtx<D> {
type Target = AtmiCtx;
fn deref(&self) -> &Self::Target {
&self.context
}
}
impl<D> AsRef<AtmiCtx> for AsyncAtmiCtx<D> {
fn as_ref(&self) -> &AtmiCtx {
&self.context
}
}
impl<D> Drop for AsyncAtmiCtx<D> {
fn drop(&mut self) {
unsafe { self.demux.release(&self.context) };
}
}
impl AtmiCtx {
#[cfg(feature = "async")]
pub fn into_async<D: AsyncReplyDriver>(self) -> AtmiResult<AsyncAtmiCtx<D>> {
AsyncAtmiCtx::new(self)
}
#[cfg(feature = "async")]
pub const ASYNC_SUPPORTED: bool = cfg!(endurox_pollable);
}
struct AsyncPendingCall<'ctx, D: AsyncReplyDriver> {
context: &'ctx AsyncAtmiCtx<D>,
cd: i32,
armed: bool,
}
impl<'ctx, D: AsyncReplyDriver> AsyncPendingCall<'ctx, D> {
fn new(context: &'ctx AsyncAtmiCtx<D>, cd: i32) -> Self {
Self {
context,
cd,
armed: true,
}
}
fn complete(&mut self) {
self.armed = false;
}
}
impl<D: AsyncReplyDriver> Drop for AsyncPendingCall<'_, D> {
fn drop(&mut self) {
if self.armed {
let _ = self.context.tpcancel(self.cd);
} else {
self.context.release_slot(self.cd);
}
}
}
fn driver_error(err: io::Error) -> AtmiError {
AtmiError::new(
raw::TPEOS,
format!("async wait on Enduro/X reply queue failed: {err}"),
)
}
#[cfg(any(feature = "async-io", feature = "tokio"))]
fn duplicate_reply_fd(reply_fd: RawFd) -> io::Result<OwnedFd> {
let duplicate = unsafe { libc::fcntl(reply_fd, libc::F_DUPFD_CLOEXEC, 0) };
if duplicate < 0 {
Err(io::Error::last_os_error())
} else {
Ok(unsafe { OwnedFd::from_raw_fd(duplicate) })
}
}
#[cfg(feature = "tokio")]
#[derive(Debug)]
pub struct TokioReplyDriver {
reply_fd: tokio::io::unix::AsyncFd<OwnedFd>,
_local: std::marker::PhantomData<std::rc::Rc<()>>,
}
#[cfg(feature = "tokio")]
impl AsyncReplyDriver for TokioReplyDriver {
type Readiness<'a> = tokio::io::unix::AsyncFdReadyGuard<'a, OwnedFd>;
fn register(reply_fd: i32) -> io::Result<Self> {
let reply_fd = duplicate_reply_fd(reply_fd)?;
let reply_fd = std::panic::catch_unwind(|| {
tokio::io::unix::AsyncFd::with_interest(reply_fd, tokio::io::Interest::READABLE)
})
.map_err(|_| io::Error::other("Tokio runtime does not have an I/O driver enabled"))??;
Ok(Self {
reply_fd,
_local: std::marker::PhantomData,
})
}
fn readable(&self) -> impl Future<Output = io::Result<Self::Readiness<'_>>> + '_ {
self.reply_fd.readable()
}
fn clear_readiness(&self, readiness: &mut Self::Readiness<'_>) {
readiness.clear_ready();
}
fn sleep_until(&self, deadline: Instant) -> impl Future<Output = ()> + '_ {
tokio::time::sleep_until(deadline.into())
}
}
#[cfg(feature = "tokio")]
pub type TokioAtmiCtx = AsyncAtmiCtx<TokioReplyDriver>;
#[cfg(feature = "tokio")]
impl AtmiCtx {
pub fn into_tokio(self) -> AtmiResult<TokioAtmiCtx> {
self.into_async()
}
pub const TOKIO_ASYNC_SUPPORTED: bool = cfg!(endurox_pollable);
}
#[cfg(feature = "async-io")]
#[derive(Debug)]
pub struct AsyncIoReplyDriver {
reply_fd: async_io::Async<OwnedFd>,
}
#[cfg(feature = "async-io")]
impl AsyncReplyDriver for AsyncIoReplyDriver {
type Readiness<'a> = ();
fn register(reply_fd: i32) -> io::Result<Self> {
let reply_fd = duplicate_reply_fd(reply_fd)?;
Ok(Self {
reply_fd: async_io::Async::new_nonblocking(reply_fd)?,
})
}
fn readable(&self) -> impl Future<Output = io::Result<Self::Readiness<'_>>> + '_ {
self.reply_fd.readable()
}
fn clear_readiness(&self, _readiness: &mut Self::Readiness<'_>) {
}
#[allow(clippy::manual_async_fn)]
fn sleep_until(&self, deadline: Instant) -> impl Future<Output = ()> + '_ {
async move {
let _ = async_io::Timer::at(deadline).await;
}
}
}
#[cfg(feature = "async-io")]
pub type AsyncIoAtmiCtx = AsyncAtmiCtx<AsyncIoReplyDriver>;
#[cfg(feature = "async-io")]
impl AtmiCtx {
pub fn into_async_io(self) -> AtmiResult<AsyncIoAtmiCtx> {
self.into_async()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::task::Wake;
#[cfg(any(feature = "async-io", feature = "tokio"))]
use std::io::{Read, Write};
#[cfg(any(feature = "async-io", feature = "tokio"))]
use std::os::fd::AsRawFd;
#[cfg(any(feature = "async-io", feature = "tokio"))]
use std::os::unix::net::UnixStream;
struct CountingWaker(AtomicUsize);
impl Wake for CountingWaker {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn reply_is_routed_to_its_own_descriptor() {
let demux = ReplyDemux::default();
demux.register_fresh(5, true);
demux.register_fresh(7, true);
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
demux.park_waker(Target::One(7), 7, &Waker::from(counter.clone()));
assert!(!demux.is_ready(7));
demux.route(7, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_ready(7), "reply must land in descriptor 7's slot");
assert!(!demux.is_ready(5), "descriptor 5 must be untouched");
assert_eq!(
counter.0.load(Ordering::SeqCst),
1,
"the parked future for 7 must be woken exactly once"
);
}
#[test]
fn reply_arriving_before_tpgetrply_is_held_for_the_caller() {
let demux = ReplyDemux::default();
demux.route(9, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_ready(9), "unclaimed reply must be parked");
demux.claim(9);
assert!(
demux.is_ready(9),
"claiming a descriptor must not discard a reply already parked for it"
);
}
#[test]
fn tpgetany_takes_any_unclaimed_reply() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
assert!(demux.any_unclaimed().is_none(), "nothing to collect yet");
demux.route(4, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_target_ready(Target::Any(any_id)));
assert_eq!(demux.any_unclaimed(), Some(4));
}
#[test]
fn tpgetany_collects_a_registered_manual_tpacall_reply() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
demux.register_fresh(7, false);
assert!(
demux.is_descriptor_busy(7),
"the descriptor is protected from reuse"
);
assert!(!demux.is_target_ready(Target::Any(any_id)));
demux.route(7, Ok(()), ParkedBuf::EMPTY);
assert_eq!(
demux.any_unclaimed(),
Some(7),
"a manual tpacall reply must stay collectable by TPGETANY"
);
assert!(demux.is_target_ready(Target::Any(any_id)));
}
#[test]
fn naming_a_descriptor_claims_it_from_tpgetany() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
demux.register_fresh(7, false);
demux.claim(7);
demux.route(7, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_ready(7), "its own collector still finds it");
assert!(
!demux.is_target_ready(Target::Any(any_id)),
"TPGETANY must not take a reply the caller asked for by number"
);
}
#[test]
fn claiming_after_the_reply_landed_keeps_it() {
let demux = ReplyDemux::default();
demux.register_fresh(7, false);
demux.route(7, Ok(()), ParkedBuf::EMPTY);
demux.claim(7);
assert!(
demux.is_ready(7),
"claiming a descriptor must not discard a reply already parked"
);
}
#[test]
fn tpgetany_does_not_steal_a_claimed_reply() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
demux.register_fresh(5, true);
demux.route(5, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_ready(5), "the owner can still collect it");
assert!(
!demux.is_target_ready(Target::Any(any_id)),
"TPGETANY must not see a reply owned by a pending tpcall"
);
demux.route(6, Ok(()), ParkedBuf::EMPTY);
assert_eq!(demux.any_unclaimed(), Some(6));
}
#[test]
fn tpgetany_waiter_is_woken_by_any_reply() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
demux.park_waker(Target::Any(any_id), any_id, &Waker::from(counter.clone()));
demux.route(3, Ok(()), ParkedBuf::EMPTY);
assert_eq!(
counter.0.load(Ordering::SeqCst),
1,
"a routed reply must wake the TPGETANY waiter"
);
}
#[test]
fn undirected_error_reaches_each_current_any_waiter_only() {
let demux = ReplyDemux::default();
let first = demux.register_any_waiter();
let second = demux.register_any_waiter();
demux.fail_all_waiting(AtmiError::new(raw::TPEOS, "queue is gone"));
assert!(
demux.take_any_error(first).is_some(),
"first TPGETANY waiter must receive the failure"
);
assert!(
demux.take_any_error(second).is_some(),
"second TPGETANY waiter must receive it too, not just the first"
);
let later = demux.register_any_waiter();
assert!(
demux.take_any_error(later).is_none(),
"a later waiter must not inherit an earlier queue failure"
);
}
#[test]
fn any_tracked_descriptor_is_unsafe_to_reuse() {
let demux = ReplyDemux::default();
demux.route(9, Ok(()), ParkedBuf::EMPTY);
assert!(
demux.is_descriptor_busy(9),
"uncollected reply blocks reuse"
);
demux.register_fresh(4, true);
demux.route(4, Ok(()), ParkedBuf::EMPTY);
assert!(
demux.is_descriptor_busy(4),
"a claimed reply must block reuse too: its future may simply not have resumed yet"
);
demux.register_fresh(6, true);
assert!(demux.is_descriptor_busy(6));
drop(demux.deregister(9));
assert!(!demux.is_descriptor_busy(9));
}
#[test]
fn undirected_error_makes_the_any_waiter_ready() {
let demux = ReplyDemux::default();
let any_id = demux.register_any_waiter();
assert!(!demux.is_target_ready(Target::Any(any_id)));
demux.fail_all_waiting(AtmiError::new(raw::TPEOS, "queue is gone"));
assert!(
demux.is_target_ready(Target::Any(any_id)),
"a pending queue error must count as readiness for its waiter"
);
}
#[test]
fn released_descriptor_becomes_reusable() {
let demux = ReplyDemux::default();
demux.route(9, Ok(()), ParkedBuf::EMPTY);
assert!(demux.is_descriptor_busy(9));
drop(demux.deregister(9));
assert!(!demux.is_descriptor_busy(9));
demux.register_fresh(9, true);
assert!(
!demux.is_ready(9),
"the reused descriptor starts from a clean slot"
);
assert!(demux.is_descriptor_busy(9), "and is tracked again");
}
#[test]
fn undirected_error_fails_every_waiter() {
let demux = ReplyDemux::default();
demux.register_fresh(5, true);
demux.register_fresh(7, true);
demux.fail_all_waiting(AtmiError::new(raw::TPEOS, "queue is gone"));
assert!(demux.is_ready(5));
assert!(demux.is_ready(7));
}
#[test]
fn deregister_clears_only_that_descriptors_slot() {
let demux = ReplyDemux::default();
demux.register_fresh(5, true);
demux.register_fresh(7, true);
assert!(demux.deregister(5).is_none());
assert!(!demux.is_ready(5), "cancelled descriptor keeps no slot");
demux.route(7, Ok(()), ParkedBuf::EMPTY);
assert!(
demux.is_ready(7),
"the live descriptor still routes normally"
);
}
#[cfg(feature = "tokio")]
#[test]
fn tokio_driver_waits_without_owning_endurox_fd() {
let (mut read_end, mut write_end) = UnixStream::pair().expect("socket pair failed");
read_end
.set_nonblocking(true)
.expect("set_nonblocking failed");
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("Tokio runtime creation failed");
let driver = runtime.block_on(async {
TokioReplyDriver::register(read_end.as_raw_fd())
.expect("Tokio driver registration failed")
});
write_end.write_all(b"x").expect("socket write failed");
runtime.block_on(async {
let mut readiness = driver.readable().await.expect("readiness wait failed");
let mut byte = [0];
read_end.read_exact(&mut byte).expect("socket read failed");
assert_eq!(byte, *b"x");
assert_eq!(
read_end
.read(&mut byte)
.expect_err("empty socket should return WouldBlock")
.kind(),
io::ErrorKind::WouldBlock
);
driver.clear_readiness(&mut readiness);
});
drop(driver);
write_end
.write_all(b"y")
.expect("second socket write failed");
let mut byte = [0];
read_end
.read_exact(&mut byte)
.expect("original descriptor was closed by driver");
assert_eq!(byte, *b"y");
}
#[cfg(feature = "async-io")]
#[test]
fn async_io_driver_waits_without_owning_endurox_fd() {
let (mut read_end, mut write_end) = UnixStream::pair().expect("socket pair failed");
read_end
.set_nonblocking(true)
.expect("set_nonblocking failed");
let driver = AsyncIoReplyDriver::register(read_end.as_raw_fd())
.expect("async-io driver registration failed");
write_end.write_all(b"x").expect("socket write failed");
async_io::block_on(async {
driver.readable().await.expect("readiness wait failed");
let mut byte = [0];
read_end.read_exact(&mut byte).expect("socket read failed");
assert_eq!(byte, *b"x");
driver.clear_readiness(&mut ());
});
drop(driver);
write_end
.write_all(b"y")
.expect("second socket write failed");
let mut byte = [0];
read_end
.read_exact(&mut byte)
.expect("original descriptor was closed by driver");
assert_eq!(byte, *b"y");
}
#[cfg(all(feature = "async-io", feature = "tokio"))]
#[test]
fn async_io_driver_future_runs_on_tokio_executor() {
let (mut read_end, mut write_end) = UnixStream::pair().expect("socket pair failed");
read_end
.set_nonblocking(true)
.expect("set_nonblocking failed");
let driver = AsyncIoReplyDriver::register(read_end.as_raw_fd())
.expect("async-io driver registration failed");
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("Tokio runtime creation failed");
write_end.write_all(b"x").expect("socket write failed");
runtime.block_on(async {
driver
.readable()
.await
.expect("Tokio did not poll async-io readiness");
let mut byte = [0];
read_end.read_exact(&mut byte).expect("socket read failed");
assert_eq!(byte, *b"x");
});
}
}