use core_storage::fs::Fs;
use core_storage::Result;
use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use crate::db::GraphDb;
#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
pub enum MaskMode {
#[default]
Omit,
Stub,
}
#[derive(Clone, Debug)]
pub struct NodeMask {
pub(crate) visible: HashSet<u32>,
mode: MaskMode,
}
impl NodeMask {
pub fn from_keys<'a, F: Fs>(db: &GraphDb<F>, keys: impl IntoIterator<Item = &'a str>) -> Self {
let visible = keys.into_iter().filter_map(|k| db.ids().get(k)).collect();
NodeMask {
visible,
mode: MaskMode::default(),
}
}
pub fn from_ids(ids: impl IntoIterator<Item = u32>) -> Self {
NodeMask {
visible: ids.into_iter().collect(),
mode: MaskMode::default(),
}
}
pub fn with_mode(self, mode: MaskMode) -> Self {
NodeMask { mode, ..self }
}
pub fn mode(&self) -> MaskMode {
self.mode
}
pub fn len(&self) -> usize {
self.visible.len()
}
pub fn is_empty(&self) -> bool {
self.visible.is_empty()
}
pub fn intersect(&self, other: &NodeMask) -> NodeMask {
NodeMask {
visible: self.visible.intersection(&other.visible).copied().collect(),
mode: MaskMode::Omit,
}
}
pub fn contains_id(&self, id: u32) -> bool {
self.visible.contains(&id)
}
pub fn contains_node<F: core_storage::fs::Fs>(
&self,
db: &crate::db::GraphDb<F>,
key: &str,
) -> bool {
db.ids()
.get(key)
.is_some_and(|id| self.visible.contains(&id))
}
}
#[derive(Default)]
pub struct RoleMaskCache {
entries: Mutex<HashMap<String, (u64, Arc<NodeMask>)>>,
}
impl RoleMaskCache {
pub fn new() -> Self {
Self::default()
}
pub fn get_or_build(
&self,
role: &str,
version: u64,
build: impl FnOnce() -> Result<NodeMask>,
) -> Result<Arc<NodeMask>> {
if let Ok(entries) = self.entries.lock() {
if let Some((v, mask)) = entries.get(role) {
if *v == version {
return Ok(Arc::clone(mask));
}
}
}
let mask = Arc::new(build()?);
if let Ok(mut entries) = self.entries.lock() {
entries.insert(role.to_string(), (version, Arc::clone(&mask)));
}
Ok(mask)
}
pub fn clear(&self) {
if let Ok(mut entries) = self.entries.lock() {
entries.clear();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_version_change_rebuilds_and_clear_empties() {
let cache = RoleMaskCache::new();
let built = std::cell::Cell::new(0u32);
let build = |ids: Vec<u32>| {
built.set(built.get() + 1);
Ok(NodeMask::from_ids(ids))
};
let m = cache.get_or_build("r", 1, || build(vec![1])).unwrap();
assert_eq!(m.len(), 1);
assert_eq!(built.get(), 1);
let m = cache.get_or_build("r", 1, || build(vec![1, 2])).unwrap();
assert_eq!(m.len(), 1, "the memoised mask is returned unchanged");
assert_eq!(built.get(), 1);
let m = cache.get_or_build("r", 2, || build(vec![1, 2])).unwrap();
assert_eq!(m.len(), 2);
assert_eq!(built.get(), 2);
let m = cache.get_or_build("other", 2, || build(vec![9])).unwrap();
assert_eq!(m.len(), 1);
assert_eq!(built.get(), 3);
cache.clear();
let _ = cache.get_or_build("r", 2, || build(vec![1, 2])).unwrap();
assert_eq!(built.get(), 4, "clear drops the entry, so it rebuilds");
}
#[test]
fn a_failed_build_is_not_cached() {
let cache = RoleMaskCache::new();
assert!(cache
.get_or_build("r", 1, || Err(core_storage::GraphError::KeyNotFound {
key: "role:r".into()
}))
.is_err());
let m = cache
.get_or_build("r", 1, || Ok(NodeMask::from_ids(vec![7])))
.unwrap();
assert_eq!(m.len(), 1);
}
}