use std::collections::{BTreeSet, HashMap, HashSet};
use crate::algs::wire::{WireOrientedArrow, cast_slice, cast_slice_mut};
use bytemuck::{Pod, Zeroable};
use crate::algs::communicator::{CommTag, Communicator, SieveCommTags, Wait};
use crate::algs::completion::size_exchange::exchange_sizes_symmetric;
use crate::mesh_error::MeshSieveError;
use crate::overlap::overlap::Overlap;
use crate::topology::cache::InvalidateCache;
use crate::topology::point::PointId;
use crate::topology::sieve::OrientedSieve;
pub trait CompletionPayload: Copy + Clone + PartialEq + Default + Send + 'static {
type Wire: Copy + Pod + Zeroable + PartialEq + Send + 'static;
fn encode(self) -> Self::Wire;
fn decode(wire: Self::Wire) -> Result<Self, MeshSieveError>;
}
impl<T> CompletionPayload for T
where
T: Copy + Clone + PartialEq + Default + Pod + Zeroable + Send + 'static,
{
type Wire = T;
fn encode(self) -> Self::Wire {
self
}
fn decode(wire: Self::Wire) -> Result<Self, MeshSieveError> {
Ok(wire)
}
}
#[repr(C)]
#[derive(Copy, Clone, PartialEq, Pod, Zeroable)]
pub struct WireRemotePayload {
rank_le: u64,
remote_le: u64,
has_remote: u64,
}
impl CompletionPayload for crate::overlap::overlap::Remote {
type Wire = WireRemotePayload;
fn encode(self) -> Self::Wire {
WireRemotePayload {
rank_le: (self.rank as u64).to_le(),
remote_le: self.remote_point.map_or(0, |p| p.get()).to_le(),
has_remote: u64::from(self.remote_point.is_some()),
}
}
fn decode(wire: Self::Wire) -> Result<Self, MeshSieveError> {
let rank = usize::try_from(u64::from_le(wire.rank_le)).map_err(|_| {
MeshSieveError::InvalidGeometry("remote rank does not fit usize".into())
})?;
let remote_point = if wire.has_remote != 0 {
Some(crate::topology::point::PointId::new(u64::from_le(
wire.remote_le,
))?)
} else {
None
};
Ok(Self { rank, remote_point })
}
}
fn remote_id_for(overlap: &Overlap, nbr: usize, local: PointId) -> Result<PointId, MeshSieveError> {
overlap
.links_to(nbr)
.find(|(p, _)| *p == local)
.and_then(|(_, rp)| rp)
.ok_or_else(|| MeshSieveError::MissingOverlap {
source: format!(
"Unresolved mapping for local {} to neighbor {}",
local.get(),
nbr
)
.into(),
})
}
fn build_wires<S>(
mesh: &S,
overlap: &Overlap,
neighbors: &[usize],
) -> Result<
HashMap<usize, Vec<WireOrientedArrow<<S::Payload as CompletionPayload>::Wire, S::Orient>>>,
MeshSieveError,
>
where
S: OrientedSieve<Point = PointId>,
S::Payload: CompletionPayload,
S::Orient: Copy + Pod + Zeroable,
{
let mut srcs_per_nbr: HashMap<usize, Vec<PointId>> = HashMap::new();
for &nbr in neighbors {
let mut srcs: Vec<PointId> = overlap.links_to(nbr).map(|(p, _)| p).collect();
srcs.sort_unstable();
srcs.dedup();
srcs_per_nbr.insert(nbr, srcs);
}
let mut est_cap: HashMap<usize, usize> = HashMap::new();
for (&nbr, srcs) in &srcs_per_nbr {
let mut sum = 0usize;
for &s in srcs {
sum += mesh.cone_points(s).count();
}
est_cap.insert(nbr, sum);
}
let mut wires: HashMap<
usize,
Vec<WireOrientedArrow<<S::Payload as CompletionPayload>::Wire, S::Orient>>,
> = HashMap::new();
for (&nbr, srcs) in &srcs_per_nbr {
let mut buf = Vec::with_capacity(*est_cap.get(&nbr).unwrap_or(&0));
for &src_local in srcs {
let src_remote = remote_id_for(overlap, nbr, src_local)?;
let payloads: std::collections::BTreeMap<_, _> = mesh.cone(src_local).collect();
let mut arrows: Vec<_> = mesh
.cone_o(src_local)
.filter_map(|(dst, orient)| {
payloads
.get(&dst)
.cloned()
.map(|payload| (dst, payload, orient))
})
.collect();
arrows.sort_unstable_by_key(|(dst, _, _)| *dst);
let mut unique_arrows: Vec<(PointId, S::Payload, S::Orient)> =
Vec::with_capacity(arrows.len());
for arrow in arrows {
if let Some(previous) = unique_arrows.last()
&& previous.0 == arrow.0
{
if previous.1 != arrow.1 || previous.2 != arrow.2 {
let kind = match (previous.1 == arrow.1, previous.2 == arrow.2) {
(false, false) => {
crate::mesh_error::RelationConflictKind::PayloadAndOrientation
}
(false, true) => crate::mesh_error::RelationConflictKind::Payload,
(true, false) => crate::mesh_error::RelationConflictKind::Orientation,
(true, true) => unreachable!(),
};
return Err(MeshSieveError::RelationConflict {
src: format!("PointId({})", src_local.get()),
dst: format!("PointId({})", arrow.0.get()),
kind,
});
}
continue;
}
unique_arrows.push(arrow);
}
let arrows = unique_arrows;
for (dst_local, payload, orient) in arrows {
let dst_remote = remote_id_for(overlap, nbr, dst_local)?;
buf.push(WireOrientedArrow::new(
src_remote.get(),
dst_remote.get(),
payload.encode(),
orient,
));
}
}
buf.sort_unstable_by_key(|w| (w.src(), w.dst()));
let mut unique: Vec<WireOrientedArrow<<S::Payload as CompletionPayload>::Wire, S::Orient>> =
Vec::with_capacity(buf.len());
for wire in buf {
if let Some(previous) = unique.last()
&& previous.src() == wire.src()
&& previous.dst() == wire.dst()
{
if previous.payload != wire.payload || previous.orientation != wire.orientation {
return Err(MeshSieveError::RelationConflict {
src: format!("PointId({})", wire.src()),
dst: format!("PointId({})", wire.dst()),
kind: match (
previous.payload == wire.payload,
previous.orientation == wire.orientation,
) {
(false, false) => {
crate::mesh_error::RelationConflictKind::PayloadAndOrientation
}
(false, true) => crate::mesh_error::RelationConflictKind::Payload,
(true, false) => crate::mesh_error::RelationConflictKind::Orientation,
(true, true) => unreachable!(),
},
});
}
continue;
}
unique.push(wire);
}
let buf = unique;
wires.insert(nbr, buf);
}
Ok(wires)
}
pub fn complete_sieve_with_tags<S, C>(
mesh: &mut S,
overlap: &Overlap,
comm: &C,
my_rank: usize,
tags: SieveCommTags,
) -> Result<(), MeshSieveError>
where
S: OrientedSieve<Point = PointId> + InvalidateCache,
S::Payload: CompletionPayload,
S::Orient: Copy + Pod + Zeroable + Send + 'static,
C: Communicator + Sync,
{
#[cfg(any(
debug_assertions,
feature = "strict-invariants",
feature = "check-invariants"
))]
overlap.validate_invariants()?;
if comm.is_no_comm() || comm.size() <= 1 {
mesh.invalidate_cache();
return Ok(());
}
let mut nb: BTreeSet<usize> = overlap.neighbor_ranks().collect();
nb.remove(&my_rank);
let neighbors: Vec<usize> = nb.into_iter().collect();
if neighbors.is_empty() {
mesh.invalidate_cache();
return Ok(());
}
let all_neighbors: HashSet<usize> = neighbors.iter().copied().collect();
let wires = build_wires(mesh, overlap, &neighbors)?;
let counts = exchange_sizes_symmetric(&wires, comm, tags.sizes, &all_neighbors)?;
let mut recvs = Vec::new();
for &nbr in &neighbors {
let n = counts.get(&nbr).copied().unwrap_or(0) as usize;
let mut buf = vec![
WireOrientedArrow::<<S::Payload as CompletionPayload>::Wire, S::Orient>::zeroed();
n
];
let h = comm.irecv_result(nbr, tags.data.as_u16(), cast_slice_mut(&mut buf))?;
recvs.push((nbr, h, buf));
}
let mut sends = Vec::new();
for &nbr in &neighbors {
let out = wires.get(&nbr).map_or(&[][..], |v| &v[..]);
sends.push(comm.isend_result(nbr, tags.data.as_u16(), cast_slice(out))?);
}
let mut maybe_err: Option<MeshSieveError> = None;
for (nbr, h, mut buf) in recvs {
match h.wait() {
Some(raw)
if raw.len()
== buf.len()
* std::mem::size_of::<
WireOrientedArrow<<S::Payload as CompletionPayload>::Wire, S::Orient>,
>() =>
{
cast_slice_mut(&mut buf).copy_from_slice(&raw);
for w in &buf {
let src = PointId::new(w.src())
.map_err(|e| MeshSieveError::MeshError(Box::new(e)))?;
let dst = PointId::new(w.dst())
.map_err(|e| MeshSieveError::MeshError(Box::new(e)))?;
let payload = S::Payload::decode(w.payload)?;
if let Err(error) = mesh.add_arrow_o(src, dst, payload, w.orientation)
&& maybe_err.is_none()
{
maybe_err = Some(error);
}
}
}
Some(raw) if maybe_err.is_none() => {
let exp = buf.len()
* std::mem::size_of::<
WireOrientedArrow<<S::Payload as CompletionPayload>::Wire, S::Orient>,
>();
maybe_err = Some(MeshSieveError::CommError {
neighbor: nbr,
source: format!("payload size mismatch: expected {exp}B, got {}B", raw.len())
.into(),
});
}
None if maybe_err.is_none() => {
maybe_err = Some(MeshSieveError::CommError {
neighbor: nbr,
source: "recv returned None".into(),
});
}
_ => {}
}
}
for h in sends {
let _ = h.wait();
}
mesh.invalidate_cache();
if let Some(e) = maybe_err {
Err(e)
} else {
Ok(())
}
}
pub fn complete_sieve<S, C>(
mesh: &mut S,
overlap: &Overlap,
comm: &C,
my_rank: usize,
) -> Result<(), MeshSieveError>
where
S: OrientedSieve<Point = PointId> + InvalidateCache,
S::Payload: CompletionPayload,
S::Orient: Copy + Pod + Zeroable + Send + 'static,
C: Communicator + Sync,
{
let tags = SieveCommTags::from_base(CommTag::new(0xC0DE));
complete_sieve_with_tags(mesh, overlap, comm, my_rank, tags)
}
pub fn complete_sieve_until_converged<S, C>(
sieve: &mut S,
overlap: &Overlap,
comm: &C,
my_rank: usize,
) -> Result<(), MeshSieveError>
where
S: OrientedSieve<Point = PointId> + InvalidateCache,
S::Payload: CompletionPayload,
S::Orient: Copy + Pod + Zeroable + Send + 'static,
C: Communicator + Sync,
{
let mut prev = None;
loop {
let before = oriented_relation_snapshot(sieve);
complete_sieve(sieve, overlap, comm, my_rank)?;
let after = oriented_relation_snapshot(sieve);
if after == before || prev.as_ref() == Some(&after) {
break;
}
prev = Some(after);
sieve.invalidate_cache();
}
Ok(())
}
fn oriented_relation_snapshot<S>(sieve: &S) -> Vec<(PointId, PointId, S::Payload, S::Orient)>
where
S: OrientedSieve<Point = PointId>,
S::Payload: CompletionPayload,
S::Orient: Copy,
{
let mut relations = Vec::new();
for src in sieve.points() {
let payloads: std::collections::BTreeMap<_, _> = sieve.cone(src).collect();
for (dst, orient) in sieve.cone_o(src) {
if let Some(payload) = payloads.get(&dst) {
relations.push((src, dst, payload.clone(), orient));
}
}
}
relations.sort_unstable_by_key(|(src, dst, _, _)| (*src, *dst));
relations
}
#[cfg(test)]
mod tests {
use super::*;
use crate::algs::communicator::Communicator;
use crate::overlap::overlap::Overlap;
use crate::topology::sieve::InMemorySieve;
#[test]
fn unresolved_mapping_errors() {
struct DummyComm;
impl Communicator for DummyComm {
type SendHandle = ();
type RecvHandle = ();
fn isend(&self, _peer: usize, _tag: u16, _buf: &[u8]) -> Self::SendHandle {}
fn irecv(&self, _peer: usize, _tag: u16, _buf: &mut [u8]) -> Self::RecvHandle {}
fn rank(&self) -> usize {
0
}
fn size(&self) -> usize {
2
}
}
let mut sieve: InMemorySieve<PointId, ()> = InMemorySieve::default();
let mut ovlp = Overlap::new();
ovlp.add_link_structural_one(PointId::new(1).unwrap(), 1); let comm = DummyComm;
let tags = SieveCommTags::from_base(CommTag::new(0x5100));
let res = complete_sieve_with_tags(&mut sieve, &ovlp, &comm, 0, tags);
assert!(matches!(res, Err(MeshSieveError::MissingOverlap { .. })));
}
}