1use std::ffi::c_void;
2use std::marker::PhantomData;
3
4use mpi::collective::CommunicatorCollectives;
5use mpi::datatype::Equivalence;
6use mpi::ffi;
7use mpi::raw::AsRaw;
8use mpi::topology::{Communicator, Rank};
9
10use crate::Error;
11
12mod sealed {
13 pub trait Sealed {}
14}
15
16pub trait RmaElement: sealed::Sealed + Equivalence + Copy + Send + Sync + 'static {
22 #[doc(hidden)]
23 const TYPE_ID: u8;
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum MemoryModel {
28 Unified,
29 Separate,
30}
31
32macro_rules! elements {
33 ($($t:ty => $id:literal),* $(,)?) => {$(
34 impl sealed::Sealed for $t {}
35 impl RmaElement for $t {
36 const TYPE_ID: u8 = $id;
37 }
38 )*};
39}
40
41elements!(
42 u8 => 1,
43 u16 => 2,
44 u32 => 3,
45 u64 => 4,
46 usize => 5,
47 i8 => 6,
48 i16 => 7,
49 i32 => 8,
50 i64 => 9,
51 isize => 10,
52 f32 => 11,
53 f64 => 12,
54);
55
56fn check(code: i32) -> Result<(), Error> {
57 if code == ffi::MPI_SUCCESS as i32 {
58 Ok(())
59 } else {
60 Err(Error::Mpi(code))
61 }
62}
63
64unsafe fn init_epoch<T>(win: ffi::MPI_Win, base: *mut T, len: usize) -> Result<MemoryModel, Error> {
71 if len > 0 && base.is_null() {
72 return Err(Error::Mpi(ffi::MPI_ERR_WIN as i32));
73 }
74 let mut value: *mut c_void = std::ptr::null_mut();
75 let mut flag = 0;
76 unsafe {
77 check(ffi::MPI_Win_get_attr(
78 win,
79 ffi::MPI_WIN_MODEL as i32,
80 &mut value as *mut *mut c_void as *mut c_void,
81 &mut flag,
82 ))?;
83 }
84 if flag == 0 || value.is_null() {
85 return Err(Error::Mpi(ffi::MPI_ERR_WIN as i32));
86 }
87 let model = unsafe {
88 if *(value as *const i32) == ffi::MPI_WIN_UNIFIED as i32 {
89 MemoryModel::Unified
90 } else {
91 MemoryModel::Separate
92 }
93 };
94 unsafe {
97 check(ffi::MPI_Win_lock_all(ffi::MPI_MODE_NOCHECK as i32, win))?;
98 }
99 if len > 0 {
102 unsafe { std::ptr::write_bytes(base, 0, len) };
103 }
104 if model == MemoryModel::Separate {
105 unsafe { check(ffi::MPI_Win_sync(win))? };
106 }
107 Ok(model)
108}
109
110pub trait CommunicatorRmaExt: Communicator {
112 fn allocate_window<T: RmaElement>(&self, len: usize) -> Result<Window<T>, Error> {
113 Window::allocate(self, len)
114 }
115}
116
117impl<C: Communicator + ?Sized> CommunicatorRmaExt for C {}
118
119pub struct Window<T: RmaElement> {
129 win: ffi::MPI_Win,
130 base: *mut T,
131 len: usize,
132 lengths: Vec<usize>,
133 rank: Rank,
134 ranks: Rank,
135 model: MemoryModel,
136 closed: bool,
137 _element: PhantomData<T>,
138}
139
140unsafe impl<T: RmaElement> Send for Window<T> {}
143unsafe impl<T: RmaElement> Sync for Window<T> {}
144
145impl<T: RmaElement> Window<T> {
146 fn allocate<C: Communicator + ?Sized>(comm: &C, len: usize) -> Result<Self, Error> {
147 if comm.test_inter() {
148 return Err(Error::Intercommunicator);
149 }
150 let threading = mpi::environment::threading_support();
151 if threading != mpi::Threading::Multiple {
152 return Err(Error::Threading(threading));
153 }
154 let width = std::mem::size_of::<T>();
155 let config = [
156 u64::try_from(len).map_err(|_| Error::SizeOverflow)?,
157 u64::try_from(width).map_err(|_| Error::SizeOverflow)?,
158 T::TYPE_ID as u64,
159 ];
160 let ranks = usize::try_from(comm.size()).map_err(|_| Error::SizeOverflow)?;
161 let configs_len = ranks.checked_mul(config.len()).ok_or(Error::SizeOverflow)?;
162 if configs_len > i32::MAX as usize {
163 return Err(Error::CountOverflow);
164 }
165 let mut configs = vec![0; configs_len];
166 comm.all_gather_into(&config[..], &mut configs[..]);
167 if configs
168 .chunks_exact(config.len())
169 .any(|c| c[1..] != config[1..])
170 {
171 return Err(Error::Window("element type differs between ranks"));
172 }
173 let lengths = configs
174 .chunks_exact(config.len())
175 .map(|c| usize::try_from(c[0]).map_err(|_| Error::SizeOverflow))
176 .collect::<Result<Vec<_>, _>>()?;
177 if lengths.iter().any(|&len| {
178 len.checked_mul(width)
179 .and_then(|bytes| ffi::MPI_Aint::try_from(bytes).ok())
180 .is_none()
181 }) {
182 return Err(Error::SizeOverflow);
183 }
184 let bytes = len.checked_mul(width).ok_or(Error::SizeOverflow)?;
185 let bytes = ffi::MPI_Aint::try_from(bytes).map_err(|_| Error::SizeOverflow)?;
186 let width = i32::try_from(width).map_err(|_| Error::SizeOverflow)?;
187 let mut base: *mut c_void = std::ptr::null_mut();
188 let mut win = unsafe { ffi::RSMPI_WIN_NULL };
190 unsafe {
191 check(ffi::MPI_Win_allocate(
192 bytes,
193 width,
194 ffi::RSMPI_INFO_NULL,
195 comm.as_raw(),
196 &mut base as *mut *mut c_void as *mut c_void,
197 &mut win,
198 ))?;
199 }
200 let model = match unsafe { init_epoch(win, base as *mut T, len) } {
202 Ok(model) => model,
203 Err(error) => {
204 unsafe {
205 let _ = ffi::MPI_Win_free(&mut win);
206 }
207 return Err(error);
208 }
209 };
210 comm.barrier();
211 Ok(Window {
212 win,
213 base: base as *mut T,
214 len,
215 lengths,
216 rank: comm.rank(),
217 ranks: comm.size(),
218 model,
219 closed: false,
220 _element: PhantomData,
221 })
222 }
223
224 pub fn len(&self) -> usize {
226 self.len
227 }
228
229 pub fn is_empty(&self) -> bool {
231 self.len == 0
232 }
233
234 pub fn memory_model(&self) -> MemoryModel {
236 self.model
237 }
238
239 pub fn sync(&self) -> Result<(), Error> {
244 unsafe { check(ffi::MPI_Win_sync(self.win)) }
245 }
246
247 pub fn read_local(&self, disp: usize, out: &mut [T]) -> Result<(), Error> {
254 self.validate(self.rank, disp, out.len())?;
255 if out.is_empty() {
256 return Ok(());
257 }
258 unsafe {
259 std::ptr::copy_nonoverlapping(self.base.add(disp), out.as_mut_ptr(), out.len());
260 }
261 Ok(())
262 }
263
264 pub fn read_local_volatile(&self, disp: usize, out: &mut [T]) -> Result<(), Error> {
266 self.validate(self.rank, disp, out.len())?;
267 for (i, value) in out.iter_mut().enumerate() {
268 unsafe {
269 *value = std::ptr::read_volatile(self.base.add(disp + i));
270 }
271 }
272 Ok(())
273 }
274
275 pub fn put(&self, dest: Rank, disp: usize, data: &[T]) -> Result<(), Error> {
287 self.validate(dest, disp, data.len())?;
288 let count = i32::try_from(data.len()).map_err(|_| Error::CountOverflow)?;
289 let datatype = T::equivalent_datatype();
290 unsafe {
291 check(ffi::MPI_Put(
292 data.as_ptr() as *const c_void,
293 count,
294 datatype.as_raw(),
295 dest,
296 disp as ffi::MPI_Aint,
297 count,
298 datatype.as_raw(),
299 self.win,
300 ))?;
301 check(ffi::MPI_Win_flush(dest, self.win))
302 }
303 }
304
305 pub fn get(&self, source: Rank, disp: usize, out: &mut [T]) -> Result<(), Error> {
314 self.validate(source, disp, out.len())?;
315 let count = i32::try_from(out.len()).map_err(|_| Error::CountOverflow)?;
316 let datatype = T::equivalent_datatype();
317 unsafe {
318 check(ffi::MPI_Get(
319 out.as_mut_ptr() as *mut c_void,
320 count,
321 datatype.as_raw(),
322 source,
323 disp as ffi::MPI_Aint,
324 count,
325 datatype.as_raw(),
326 self.win,
327 ))?;
328 check(ffi::MPI_Win_flush(source, self.win))
329 }
330 }
331
332 pub fn fetch_add(&self, dest: Rank, disp: usize, value: T) -> Result<T, Error> {
339 self.validate(dest, disp, 1)?;
340 let datatype = T::equivalent_datatype();
341 let mut previous = std::mem::MaybeUninit::<T>::uninit();
342 unsafe {
343 check(ffi::MPI_Fetch_and_op(
344 &value as *const T as *const c_void,
345 previous.as_mut_ptr() as *mut c_void,
346 datatype.as_raw(),
347 dest,
348 disp as ffi::MPI_Aint,
349 ffi::RSMPI_SUM,
350 self.win,
351 ))?;
352 check(ffi::MPI_Win_flush(dest, self.win))?;
353 Ok(previous.assume_init())
354 }
355 }
356
357 pub fn close(mut self) -> Result<(), Error> {
362 let result = self.finish();
363 if result.is_err() {
364 self.closed = true;
365 }
366 result
367 }
368
369 fn validate(&self, rank: Rank, start: usize, len: usize) -> Result<(), Error> {
370 if rank < 0 || rank >= self.ranks {
371 return Err(Error::Rank(rank));
372 }
373 let end = start.checked_add(len).ok_or(Error::SizeOverflow)?;
374 let window = self.lengths[rank as usize];
375 if end > window {
376 return Err(Error::Range { start, len, window });
377 }
378 Ok(())
379 }
380
381 fn finish(&mut self) -> Result<(), Error> {
382 if self.closed {
383 return Ok(());
384 }
385 unsafe {
386 check(ffi::MPI_Win_unlock_all(self.win))?;
387 check(ffi::MPI_Win_free(&mut self.win))?;
388 }
389 self.closed = true;
390 Ok(())
391 }
392}
393
394impl<T: RmaElement> Drop for Window<T> {
395 fn drop(&mut self) {
396 let _ = self.finish();
399 }
400}