Skip to main content

mpi_rma/
window.rs

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
16/// Plain scalar values whose in-memory representation is safe for RMA.
17///
18/// The trait is sealed: arbitrary `Equivalence` implementations may describe
19/// datatypes whose extent exceeds the Rust object and cannot safely back a
20/// contiguous window.
21pub 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
64/// Enter the passive-target epoch for `win`, zero its storage, and report
65/// the memory model. The caller owns the handle and must free it on error.
66///
67/// # Safety
68/// `win` must be a freshly allocated handle owning `base` for `len`
69/// elements of `T`, with no access epoch started.
70unsafe 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    // NOCHECK: no conflicting lock can exist yet, since all ranks are still
95    // inside the collective constructor.
96    unsafe {
97        check(ffi::MPI_Win_lock_all(ffi::MPI_MODE_NOCHECK as i32, win))?;
98    }
99    // RmaElement is sealed to numeric scalars, for which all-zero is a valid
100    // value. Win_sync publishes it under the separate model.
101    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
110/// Method-style RMA extension for every rsmpi communicator.
111pub 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
119/// MPI-allocated homogeneous memory exposed to the communicator.
120///
121/// `Send + Sync` once constructed. Construction requires
122/// `MPI_THREAD_MULTIPLE`; the access epoch is entered inside `allocate` and
123/// left in [`close`](Self::close) or `Drop`, and no public API exposes
124/// epoch transitions. All reads and writes are issued through the
125/// `put`/`get`/`fetch_add` family, each completed at the target before
126/// returning; no Rust reference into remotely mutable storage is ever
127/// exposed.
128pub 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
140// SAFETY: the only constructor rejects anything below MPI_THREAD_MULTIPLE.
141// Storage is reached only through MPI calls, never via Rust references.
142unsafe 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        // SAFETY: read-only MPI null handle used as an out-parameter seed.
189        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        // The window handle is live from here on; any failure must free it.
201        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    /// Number of locally allocated elements.
225    pub fn len(&self) -> usize {
226        self.len
227    }
228
229    /// Whether this rank allocated an empty window.
230    pub fn is_empty(&self) -> bool {
231        self.len == 0
232    }
233
234    /// Memory model reported by `MPI_WIN_MODEL`.
235    pub fn memory_model(&self) -> MemoryModel {
236        self.model
237    }
238
239    /// Make remote updates visible to local memory accesses.
240    ///
241    /// Required after a remote put on a separate-model window before any
242    /// local read can observe the new contents.
243    pub fn sync(&self) -> Result<(), Error> {
244        unsafe { check(ffi::MPI_Win_sync(self.win)) }
245    }
246
247    /// Copy a region from this rank's local window storage.
248    ///
249    /// Concurrent remote updates to the same region yield undefined
250    /// contents; the caller is responsible for the ordering. Volatile
251    /// reads of single scalars are provided by
252    /// [`read_local_volatile`](Self::read_local_volatile).
253    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    /// Volatile local copy for scalars concurrently mutated by an RMA operation.
265    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    /// Put a contiguous region and complete the transfer at the target before return.
276    ///
277    /// The transfer reaches remote completion (target's public window
278    /// visible to subsequent `get` from any process) before this returns.
279    ///
280    /// # Errors
281    /// - [`Error::Rank`] if `dest` is outside the window communicator.
282    /// - [`Error::Range`] if `disp + data.len()` exceeds the target's
283    ///   window length.
284    /// - [`Error::CountOverflow`] if `data.len()` does not fit an `MPI Count`.
285    /// - [`Error::Mpi`] if the MPI call fails.
286    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    /// Get a contiguous region and complete the transfer before return.
306    ///
307    /// # Errors
308    /// - [`Error::Rank`] if `source` is outside the window communicator.
309    /// - [`Error::Range`] if `disp + out.len()` exceeds the source's
310    ///   window length.
311    /// - [`Error::CountOverflow`] if `out.len()` does not fit an `MPI Count`.
312    /// - [`Error::Mpi`] if the MPI call fails.
313    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    /// Atomically add `value` at the target and return the previous value.
333    ///
334    /// # Errors
335    /// - [`Error::Rank`] if `dest` is outside the window communicator.
336    /// - [`Error::Range`] if `disp` is outside the target's window length.
337    /// - [`Error::Mpi`] if the MPI call fails.
338    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    /// Close the window. Collective over the window group.
358    ///
359    /// Prefer this over relying on `Drop`: the destructor's shutdown runs
360    /// the same steps but the collective boundary is implicit.
361    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        // MPI_Win_free is collective. Well-structured MPI programs drop
397        // windows symmetrically; explicit `close` makes that boundary visible.
398        let _ = self.finish();
399    }
400}