use std::sync::{Arc, OnceLock};
use crate::persistent_artrie::core::key_encoding::KeyEncoding;
use crate::persistent_artrie::core::overlay::node::{Child, OverlayNode};
use crate::persistent_artrie::core::overlay::OverlayFaulter;
use crate::value::DictionaryValue;
use crate::{DictionaryNode, MappedDictionaryNode, SnapshotTraversalCursor};
struct OverlaySnapshotCursorArena<K: KeyEncoding, V: DictionaryValue> {
nodes: Vec<OverlaySnapshotCursorNode<K, V>>,
edges: Box<[(K::Token, usize)]>,
}
struct OverlaySnapshotCursorNode<K: KeyEncoding, V: DictionaryValue> {
overlay: Arc<OverlayNode<K, V>>,
edge_start: usize,
edge_len: usize,
parent: Option<(usize, K::Token)>,
is_final: bool,
value: Option<V>,
}
impl<K: KeyEncoding, V: DictionaryValue> OverlaySnapshotCursorArena<K, V> {
fn build(
root: Arc<OverlayNode<K, V>>,
overlay_faulter: &Option<Arc<dyn OverlayFaulter<K, V>>>,
) -> Self {
let mut nodes = vec![OverlaySnapshotCursorNode {
is_final: root.is_final(),
value: root.get_value(),
overlay: root,
edge_start: 0,
edge_len: 0,
parent: None,
}];
let mut node_edges = vec![Vec::new()];
let mut node_index = 0usize;
while node_index < nodes.len() {
let overlay = Arc::clone(&nodes[node_index].overlay);
let mut edges = Vec::with_capacity(overlay.num_children());
for (&unit, child) in overlay.iter_children() {
let Some(label) = K::unit_to_token(unit) else {
continue;
};
let Some(child_overlay) =
OverlayDictionaryNode::<K, V>::resolve_overlay_child(child, overlay_faulter)
else {
continue;
};
let child_index = nodes.len();
nodes.push(OverlaySnapshotCursorNode {
is_final: child_overlay.is_final(),
value: child_overlay.get_value(),
overlay: child_overlay,
edge_start: 0,
edge_len: 0,
parent: Some((node_index, label)),
});
node_edges.push(Vec::new());
edges.push((label, child_index));
}
node_edges[node_index] = edges;
node_index += 1;
}
let edge_count = node_edges.iter().map(Vec::len).sum();
let mut flat_edges = Vec::with_capacity(edge_count);
for (node, edges) in nodes.iter_mut().zip(node_edges) {
node.edge_start = flat_edges.len();
node.edge_len = edges.len();
flat_edges.extend(edges);
}
Self {
nodes,
edges: flat_edges.into_boxed_slice(),
}
}
#[inline]
fn node(&self, cursor: SnapshotTraversalCursor) -> Option<&OverlaySnapshotCursorNode<K, V>> {
self.nodes.get(cursor.index())
}
#[inline]
fn node_edges(&self, node: &OverlaySnapshotCursorNode<K, V>) -> &[(K::Token, usize)] {
&self.edges[node.edge_start..node.edge_start + node.edge_len]
}
}
#[derive(Clone)]
pub struct OverlayDictionaryNode<K: KeyEncoding, V: DictionaryValue = ()> {
overlay: Option<Arc<OverlayNode<K, V>>>,
overlay_faulter: Option<Arc<dyn OverlayFaulter<K, V>>>,
snapshot_cursor_arena: Arc<OnceLock<OverlaySnapshotCursorArena<K, V>>>,
}
impl<K: KeyEncoding, V: DictionaryValue> std::fmt::Debug for OverlayDictionaryNode<K, V> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OverlayDictionaryNode")
.field("overlay", &self.overlay)
.field("has_faulter", &self.overlay_faulter.is_some())
.field(
"has_snapshot_cursor_arena",
&self.snapshot_cursor_arena.get().is_some(),
)
.finish()
}
}
impl<K: KeyEncoding, V: DictionaryValue> OverlayDictionaryNode<K, V> {
pub(crate) fn from_overlay_root(
node: Arc<OverlayNode<K, V>>,
overlay_faulter: Option<Arc<dyn OverlayFaulter<K, V>>>,
) -> Self {
Self {
overlay: Some(node),
overlay_faulter,
snapshot_cursor_arena: Arc::new(OnceLock::new()),
}
}
pub(crate) fn from_overlay_node(
node: Arc<OverlayNode<K, V>>,
overlay_faulter: Option<Arc<dyn OverlayFaulter<K, V>>>,
) -> Self {
Self {
overlay: Some(node),
overlay_faulter,
snapshot_cursor_arena: Arc::new(OnceLock::new()),
}
}
#[inline]
fn resolve_overlay_child(
child: &Child<K, V>,
overlay_faulter: &Option<Arc<dyn OverlayFaulter<K, V>>>,
) -> Option<Arc<OverlayNode<K, V>>> {
if let Some(child_arc) = child.as_in_mem() {
return Some(Arc::clone(child_arc));
}
let on_disk = child.as_on_disk()?;
if on_disk.is_null() {
return None;
}
overlay_faulter.as_ref()?.fault_overlay_slot(on_disk)
}
pub(crate) fn overlay_child_node(
child: &Child<K, V>,
overlay_faulter: &Option<Arc<dyn OverlayFaulter<K, V>>>,
) -> Option<Self> {
Self::resolve_overlay_child(child, overlay_faulter)
.map(|node| Self::from_overlay_node(node, overlay_faulter.clone()))
}
#[inline]
fn snapshot_cursor_arena(&self) -> Option<&OverlaySnapshotCursorArena<K, V>> {
let root = Arc::clone(self.overlay.as_ref()?);
Some(
self.snapshot_cursor_arena
.get_or_init(|| OverlaySnapshotCursorArena::build(root, &self.overlay_faulter)),
)
}
}
impl<K: KeyEncoding, V: DictionaryValue> DictionaryNode for OverlayDictionaryNode<K, V> {
type Unit = K::Token;
type SnapshotCursor = crate::SnapshotTraversalCursor;
type SnapshotGraphValueHandle = crate::SnapshotTraversalCursor;
#[inline]
fn snapshot_root_cursor(&self) -> Option<Self::SnapshotCursor> {
let arena = self.snapshot_cursor_arena()?;
(!arena.nodes.is_empty())
.then(|| SnapshotTraversalCursor::from_index(0))
.flatten()
}
#[inline]
fn snapshot_cursor_requires_full_projection(&self) -> bool {
self.snapshot_cursor_arena.get().is_none()
}
#[inline]
fn contains_snapshot_cursor(&self, cursor: Self::SnapshotCursor) -> bool {
self.snapshot_cursor_arena
.get()
.is_some_and(|arena| cursor.index() < arena.nodes.len())
}
#[inline]
fn supports_snapshot_cursor_nodes(&self) -> bool {
true
}
#[inline]
fn supports_snapshot_cursor_key_units(&self) -> bool {
true
}
#[inline]
unsafe fn snapshot_cursor_key_units(
&self,
cursor: Self::SnapshotCursor,
) -> Option<Vec<Self::Unit>> {
let arena = self.snapshot_cursor_arena.get()?;
let mut index = cursor.index();
arena.nodes.get(index)?;
let mut reverse = Vec::new();
while let Some((parent, label)) = arena.nodes[index].parent {
reverse.push(label);
index = parent;
}
reverse.reverse();
Some(reverse)
}
#[inline]
unsafe fn snapshot_cursor_node(&self, cursor: Self::SnapshotCursor) -> Option<Self> {
let node = self.snapshot_cursor_arena.get()?.node(cursor)?;
Some(Self::from_overlay_node(
Arc::clone(&node.overlay),
self.overlay_faulter.clone(),
))
}
#[inline]
unsafe fn filter_map_snapshot_cursor_edges_and_finality<T, P, F>(
&self,
cursor: Self::SnapshotCursor,
mut project: P,
mut visitor: F,
) -> Option<bool>
where
P: FnMut(Self::Unit) -> Option<T>,
F: FnMut(Self::Unit, Self::SnapshotCursor, T),
{
let arena = self.snapshot_cursor_arena.get()?;
let node = arena.node(cursor)?;
for &(label, child_index) in arena.node_edges(node) {
let Some(projected) = project(label) else {
continue;
};
visitor(
label,
SnapshotTraversalCursor::from_index(child_index)?,
projected,
);
}
Some(node.is_final)
}
#[inline]
unsafe fn snapshot_cursor_is_final(&self, cursor: Self::SnapshotCursor) -> Option<bool> {
Some(self.snapshot_cursor_arena.get()?.node(cursor)?.is_final)
}
#[inline]
unsafe fn snapshot_cursor_transition(
&self,
cursor: Self::SnapshotCursor,
wanted: Self::Unit,
) -> Option<Option<Self::SnapshotCursor>> {
let arena = self.snapshot_cursor_arena.get()?;
let node = arena.node(cursor)?;
let edges = arena.node_edges(node);
let found = edges
.binary_search_by_key(&wanted, |&(label, _)| label)
.ok()
.and_then(|edge_index| SnapshotTraversalCursor::from_index(edges[edge_index].1));
Some(found)
}
#[inline]
fn supports_efficient_snapshot_cursor_edge_paging(&self) -> bool {
true
}
#[inline]
unsafe fn visit_snapshot_cursor_edge_page<F>(
&self,
cursor: Self::SnapshotCursor,
start: usize,
capacity: usize,
mut visitor: F,
) -> Option<(bool, usize)>
where
F: FnMut(Self::Unit, Self::SnapshotCursor),
{
let arena = self.snapshot_cursor_arena.get()?;
let node = arena.node(cursor)?;
let edges = arena.node_edges(node);
let total = edges.len();
let end = start.saturating_add(capacity).min(total);
if start < end {
for &(label, child_index) in &edges[start..end] {
visitor(label, SnapshotTraversalCursor::from_index(child_index)?);
}
}
Some((node.is_final, total))
}
fn is_final(&self) -> bool {
match &self.overlay {
Some(node) => node.is_final(),
None => false,
}
}
fn transition(&self, label: K::Token) -> Option<Self> {
let node = self.overlay.as_ref()?;
let child = node.find_child(K::token_to_unit(label))?;
Self::overlay_child_node(child, &self.overlay_faulter)
}
fn edges(&self) -> Box<dyn Iterator<Item = (K::Token, Self)> + '_> {
let Some(node) = &self.overlay else {
return Box::new(std::iter::empty());
};
let mut edges = Vec::with_capacity(node.num_children());
for (&unit, child) in node.iter_children() {
let Some(token) = K::unit_to_token(unit) else {
continue;
};
if let Some(child_node) = Self::overlay_child_node(child, &self.overlay_faulter) {
edges.push((token, child_node));
}
}
Box::new(edges.into_iter())
}
#[inline]
fn for_each_edge<F>(&self, mut visitor: F)
where
F: FnMut(K::Token, Self),
{
let Some(node) = &self.overlay else {
return;
};
for (&unit, child) in node.iter_children() {
let Some(token) = K::unit_to_token(unit) else {
continue;
};
if let Some(child_node) = Self::overlay_child_node(child, &self.overlay_faulter) {
visitor(token, child_node);
}
}
}
#[inline]
fn filter_map_edges<T, P, F>(&self, mut project: P, mut visitor: F)
where
P: FnMut(K::Token) -> Option<T>,
F: FnMut(K::Token, Self, T),
{
let Some(node) = &self.overlay else {
return;
};
for (&unit, child) in node.iter_children() {
let Some(token) = K::unit_to_token(unit) else {
continue;
};
let Some(projected) = project(token) else {
continue;
};
if let Some(child_node) = Self::overlay_child_node(child, &self.overlay_faulter) {
visitor(token, child_node, projected);
}
}
}
fn edge_count(&self) -> Option<usize> {
self.overlay.as_ref().map(|node| node.num_children())
}
}
impl<K: KeyEncoding, V: DictionaryValue> MappedDictionaryNode for OverlayDictionaryNode<K, V> {
type Value = V;
fn value(&self) -> Option<V> {
self.overlay.as_ref().and_then(|node| node.get_value())
}
#[inline]
fn supports_snapshot_cursor_values(&self) -> bool {
true
}
#[inline]
unsafe fn snapshot_cursor_value(&self, cursor: Self::SnapshotCursor) -> Option<Option<V>> {
Some(
self.snapshot_cursor_arena
.get()?
.node(cursor)?
.value
.clone(),
)
}
}
#[allow(dead_code)]
fn _assert_overlay_dictionary_node_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
use crate::persistent_artrie::core::key_encoding::{ByteKey, CharKey};
assert_send_sync::<OverlayDictionaryNode<ByteKey, ()>>();
assert_send_sync::<OverlayDictionaryNode<ByteKey, u64>>();
assert_send_sync::<OverlayDictionaryNode<CharKey, ()>>();
assert_send_sync::<OverlayDictionaryNode<CharKey, u64>>();
}
#[cfg(test)]
mod snapshot_cursor_tests {
use std::collections::BTreeMap;
use std::sync::{Arc, Barrier};
use super::*;
use crate::persistent_artrie::core::key_encoding::{ByteKey, CharKey, U64Key};
fn insert<K: KeyEncoding>(
root: &Arc<OverlayNode<K, u64>>,
units: &[K::Unit],
value: u64,
) -> Arc<OverlayNode<K, u64>> {
if let Some((&head, tail)) = units.split_first() {
let child = root
.find_child(head)
.and_then(Child::as_in_mem)
.cloned()
.unwrap_or_else(|| Arc::new(OverlayNode::new()));
let child = insert(&child, tail, value);
Arc::new(root.with_child(head, Child::InMem(child)))
} else {
Arc::new(root.as_final().with_value(value))
}
}
fn follow<K: KeyEncoding>(
owner: &OverlayDictionaryNode<K, u64>,
mut cursor: SnapshotTraversalCursor,
units: &[K::Token],
) -> Option<SnapshotTraversalCursor> {
for &unit in units {
cursor = unsafe { owner.snapshot_cursor_transition(cursor, unit)? }?;
}
Some(cursor)
}
fn assert_exact_cursor_surface<K: KeyEncoding>(
root: Arc<OverlayNode<K, u64>>,
expected: BTreeMap<Vec<K::Token>, u64>,
) {
let owner = OverlayDictionaryNode::<K, u64>::from_overlay_root(root, None);
assert!(owner.snapshot_cursor_requires_full_projection());
let root_cursor = owner.snapshot_root_cursor().expect("direct root cursor");
assert!(!owner.snapshot_cursor_requires_full_projection());
assert!(owner.contains_snapshot_cursor(root_cursor));
assert!(owner.supports_snapshot_cursor_nodes());
assert!(owner.supports_snapshot_cursor_key_units());
assert!(owner.supports_snapshot_cursor_values());
assert!(owner.supports_efficient_snapshot_cursor_edge_paging());
let mut pending = vec![(Vec::new(), root_cursor)];
while let Some((path, cursor)) = pending.pop() {
let expected_value = expected.get(&path).copied();
assert_eq!(
unsafe { owner.snapshot_cursor_key_units(cursor) },
Some(path.clone())
);
assert_eq!(
unsafe { owner.snapshot_cursor_is_final(cursor) },
Some(expected_value.is_some())
);
assert_eq!(
unsafe { owner.snapshot_cursor_value(cursor) },
Some(expected_value)
);
let materialized = unsafe { owner.snapshot_cursor_node(cursor) }.expect("cursor node");
assert_eq!(materialized.is_final(), expected_value.is_some());
assert_eq!(materialized.value(), expected_value);
let mut edges = Vec::new();
let finality = unsafe {
owner.filter_map_snapshot_cursor_edges_and_finality(
cursor,
Some,
|label, child, label_again| {
assert_eq!(label, label_again);
edges.push((label, child));
},
)
};
assert_eq!(finality, Some(expected_value.is_some()));
assert!(edges.windows(2).all(|pair| pair[0].0 < pair[1].0));
let mut paged = Vec::new();
for start in 0..=edges.len() {
let mut page = Vec::new();
let metadata = unsafe {
owner.visit_snapshot_cursor_edge_page(cursor, start, 1, |label, child| {
page.push((label, child));
})
};
assert_eq!(metadata, Some((expected_value.is_some(), edges.len())));
if start < edges.len() {
assert_eq!(page, vec![edges[start]]);
paged.extend(page);
} else {
assert!(page.is_empty());
}
}
assert_eq!(paged, edges);
for &(label, child) in edges.iter().rev() {
assert_eq!(
unsafe { owner.snapshot_cursor_transition(cursor, label) },
Some(Some(child))
);
let mut child_path = path.clone();
child_path.push(label);
pending.push((child_path, child));
}
}
for (term, value) in expected {
let cursor = follow(&owner, root_cursor, &term).expect("expected term cursor");
assert_eq!(
unsafe { owner.snapshot_cursor_value(cursor) },
Some(Some(value))
);
}
}
#[test]
fn byte_char_u64_and_vocabulary_overlay_cursors_are_exact_and_ordered() {
let mut byte_root = Arc::new(OverlayNode::<ByteKey, u64>::new());
let byte_terms = [
(b"z".as_slice(), 1),
(b"ant".as_slice(), 2),
(b"an".as_slice(), 3),
];
for (term, value) in byte_terms {
byte_root = insert(&byte_root, term, value);
}
assert_exact_cursor_surface::<ByteKey>(
byte_root,
BTreeMap::from([
(b"an".to_vec(), 3),
(b"ant".to_vec(), 2),
(b"z".to_vec(), 1),
]),
);
let mut char_root = Arc::new(OverlayNode::<CharKey, u64>::new());
let char_terms = [
(vec!['雪' as u32], vec!['雪'], 5),
(vec!['a' as u32, 'β' as u32], vec!['a', 'β'], 7),
(vec!['a' as u32], vec!['a'], 11),
];
let mut char_expected = BTreeMap::new();
for (units, tokens, value) in char_terms {
char_root = insert(&char_root, &units, value);
char_expected.insert(tokens, value);
}
assert_exact_cursor_surface::<CharKey>(char_root, char_expected);
let mut u64_root = Arc::new(OverlayNode::<U64Key, u64>::new());
let u64_terms = [(vec![99], 13), (vec![1, 8], 17), (vec![1], 19)];
let mut u64_expected = BTreeMap::new();
for (term, value) in u64_terms {
u64_root = insert(&u64_root, &term, value);
u64_expected.insert(term, value);
}
assert_exact_cursor_surface::<U64Key>(u64_root, u64_expected);
}
#[test]
fn overlay_cursor_revisions_are_isolated_during_concurrent_capture_and_writes() {
let empty = Arc::new(OverlayNode::<ByteKey, u64>::new());
let old_root = insert(&empty, b"ant", 7);
let old_owner = Arc::new(OverlayDictionaryNode::<ByteKey, u64>::from_overlay_root(
Arc::clone(&old_root),
None,
));
let participants = 9;
let barrier = Arc::new(Barrier::new(participants));
let mut threads = Vec::new();
for _ in 0..8 {
let owner = Arc::clone(&old_owner);
let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || {
barrier.wait();
for _ in 0..128 {
let root = owner.snapshot_root_cursor().expect("old root cursor");
let ant = follow(&owner, root, b"ant").expect("old term");
assert_eq!(unsafe { owner.snapshot_cursor_value(ant) }, Some(Some(7)));
assert_eq!(
unsafe { owner.snapshot_cursor_transition(root, b'z') },
Some(None)
);
}
}));
}
barrier.wait();
let mut newest = old_root;
for value in 0..128 {
newest = insert(&newest, b"zoo", value);
}
let fresh_owner = OverlayDictionaryNode::<ByteKey, u64>::from_overlay_root(newest, None);
let fresh_root = fresh_owner
.snapshot_root_cursor()
.expect("fresh root cursor");
let zoo = follow(&fresh_owner, fresh_root, b"zoo").expect("fresh term");
assert_eq!(
unsafe { fresh_owner.snapshot_cursor_value(zoo) },
Some(Some(127))
);
for thread in threads {
thread.join().expect("snapshot reader");
}
}
}