use std::collections::HashMap;
use rabs_protocol::result_identity::ObjectId;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct GraphBounds {
pub max_depth: usize,
pub max_fanout: usize,
pub max_nodes: usize,
}
pub const DEFAULT_BOUNDS: GraphBounds = GraphBounds {
max_depth: 64,
max_fanout: 65_536,
max_nodes: 1_048_576,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ManifestNode {
pub id: ObjectId,
pub references: Vec<ObjectId>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClosureError {
Cycle(ObjectId),
DepthExceeded,
FanoutExceeded(ObjectId),
NodeCountExceeded,
DanglingReference(ObjectId),
DuplicateNode(ObjectId),
}
#[derive(Clone, Copy)]
enum VisitState {
Unseen,
Active,
Complete(usize),
}
struct Frame {
node: usize,
next_child: usize,
height: usize,
}
pub fn validate_closure(
root: &ObjectId,
nodes: &[ManifestNode],
bounds: GraphBounds,
) -> Result<(), ClosureError> {
if nodes.len() > bounds.max_nodes {
return Err(ClosureError::NodeCountExceeded);
}
let mut index = HashMap::with_capacity(nodes.len());
for (position, node) in nodes.iter().enumerate() {
if node.references.len() > bounds.max_fanout {
return Err(ClosureError::FanoutExceeded(node.id.clone()));
}
if index.insert(&node.id, position).is_some() {
return Err(ClosureError::DuplicateNode(node.id.clone()));
}
}
let Some(&root_index) = index.get(root) else {
return Err(ClosureError::DanglingReference(root.clone()));
};
let mut states = vec![VisitState::Unseen; nodes.len()];
states[root_index] = VisitState::Active;
let mut stack = vec![Frame {
node: root_index,
next_child: 0,
height: 0,
}];
while let Some(mut frame) = stack.pop() {
let depth = stack.len();
let node = &nodes[frame.node];
if let Some(child) = node.references.get(frame.next_child) {
frame.next_child += 1;
let child_depth = depth.checked_add(1).ok_or(ClosureError::DepthExceeded)?;
if child_depth > bounds.max_depth {
return Err(ClosureError::DepthExceeded);
}
let Some(&child_index) = index.get(child) else {
return Err(ClosureError::DanglingReference(child.clone()));
};
match states[child_index] {
VisitState::Active => return Err(ClosureError::Cycle(child.clone())),
VisitState::Complete(height) => {
if height > bounds.max_depth - child_depth {
return Err(ClosureError::DepthExceeded);
}
let through_child = height.checked_add(1).ok_or(ClosureError::DepthExceeded)?;
frame.height = frame.height.max(through_child);
stack.push(frame);
}
VisitState::Unseen => {
stack.push(frame);
states[child_index] = VisitState::Active;
stack.push(Frame {
node: child_index,
next_child: 0,
height: 0,
});
}
}
} else {
if frame.height > bounds.max_depth - depth {
return Err(ClosureError::DepthExceeded);
}
states[frame.node] = VisitState::Complete(frame.height);
if let Some(parent) = stack.last_mut() {
let through_child = frame
.height
.checked_add(1)
.ok_or(ClosureError::DepthExceeded)?;
parent.height = parent.height.max(through_child);
}
}
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PackMember {
pub offset: u64,
pub length: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PackError {
Overlap,
OutOfBounds,
EmptyMember,
}
pub fn validate_pack_ranges(members: &[PackMember], pack_len: u64) -> Result<(), PackError> {
let mut sorted: Vec<&PackMember> = members.iter().collect();
sorted.sort_by_key(|m| m.offset);
let mut previous_end: u64 = 0;
for member in sorted {
if member.length == 0 {
return Err(PackError::EmptyMember);
}
let end = member
.offset
.checked_add(member.length)
.ok_or(PackError::OutOfBounds)?;
if end > pack_len {
return Err(PackError::OutOfBounds);
}
if member.offset < previous_end {
return Err(PackError::Overlap);
}
previous_end = end;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use rabs_protocol::result_identity::{DigestAlgorithm, TypedDigest};
fn id(tag: u8) -> ObjectId {
ObjectId(TypedDigest {
algorithm: DigestAlgorithm::Sha256V1,
domain: "rabs.object.v1",
bytes: [tag; 32],
})
}
fn node(tag: u8, refs: &[u8]) -> ManifestNode {
ManifestNode {
id: id(tag),
references: refs.iter().map(|t| id(*t)).collect(),
}
}
#[test]
fn clean_dags_validate_including_shared_subtrees() {
let nodes = vec![node(1, &[2, 3]), node(2, &[4]), node(3, &[4]), node(4, &[])];
assert_eq!(validate_closure(&id(1), &nodes, DEFAULT_BOUNDS), Ok(()));
}
#[test]
fn shared_subtrees_respect_longest_path_in_either_visit_order() {
for references in [[2, 3], [3, 2]] {
let mut nodes = vec![
node(1, &references),
node(2, &[4]),
node(3, &[2]),
node(4, &[5]),
node(5, &[]),
];
for _ in 0..2 {
assert_eq!(
validate_closure(
&id(1),
&nodes,
GraphBounds {
max_depth: 3,
..DEFAULT_BOUNDS
},
),
Err(ClosureError::DepthExceeded)
);
assert_eq!(
validate_closure(
&id(1),
&nodes,
GraphBounds {
max_depth: 4,
..DEFAULT_BOUNDS
},
),
Ok(())
);
nodes.reverse();
}
}
}
#[test]
fn duplicate_identities_cannot_hide_a_different_graph() {
for duplicate in [node(1, &[]), node(1, &[1]), node(1, &[99])] {
let mut nodes = vec![node(1, &[]), duplicate];
for _ in 0..2 {
assert_eq!(
validate_closure(&id(1), &nodes, DEFAULT_BOUNDS),
Err(ClosureError::DuplicateNode(id(1)))
);
nodes.reverse();
}
}
}
#[test]
fn input_budget_and_zero_depth_are_enforced() {
let root_only = [node(1, &[])];
let zero_depth = GraphBounds {
max_depth: 0,
max_fanout: 1,
max_nodes: 1,
};
assert_eq!(validate_closure(&id(1), &root_only, zero_depth), Ok(()));
assert_eq!(
validate_closure(
&id(1),
&root_only,
GraphBounds {
max_nodes: 0,
..zero_depth
},
),
Err(ClosureError::NodeCountExceeded)
);
assert_eq!(
validate_closure(&id(1), &[], zero_depth),
Err(ClosureError::DanglingReference(id(1)))
);
assert_eq!(
validate_closure(&id(1), &[node(1, &[1])], zero_depth),
Err(ClosureError::DepthExceeded)
);
assert_eq!(
validate_closure(&id(1), &[node(1, &[]), node(2, &[])], zero_depth),
Err(ClosureError::NodeCountExceeded)
);
}
#[test]
fn large_depth_policy_does_not_recurse_on_process_stack() {
fn numbered_id(number: usize) -> ObjectId {
let mut object = id(0);
object.0.bytes[..8].copy_from_slice(&u64::try_from(number).unwrap().to_be_bytes());
object
}
let count = 20_000;
let nodes: Vec<_> = (0..count)
.map(|number| ManifestNode {
id: numbered_id(number),
references: if number + 1 == count {
Vec::new()
} else {
vec![numbered_id(number + 1)]
},
})
.collect();
let result = std::thread::Builder::new()
.stack_size(64 * 1024)
.spawn(move || {
validate_closure(
&numbered_id(0),
&nodes,
GraphBounds {
max_depth: usize::MAX,
max_fanout: 1,
max_nodes: count,
},
)
})
.unwrap()
.join()
.unwrap();
assert_eq!(result, Ok(()));
}
#[test]
fn cycle_corpus_rejected_before_heavy_traversal() {
let self_cycle = vec![node(1, &[1])];
assert_eq!(
validate_closure(&id(1), &self_cycle, DEFAULT_BOUNDS),
Err(ClosureError::Cycle(id(1)))
);
let chained = vec![node(1, &[2]), node(2, &[3]), node(3, &[2])];
assert_eq!(
validate_closure(&id(1), &chained, DEFAULT_BOUNDS),
Err(ClosureError::Cycle(id(2)))
);
}
#[test]
fn bounds_and_closure_holes_reject() {
let tight = GraphBounds {
max_depth: 3,
max_fanout: 10,
max_nodes: 100,
};
let chain = vec![
node(1, &[2]),
node(2, &[3]),
node(3, &[4]),
node(4, &[5]),
node(5, &[]),
];
assert_eq!(
validate_closure(&id(1), &chain, tight),
Err(ClosureError::DepthExceeded)
);
let wide_refs: Vec<u8> = (10..=30).collect();
let mut wide = vec![ManifestNode {
id: id(1),
references: wide_refs.iter().map(|t| id(*t)).collect(),
}];
wide.extend(wide_refs.iter().map(|t| node(*t, &[])));
assert_eq!(
validate_closure(&id(1), &wide, tight),
Err(ClosureError::FanoutExceeded(id(1)))
);
let dangling = vec![node(1, &[2])];
assert_eq!(
validate_closure(&id(1), &dangling, DEFAULT_BOUNDS),
Err(ClosureError::DanglingReference(id(2)))
);
}
#[test]
fn pack_range_corpus_rejected() {
let ok = [
PackMember {
offset: 0,
length: 10,
},
PackMember {
offset: 10,
length: 5,
},
PackMember {
offset: 20,
length: 4,
},
];
assert_eq!(validate_pack_ranges(&ok, 24), Ok(()));
let overlap = [
PackMember {
offset: 8,
length: 5,
},
PackMember {
offset: 0,
length: 10,
},
];
assert_eq!(validate_pack_ranges(&overlap, 100), Err(PackError::Overlap));
let oob = [PackMember {
offset: 20,
length: 10,
}];
assert_eq!(validate_pack_ranges(&oob, 25), Err(PackError::OutOfBounds));
let wrap = [PackMember {
offset: u64::MAX - 1,
length: 10,
}];
assert_eq!(
validate_pack_ranges(&wrap, u64::MAX),
Err(PackError::OutOfBounds)
);
let empty = [PackMember {
offset: 0,
length: 0,
}];
assert_eq!(
validate_pack_ranges(&empty, 10),
Err(PackError::EmptyMember)
);
}
}