use smallvec::SmallVec;
use crate::dict::NameDict;
use crate::{NodeId, RelationshipId};
const INLINE_BYTES: usize = 16;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct AdjEntry {
pub type_id: u32,
pub neighbour: NodeId,
pub rel: RelationshipId,
}
#[derive(Clone, Default)]
pub(super) struct AdjList {
bytes: SmallVec<u8, INLINE_BYTES>,
}
#[inline]
fn width(value: u64) -> usize {
((64 - value.leading_zeros() as usize).div_ceil(8)).max(1)
}
#[inline]
fn type_width(code: u8) -> usize {
match code {
0 => 1,
1 => 2,
_ => 4,
}
}
#[inline]
fn read(bytes: &[u8], pos: usize, n: usize) -> u64 {
if let Some(word) = bytes.get(pos..pos + 8) {
let word = u64::from_le_bytes(word.try_into().unwrap());
word & (u64::MAX >> (64 - 8 * n))
} else {
let mut buf = [0u8; 8];
buf[..n].copy_from_slice(&bytes[pos..pos + n]);
u64::from_le_bytes(buf)
}
}
impl AdjList {
pub(super) fn new() -> Self {
Self::default()
}
#[inline]
pub(super) fn is_empty(&self) -> bool {
self.bytes.is_empty()
}
pub(super) fn clear(&mut self) {
self.bytes.clear();
}
pub(super) fn push(&mut self, entry: AdjEntry) {
let nb = width(entry.neighbour);
let rb = width(entry.rel);
let (tcode, tb) = match entry.type_id {
0..=0xFF => (0u8, 1),
0x100..=0xFFFF => (1, 2),
_ => (2, 4),
};
let header = (nb as u8 - 1) | ((rb as u8 - 1) << 3) | (tcode << 6);
self.bytes.reserve(1 + tb + nb + rb);
self.bytes.push(header);
self.bytes
.extend_from_slice(&entry.type_id.to_le_bytes()[..tb]);
self.bytes
.extend_from_slice(&entry.neighbour.to_le_bytes()[..nb]);
self.bytes.extend_from_slice(&entry.rel.to_le_bytes()[..rb]);
}
pub(super) fn remove(&mut self, rel: RelationshipId) -> bool {
let mut iter = self.iter();
loop {
let start = iter.pos;
let Some(entry) = iter.next() else {
return false;
};
if entry.rel == rel {
let end = iter.pos;
self.bytes.drain(start..end);
return true;
}
}
}
#[inline]
pub(super) fn iter(&self) -> AdjIter<'_> {
AdjIter {
bytes: &self.bytes,
pos: 0,
}
}
pub(super) fn count(&self) -> usize {
let bytes = self.bytes.as_slice();
let mut pos = 0;
let mut n = 0;
while pos < bytes.len() {
let h = bytes[pos];
pos += 1 + type_width(h >> 6) + (h & 7) as usize + 1 + ((h >> 3) & 7) as usize + 1;
n += 1;
}
n
}
pub(super) fn heap_bytes(&self) -> usize {
if self.bytes.spilled() {
self.bytes.capacity()
} else {
0
}
}
}
impl std::fmt::Debug for AdjList {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_list().entries(self.iter()).finish()
}
}
pub(super) struct AdjIter<'a> {
bytes: &'a [u8],
pos: usize,
}
impl Iterator for AdjIter<'_> {
type Item = AdjEntry;
#[inline]
fn next(&mut self) -> Option<AdjEntry> {
let h = *self.bytes.get(self.pos)?;
let tb = type_width(h >> 6);
let nb = (h & 7) as usize + 1;
let rb = ((h >> 3) & 7) as usize + 1;
let p = self.pos + 1;
let entry = AdjEntry {
type_id: read(self.bytes, p, tb) as u32,
neighbour: read(self.bytes, p + tb, nb),
rel: read(self.bytes, p + tb + nb, rb),
};
self.pos = p + tb + nb + rb;
Some(entry)
}
}
pub(super) enum TypeFilter {
Any,
One(u32),
Many(SmallVec<u32, 4>),
Nothing,
}
impl TypeFilter {
#[inline]
pub(super) fn resolve(dict: &NameDict, types: &[String]) -> Self {
match types {
[] => TypeFilter::Any,
[single] => match dict.id_of(single) {
Some(id) => TypeFilter::One(id),
None => TypeFilter::Nothing,
},
many => {
let ids: SmallVec<u32, 4> = many.iter().filter_map(|t| dict.id_of(t)).collect();
match ids.as_slice() {
[] => TypeFilter::Nothing,
[one] => TypeFilter::One(*one),
_ => TypeFilter::Many(ids),
}
}
}
}
#[inline]
pub(super) fn matches(&self, type_id: u32) -> bool {
match self {
TypeFilter::Any => true,
TypeFilter::One(id) => *id == type_id,
TypeFilter::Many(ids) => ids.contains(&type_id),
TypeFilter::Nothing => false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn entries() -> Vec<AdjEntry> {
let ids = [
0u64,
1,
255,
256,
65_535,
65_536,
16_777_215,
16_777_216,
u32::MAX as u64,
u32::MAX as u64 + 1,
(1 << 48) - 1,
1 << 56,
u64::MAX,
];
let types = [0u32, 1, 255, 256, 65_535, 65_536, u32::MAX];
let mut out = Vec::new();
for (i, &neighbour) in ids.iter().enumerate() {
for (j, &type_id) in types.iter().enumerate() {
let rel = ids[(i * 7 + j * 3) % ids.len()] ^ (out.len() as u64);
out.push(AdjEntry {
type_id,
neighbour,
rel,
});
}
}
out
}
#[test]
fn every_width_round_trips() {
let entries = entries();
let mut list = AdjList::new();
for &e in &entries {
list.push(e);
}
assert_eq!(list.iter().collect::<Vec<_>>(), entries);
assert_eq!(list.count(), entries.len());
}
#[test]
fn a_single_short_entry_reads_without_eight_bytes_behind_it() {
let mut list = AdjList::new();
let e = AdjEntry {
type_id: 3,
neighbour: 7,
rel: 9,
};
list.push(e);
assert_eq!(list.iter().collect::<Vec<_>>(), vec![e]);
assert_eq!(list.heap_bytes(), 0);
}
#[test]
fn remove_takes_exactly_one_entry_and_keeps_order() {
let entries = entries();
let mut list = AdjList::new();
for &e in &entries {
list.push(e);
}
let mut expected = entries.clone();
for index in [expected.len() / 2, 0, expected.len() - 3] {
let gone = expected.remove(index);
assert!(list.remove(gone.rel));
assert!(!list.remove(gone.rel));
assert_eq!(list.iter().collect::<Vec<_>>(), expected);
assert_eq!(list.count(), expected.len());
}
for e in expected {
assert!(list.remove(e.rel));
}
assert!(list.is_empty());
}
#[test]
fn two_entries_of_a_16m_graph_stay_inline() {
let mut list = AdjList::new();
for rel in [16_000_000u64, 16_000_001] {
list.push(AdjEntry {
type_id: 200,
neighbour: 15_999_999,
rel,
});
}
assert_eq!(list.heap_bytes(), 0);
assert_eq!(std::mem::size_of::<AdjList>(), 24);
}
#[test]
fn a_type_filter_resolves_names_once() {
let mut dict = NameDict::default();
for i in 0..20 {
dict.id_or_insert_str(&format!("T{i}"));
}
let filter = |types: &[&str]| {
let types: Vec<String> = types.iter().map(|t| t.to_string()).collect();
TypeFilter::resolve(&dict, &types)
};
assert!(filter(&[]).matches(19));
let one = filter(&["T4"]);
assert!(one.matches(4) && !one.matches(5));
let many = filter(&["T4", "nope", "T9"]);
assert!(many.matches(4) && many.matches(9) && !many.matches(5));
assert!(!filter(&["nope"]).matches(0));
}
}