use std::marker::PhantomData;
use crate::point_to_point::Status;
use crate::transport;
use crate::{Count, Rank, Tag};
pub unsafe trait Scope<'a> {}
#[derive(Clone, Copy, Debug)]
pub struct StaticScope;
unsafe impl Scope<'static> for StaticScope {}
pub struct LocalScope<'a> {
_invariant: PhantomData<std::cell::Cell<&'a ()>>,
}
unsafe impl<'a> Scope<'a> for &LocalScope<'a> {}
pub fn scope<'a, F, R>(f: F) -> R
where
F: FnOnce(&LocalScope<'a>) -> R,
{
let scope = LocalScope {
_invariant: PhantomData,
};
f(&scope)
}
fn complete_recv(ctx: u32, source: Rank, tag: Tag, ptr: *mut u8, len: usize) -> Status {
let (src, t, count, _dt, payload) = transport::runtime().recv(ctx, source, tag);
let n = len.min(payload.len());
unsafe {
std::ptr::copy_nonoverlapping(payload.as_ptr(), ptr, n);
}
Status::new(src, t, count as Count, payload.len())
}
enum State {
Completed { status: Status },
PendingRecv {
ctx: u32,
source: Rank,
tag: Tag,
ptr: *mut u8,
len: usize,
},
PendingJoin {
handle: Option<std::thread::JoinHandle<()>>,
},
Consumed,
}
pub struct Request<'a, D: ?Sized = [u8], S = StaticScope> {
state: State,
_life: PhantomData<&'a mut ()>,
_data: PhantomData<*mut D>,
_scope: PhantomData<S>,
}
impl<'a, D: ?Sized, S: Scope<'a>> Request<'a, D, S> {
pub(crate) fn completed(_scope: S) -> Request<'a, D, S> {
Request {
state: State::Completed {
status: Status::new(0, 0, 0, 0),
},
_life: PhantomData,
_data: PhantomData,
_scope: PhantomData,
}
}
pub(crate) fn pending_recv(
_scope: S,
ptr: *mut u8,
len: usize,
ctx: u32,
source: Rank,
tag: Tag,
) -> Request<'a, D, S> {
Request {
state: State::PendingRecv {
ctx,
source,
tag,
ptr,
len,
},
_life: PhantomData,
_data: PhantomData,
_scope: PhantomData,
}
}
pub(crate) fn from_join(_scope: S, handle: std::thread::JoinHandle<()>) -> Request<'a, D, S> {
Request {
state: State::PendingJoin {
handle: Some(handle),
},
_life: PhantomData,
_data: PhantomData,
_scope: PhantomData,
}
}
fn ready(&self) -> bool {
match &self.state {
State::Completed { .. } => true,
State::Consumed => true,
State::PendingRecv {
ctx, source, tag, ..
} => transport::runtime().probe(*ctx, *source, *tag).is_some(),
State::PendingJoin { handle } => {
handle.as_ref().map(|h| h.is_finished()).unwrap_or(true)
}
}
}
pub fn wait(mut self) -> Status {
let state = std::mem::replace(&mut self.state, State::Consumed);
match state {
State::Completed { status } => status,
State::PendingRecv {
ctx,
source,
tag,
ptr,
len,
} => complete_recv(ctx, source, tag, ptr, len),
State::PendingJoin { mut handle } => {
if let Some(h) = handle.take() {
let _ = h.join();
}
Status::new(0, 0, 0, 0)
}
State::Consumed => unreachable!("request already consumed"),
}
}
pub fn wait_without_status(self) {
let _ = self.wait();
}
pub fn test(mut self) -> Result<Status, Request<'a, D, S>> {
let is_ready = self.ready();
if !is_ready {
return Err(self);
}
let state = std::mem::replace(&mut self.state, State::Consumed);
let status = match state {
State::Completed { status } => status,
State::PendingRecv {
ctx,
source,
tag,
ptr,
len,
} => complete_recv(ctx, source, tag, ptr, len),
State::PendingJoin { mut handle } => {
if let Some(h) = handle.take() {
let _ = h.join();
}
Status::new(0, 0, 0, 0)
}
State::Consumed => unreachable!(),
};
Ok(status)
}
pub fn cancel(mut self) {
self.state = State::Consumed;
}
}
impl<D: ?Sized, S> Drop for Request<'_, D, S> {
fn drop(&mut self) {
match &mut self.state {
State::PendingJoin { handle } => {
if let Some(h) = handle.take() {
let _ = h.join();
}
}
State::Completed { .. } | State::PendingRecv { .. } => {
if !std::thread::panicking() {
panic!(
"an in-flight mpi::request::Request was dropped; complete it with \
wait()/test()/cancel() or hold it in a WaitGuard"
);
}
}
State::Consumed => {}
}
}
}
pub struct WaitGuard<'a, D: ?Sized = [u8], S = StaticScope>(Option<Request<'a, D, S>>);
impl<'a, D: ?Sized, S: Scope<'a>> From<Request<'a, D, S>> for WaitGuard<'a, D, S> {
fn from(r: Request<'a, D, S>) -> Self {
WaitGuard(Some(r))
}
}
impl<'a, D: ?Sized, S: Scope<'a>> WaitGuard<'a, D, S> {
pub fn wait(mut self) -> Status {
self.0.take().unwrap().wait()
}
}
impl<D: ?Sized, S> Drop for WaitGuard<'_, D, S> {
fn drop(&mut self) {
if let Some(r) = self.0.take() {
let mut r = std::mem::ManuallyDrop::new(r);
let state = std::mem::replace(&mut r.state, State::Consumed);
match state {
State::PendingRecv {
ctx,
source,
tag,
ptr,
len,
} => {
let _ = complete_recv(ctx, source, tag, ptr, len);
}
State::PendingJoin { mut handle } => {
if let Some(h) = handle.take() {
let _ = h.join();
}
}
State::Completed { .. } | State::Consumed => {}
}
}
}
}
pub struct CancelGuard<'a, D: ?Sized = [u8], S = StaticScope>(Option<Request<'a, D, S>>);
impl<'a, D: ?Sized, S: Scope<'a>> From<Request<'a, D, S>> for CancelGuard<'a, D, S> {
fn from(r: Request<'a, D, S>) -> Self {
CancelGuard(Some(r))
}
}
impl<D: ?Sized, S> Drop for CancelGuard<'_, D, S> {
fn drop(&mut self) {
if let Some(r) = self.0.take() {
let mut r = std::mem::ManuallyDrop::new(r);
if let State::PendingJoin { handle } = &mut r.state {
if let Some(h) = handle.take() {
let _ = h.join();
}
}
r.state = State::Consumed;
}
}
}
pub fn wait_any<'a, D: ?Sized, S: Scope<'a>>(
requests: &mut Vec<Request<'a, D, S>>,
) -> Option<(usize, Status)> {
if requests.is_empty() {
return None;
}
loop {
for i in 0..requests.len() {
if requests[i].ready() {
let r = requests.remove(i);
return Some((i, r.wait()));
}
}
std::thread::yield_now();
}
}
pub fn wait_all<'a, D: ?Sized, S: Scope<'a>>(requests: Vec<Request<'a, D, S>>) -> Vec<Status> {
requests.into_iter().map(|r| r.wait()).collect()
}
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
struct GReqState {
done: AtomicBool,
status: std::sync::Mutex<Option<Status>>,
}
pub struct GeneralizedRequest {
state: Arc<GReqState>,
}
pub struct GeneralizedRequestCompleter {
state: Arc<GReqState>,
}
impl GeneralizedRequest {
pub fn start() -> (GeneralizedRequest, GeneralizedRequestCompleter) {
let state = Arc::new(GReqState {
done: AtomicBool::new(false),
status: std::sync::Mutex::new(None),
});
(
GeneralizedRequest {
state: Arc::clone(&state),
},
GeneralizedRequestCompleter { state },
)
}
pub fn is_complete(&self) -> bool {
self.state.done.load(Ordering::Acquire)
}
pub fn wait(self) -> Status {
while !self.state.done.load(Ordering::Acquire) {
std::thread::yield_now();
}
self.state
.status
.lock()
.unwrap()
.unwrap_or(Status::new(0, 0, 0, 0))
}
pub fn test(self) -> Result<Status, GeneralizedRequest> {
if self.state.done.load(Ordering::Acquire) {
Ok(self
.state
.status
.lock()
.unwrap()
.unwrap_or(Status::new(0, 0, 0, 0)))
} else {
Err(self)
}
}
}
impl GeneralizedRequestCompleter {
pub fn complete(self) {
*self.state.status.lock().unwrap() = Some(Status::new(0, 0, 0, 0));
self.state.done.store(true, Ordering::Release);
}
}
enum PersistentKind {
Send {
src: Rank,
dest_world: i32,
dt: u32,
count: u64,
},
Recv {
source: Rank,
},
}
pub struct PersistentRequest<'a> {
ctx: u32,
tag: Tag,
kind: PersistentKind,
ptr: *mut u8,
len: usize,
last: Option<Status>,
_life: PhantomData<&'a mut ()>,
}
impl<'a> PersistentRequest<'a> {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new_send(
ctx: u32,
src: Rank,
dest_world: i32,
tag: Tag,
dt: u32,
count: u64,
ptr: *const u8,
len: usize,
) -> PersistentRequest<'a> {
PersistentRequest {
ctx,
tag,
kind: PersistentKind::Send {
src,
dest_world,
dt,
count,
},
ptr: ptr as *mut u8,
len,
last: None,
_life: PhantomData,
}
}
pub(crate) fn new_recv(
ctx: u32,
source: Rank,
tag: Tag,
ptr: *mut u8,
len: usize,
) -> PersistentRequest<'a> {
PersistentRequest {
ctx,
tag,
kind: PersistentKind::Recv { source },
ptr,
len,
last: None,
_life: PhantomData,
}
}
pub fn start(&mut self) {
match self.kind {
PersistentKind::Send {
src,
dest_world,
dt,
count,
} => {
let bytes = unsafe { std::slice::from_raw_parts(self.ptr, self.len) };
transport::runtime()
.send(self.ctx, src, dest_world, self.tag, count, dt, bytes)
.expect("persistent send failed");
self.last = Some(Status::new(dest_world, self.tag, count as Count, self.len));
}
PersistentKind::Recv { .. } => {}
}
}
pub fn wait(&mut self) -> Status {
match self.kind {
PersistentKind::Send { .. } => self.last.take().unwrap_or(Status::new(0, 0, 0, 0)),
PersistentKind::Recv { source } => {
let (s, t, count, _dt, payload) =
transport::runtime().recv(self.ctx, source, self.tag);
let n = self.len.min(payload.len());
unsafe {
std::ptr::copy_nonoverlapping(payload.as_ptr(), self.ptr, n);
}
Status::new(s, t, count as Count, payload.len())
}
}
}
}