use crate::debug_invariants::DebugInvariants;
use crate::mesh_error::MeshSieveError;
use crate::topology::cache::InvalidateCache;
use crate::topology::point::PointId;
use std::collections::HashMap;
#[derive(Clone, Debug, Default, serde::Serialize, serde::Deserialize)]
pub struct Atlas {
map: HashMap<PointId, (usize, usize)>,
order: Vec<PointId>,
total_len: usize,
version: u64,
}
impl InvalidateCache for Atlas {
fn invalidate_cache(&mut self) {
}
}
impl Atlas {
pub fn try_insert(&mut self, p: PointId, len: usize) -> Result<usize, MeshSieveError> {
if len == 0 {
return Err(MeshSieveError::ZeroLengthSlice);
}
if self.map.contains_key(&p) {
return Err(MeshSieveError::DuplicatePoint(p));
}
let offset = self.total_len;
self.map.insert(p, (offset, len));
self.order.push(p);
self.total_len += len;
self.version = self.version.wrapping_add(1);
InvalidateCache::invalidate_cache(self);
#[cfg(any(
debug_assertions,
feature = "strict-invariants",
feature = "check-invariants"
))]
self.debug_assert_invariants();
Ok(offset)
}
#[inline]
pub fn get(&self, p: PointId) -> Option<(usize, usize)> {
self.map.get(&p).copied()
}
#[inline]
pub fn contains(&self, p: PointId) -> bool {
self.map.contains_key(&p)
}
#[inline]
pub fn len(&self) -> usize {
debug_assert_eq!(self.order.len(), self.map.len());
self.order.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
debug_assert_eq!(self.order.is_empty(), self.map.is_empty());
self.order.is_empty()
}
#[inline]
pub fn total_len(&self) -> usize {
self.total_len
}
#[inline]
pub fn version(&self) -> u64 {
self.version
}
#[inline]
pub fn atlas_map(&self) -> Vec<(usize, usize)> {
self.order.iter().map(|&p| self.map[&p]).collect()
}
#[inline]
pub fn atlas_entries(&self) -> Vec<(PointId, (usize, usize))> {
self.order.iter().map(|&p| (p, self.map[&p])).collect()
}
pub fn iter_entries<'a>(&'a self) -> impl Iterator<Item = (PointId, (usize, usize))> + 'a {
self.order.iter().map(move |&p| (p, self.map[&p]))
}
pub fn iter_spans<'a>(&'a self) -> impl Iterator<Item = (usize, usize)> + 'a {
self.order.iter().map(move |&p| self.map[&p])
}
#[inline]
pub fn points<'a>(&'a self) -> impl Iterator<Item = PointId> + 'a {
self.order.iter().copied()
}
pub fn build_scatter_plan(&self) -> crate::data::section::ScatterPlan {
crate::data::section::ScatterPlan {
atlas_version: self.version,
spans: self.atlas_map(),
}
}
pub fn remove_point(&mut self, p: PointId) -> Result<(), MeshSieveError> {
let existed = self.map.remove(&p).is_some();
if !existed {
return Err(MeshSieveError::MissingAtlasPoint(p));
}
self.order.retain(|&x| x != p);
let mut next_offset = 0usize;
let mut new_offsets = Vec::with_capacity(self.order.len());
for &pt in &self.order {
let len = match self.map.get(&pt) {
Some(&(_, len)) => len,
None => {
return Err(MeshSieveError::MissingAtlasPoint(pt));
}
};
new_offsets.push((pt, next_offset, len));
next_offset += len;
}
for (pt, off, len) in new_offsets {
self.map.insert(pt, (off, len));
}
self.total_len = next_offset;
self.version = self.version.wrapping_add(1);
InvalidateCache::invalidate_cache(self);
#[cfg(any(
debug_assertions,
feature = "strict-invariants",
feature = "check-invariants"
))]
self.debug_assert_invariants();
Ok(())
}
}
impl DebugInvariants for Atlas {
fn debug_assert_invariants(&self) {
crate::debug_invariants!(self.validate_invariants(), "Atlas invalid");
}
fn validate_invariants(&self) -> Result<(), MeshSieveError> {
use std::collections::HashSet;
let set: HashSet<_> = self.order.iter().copied().collect();
if set.len() != self.order.len() {
let mut seen = HashSet::new();
let dup = self
.order
.iter()
.copied()
.find(|p| !seen.insert(*p))
.unwrap();
return Err(MeshSieveError::DuplicatePoint(dup));
}
if let Some(&p) = self.order.iter().find(|&&p| !self.map.contains_key(&p)) {
return Err(MeshSieveError::MissingAtlasPoint(p));
}
if let Some(&p) = self.map.keys().find(|p| !set.contains(p)) {
return Err(MeshSieveError::DuplicatePoint(p));
}
for &p in &self.order {
let (off, len) = self.map[&p];
if len == 0 {
return Err(MeshSieveError::ZeroLengthSlice);
}
let _end = off
.checked_add(len)
.ok_or_else(|| MeshSieveError::ScatterChunkMismatch { offset: off, len })?;
}
let mut expected_off = 0usize;
let mut sum = 0usize;
for &p in &self.order {
let (off, len) = self.map[&p];
if off != expected_off {
return Err(MeshSieveError::AtlasContiguityMismatch {
point: p,
expected: expected_off,
found: off,
});
}
expected_off = off + len; sum = sum
.checked_add(len)
.ok_or_else(|| MeshSieveError::ScatterLengthMismatch {
expected: usize::MAX,
found: 0,
})?;
}
if sum != self.total_len {
return Err(MeshSieveError::ScatterLengthMismatch {
expected: sum,
found: self.total_len,
});
}
Ok(())
}
}
#[cfg(test)]
impl Atlas {
pub fn force_offset(&mut self, p: PointId, new_off: usize) {
if let Some((_, len)) = self.map.get(&p).copied() {
self.map.insert(p, (new_off, len));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::topology::point::PointId;
#[test]
fn insert_and_lookup() {
let mut a = Atlas::default();
let p1 = PointId::new(1).unwrap();
let off1 = a.try_insert(p1, 3);
assert_eq!(off1.unwrap(), 0);
let p2 = PointId::new(2).unwrap();
let off2 = a.try_insert(p2, 5);
assert_eq!(off2.unwrap(), 3);
assert_eq!(a.get(p1), Some((0, 3)));
assert_eq!(a.get(p2), Some((3, 5)));
assert_eq!(a.total_len(), 8);
assert_eq!(a.points().collect::<Vec<_>>(), vec![p1, p2]);
}
#[test]
fn zero_len_rejected() {
let mut a = Atlas::default();
assert_eq!(
a.try_insert(PointId::new(7).unwrap(), 0).unwrap_err(),
MeshSieveError::ZeroLengthSlice
);
}
#[test]
fn atlas_cache_cleared_on_insert() {
use crate::topology::cache::InvalidateCache;
use crate::topology::point::PointId;
let mut atlas = Atlas::default();
let _ = atlas.try_insert(PointId::new(1).unwrap(), 2);
InvalidateCache::invalidate_cache(&mut atlas); let _ = atlas.try_insert(PointId::new(2).unwrap(), 1);
assert_eq!(atlas.get(PointId::new(1).unwrap()), Some((0, 2)));
assert_eq!(atlas.get(PointId::new(2).unwrap()), Some((2, 1)));
}
#[test]
fn remove_point_recomputes_offsets() {
let mut a = Atlas::default();
let p1 = PointId::new(1).unwrap();
let p2 = PointId::new(2).unwrap();
let p3 = PointId::new(3).unwrap();
let _ = a.try_insert(p1, 3);
let _ = a.try_insert(p2, 5);
let _ = a.try_insert(p3, 2);
a.remove_point(p2).unwrap();
assert_eq!(a.get(p1), Some((0, 3)));
assert_eq!(a.get(p3), Some((3, 2)));
assert_eq!(a.total_len(), 5);
assert_eq!(a.points().collect::<Vec<_>>(), vec![p1, p3]);
}
#[test]
fn duplicate_insert_panics() {
let mut a = Atlas::default();
let p = PointId::new(42).unwrap();
let _ = a.try_insert(p, 1);
assert_eq!(a.try_insert(p, 2), Err(MeshSieveError::DuplicatePoint(p)));
}
#[test]
fn get_missing_point_returns_none() {
let a = Atlas::default();
let p = PointId::new(99).unwrap();
assert_eq!(a.get(p), None);
assert_eq!(a.total_len(), 0);
assert!(a.points().next().is_none());
}
#[test]
fn remove_first_and_last_points() {
let mut a = Atlas::default();
let p1 = PointId::new(1).unwrap();
let p2 = PointId::new(2).unwrap();
let p3 = PointId::new(3).unwrap();
a.try_insert(p1, 2).unwrap();
a.try_insert(p2, 4).unwrap();
a.try_insert(p3, 1).unwrap();
a.remove_point(p1).unwrap();
assert_eq!(a.points().collect::<Vec<_>>(), vec![p2, p3]);
assert_eq!(a.get(p2), Some((0, 4)));
assert_eq!(a.get(p3), Some((4, 1)));
a.remove_point(p3).unwrap();
assert_eq!(a.points().collect::<Vec<_>>(), vec![p2]);
assert_eq!(a.get(p2), Some((0, 4)));
assert_eq!(a.total_len(), 4);
}
#[test]
fn clear_all_then_reinsert() {
let mut a = Atlas::default();
let pts = [PointId::new(10).unwrap(), PointId::new(20).unwrap()];
for &p in &pts {
a.try_insert(p, 5).unwrap();
}
for &p in &pts {
a.remove_point(p).unwrap();
}
assert!(a.points().next().is_none());
assert_eq!(a.total_len(), 0);
let off = a.try_insert(PointId::new(30).unwrap(), 7).unwrap();
assert_eq!(off, 0);
assert_eq!(
a.points().collect::<Vec<_>>(),
vec![PointId::new(30).unwrap()]
);
}
#[test]
fn serde_roundtrip() {
let mut a = Atlas::default();
a.try_insert(PointId::new(5).unwrap(), 3).unwrap();
a.try_insert(PointId::new(6).unwrap(), 2).unwrap();
let ser = serde_json::to_string(&a).expect("serialize");
let de: Atlas = serde_json::from_str(&ser).expect("deserialize");
assert_eq!(de.get(PointId::new(5).unwrap()), Some((0, 3)));
assert_eq!(de.get(PointId::new(6).unwrap()), Some((3, 2)));
assert_eq!(
de.points().collect::<Vec<_>>(),
vec![PointId::new(5).unwrap(), PointId::new(6).unwrap()]
);
}
fn pid(id: u64) -> PointId {
PointId::new(id).unwrap()
}
#[test]
fn validate_fails_when_order_missing_map_key() {
let mut a = Atlas::default();
let p1 = pid(1);
let p2 = pid(2);
a.try_insert(p1, 1).unwrap();
a.try_insert(p2, 2).unwrap();
a.order.retain(|&x| x != p2);
let e = a.validate_invariants().unwrap_err();
assert!(matches!(e, MeshSieveError::DuplicatePoint(pp) if pp == p2));
}
#[test]
fn validate_fails_when_map_missing_order_key() {
let mut a = Atlas::default();
let p1 = pid(1);
a.try_insert(p1, 3).unwrap();
a.map.remove(&p1);
let e = a.validate_invariants().unwrap_err();
assert!(matches!(e, MeshSieveError::MissingAtlasPoint(pp) if pp == p1));
}
#[test]
fn remove_point_errors_if_absent() {
let mut a = Atlas::default();
let p1 = pid(1);
let p2 = pid(2);
a.try_insert(p1, 1).unwrap();
let e = a.remove_point(p2).unwrap_err();
assert!(matches!(e, MeshSieveError::MissingAtlasPoint(pp) if pp == p2));
a.remove_point(p1).unwrap();
assert!(a.is_empty());
assert_eq!(a.total_len(), 0);
}
}