use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::{Path, PathBuf};
use crate::csr::CsrSnapshot;
use crate::error::Error;
const MAGIC: &[u8; 8] = b"ISSNCSR2";
const FLAG_WEIGHTED: u64 = 1;
const FLAG_NEGATIVE_WEIGHT: u64 = 2;
struct Fnv(u64);
impl Fnv {
fn new() -> Self {
Fnv(0xcbf2_9ce4_8422_2325)
}
fn update(&mut self, bytes: &[u8]) {
for &b in bytes {
self.0 ^= b as u64;
self.0 = self.0.wrapping_mul(0x0000_0100_0000_01b3);
}
}
}
struct SumWriter<W: Write> {
inner: W,
sum: Fnv,
}
impl<W: Write> SumWriter<W> {
fn put(&mut self, bytes: &[u8]) -> std::io::Result<()> {
self.sum.update(bytes);
self.inner.write_all(bytes)
}
fn put_u64(&mut self, v: u64) -> std::io::Result<()> {
self.put(&v.to_le_bytes())
}
fn put_u64s(&mut self, vs: &[u64]) -> std::io::Result<()> {
for &v in vs {
self.put(&v.to_le_bytes())?;
}
Ok(())
}
fn put_usizes(&mut self, vs: &[usize]) -> std::io::Result<()> {
for &v in vs {
self.put(&(v as u64).to_le_bytes())?;
}
Ok(())
}
fn put_u32s(&mut self, vs: &[u32]) -> std::io::Result<()> {
for &v in vs {
self.put(&v.to_le_bytes())?;
}
Ok(())
}
fn put_f64s(&mut self, vs: &[f64]) -> std::io::Result<()> {
for &v in vs {
self.put(&v.to_le_bytes())?;
}
Ok(())
}
}
struct SumReader<R: Read> {
inner: R,
sum: Fnv,
}
impl<R: Read> SumReader<R> {
fn get<const N: usize>(&mut self) -> std::io::Result<[u8; N]> {
let mut buf = [0u8; N];
self.inner.read_exact(&mut buf)?;
self.sum.update(&buf);
Ok(buf)
}
fn get_u64(&mut self) -> std::io::Result<u64> {
Ok(u64::from_le_bytes(self.get::<8>()?))
}
fn get_u64s(&mut self, n: usize) -> std::io::Result<Vec<u64>> {
let mut out = Vec::with_capacity(n);
for _ in 0..n {
out.push(self.get_u64()?);
}
Ok(out)
}
fn get_usizes(&mut self, n: usize) -> std::io::Result<Vec<usize>> {
let mut out = Vec::with_capacity(n);
for _ in 0..n {
out.push(self.get_u64()? as usize);
}
Ok(out)
}
fn get_u32s(&mut self, n: usize) -> std::io::Result<Vec<u32>> {
let mut out = Vec::with_capacity(n);
for _ in 0..n {
out.push(u32::from_le_bytes(self.get::<4>()?));
}
Ok(out)
}
fn get_f64s(&mut self, n: usize) -> std::io::Result<Vec<f64>> {
let mut out = Vec::with_capacity(n);
for _ in 0..n {
out.push(f64::from_le_bytes(self.get::<8>()?));
}
Ok(out)
}
}
impl<W: Write> Write for SumWriter<W> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let n = self.inner.write(buf)?;
self.sum.update(&buf[..n]);
Ok(n)
}
fn flush(&mut self) -> std::io::Result<()> {
self.inner.flush()
}
}
impl<R: Read> Read for SumReader<R> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let n = self.inner.read(buf)?;
self.sum.update(&buf[..n]);
Ok(n)
}
}
fn csr_path(dir: &Path) -> PathBuf {
dir.join("csr.cache")
}
pub(crate) fn save_csr(
dir: &Path,
snap: &CsrSnapshot,
db_id: [u8; 16],
commit_gen: u64,
) -> Result<(), Error> {
let tmp = dir.join("csr.cache.tmp");
let write = || -> std::io::Result<()> {
let mut w = SumWriter {
inner: BufWriter::new(File::create(&tmp)?),
sum: Fnv::new(),
};
w.put(MAGIC)?;
w.put(&db_id)?;
w.put_u64(commit_gen)?;
let mut flags = 0u64;
if snap.edge_weight.is_some() {
flags |= FLAG_WEIGHTED;
}
if snap.has_negative_weight {
flags |= FLAG_NEGATIVE_WEIGHT;
}
w.put_u64(flags)?;
w.put_u64(snap.dense_to_id.len() as u64)?;
w.put_u64(snap.col_idx.len() as u64)?;
w.put_u64s(&snap.dense_to_id)?;
w.put_usizes(&snap.row_ptr)?;
w.put_u32s(&snap.col_idx)?;
w.put_u32s(&snap.edge_type)?;
w.put_u64s(&snap.edge_id)?;
w.put_usizes(&snap.in_row_ptr)?;
w.put_u32s(&snap.in_col_idx)?;
w.put_u32s(&snap.in_edge_type)?;
w.put_u64s(&snap.in_edge_id)?;
if let Some(weights) = &snap.edge_weight {
w.put_f64s(weights)?;
}
let sum = w.sum.0;
w.inner.write_all(&sum.to_le_bytes())?;
w.inner.flush()?;
Ok(())
};
write().map_err(Error::Io)?;
std::fs::rename(&tmp, csr_path(dir)).map_err(Error::Io)?;
Ok(())
}
fn csr_file_len(n: u64, e: u64, weighted: bool) -> Option<u64> {
let mut total = 8u64 + 16 + 8 + 8 + 8 + 8 + 8;
let u64_elems = n
.checked_add(n.checked_add(1)?.checked_mul(2)?)?
.checked_add(e.checked_mul(2)?)?;
total = total.checked_add(u64_elems.checked_mul(8)?)?;
total = total.checked_add(e.checked_mul(4)?.checked_mul(4)?)?;
if weighted {
total = total.checked_add(e.checked_mul(8)?)?;
}
Some(total)
}
pub(crate) fn load_csr(
dir: &Path,
db_id: [u8; 16],
expected_gen: u64,
want_weights: bool,
) -> Option<CsrSnapshot> {
let file = File::open(csr_path(dir)).ok()?;
let file_len = file.metadata().ok()?.len();
let mut r = SumReader {
inner: BufReader::new(file),
sum: Fnv::new(),
};
let read = |r: &mut SumReader<BufReader<File>>| -> std::io::Result<Option<CsrSnapshot>> {
let magic = r.get::<8>()?;
if &magic != MAGIC {
return Ok(None);
}
if r.get::<16>()? != db_id {
return Ok(None);
}
let file_gen = r.get_u64()?;
if file_gen != expected_gen {
return Ok(None);
}
let flags = r.get_u64()?;
let weighted = flags & FLAG_WEIGHTED != 0;
if want_weights && !weighted {
return Ok(None);
}
let n64 = r.get_u64()?;
let e64 = r.get_u64()?;
match csr_file_len(n64, e64, weighted) {
Some(required) if required <= file_len => {}
_ => return Ok(None),
}
let n = n64 as usize;
let e = e64 as usize;
let dense_to_id = r.get_u64s(n)?;
let row_ptr = r.get_usizes(n + 1)?;
let col_idx = r.get_u32s(e)?;
let edge_type = r.get_u32s(e)?;
let edge_id = r.get_u64s(e)?;
let in_row_ptr = r.get_usizes(n + 1)?;
let in_col_idx = r.get_u32s(e)?;
let in_edge_type = r.get_u32s(e)?;
let in_edge_id = r.get_u64s(e)?;
let edge_weight = if weighted { Some(r.get_f64s(e)?) } else { None };
let expected_sum = r.sum.0;
let mut sum_buf = [0u8; 8];
r.inner.read_exact(&mut sum_buf)?;
if u64::from_le_bytes(sum_buf) != expected_sum {
return Ok(None);
}
let id_to_dense = dense_to_id
.iter()
.enumerate()
.map(|(d, &id)| (id, d as u32))
.collect();
Ok(Some(CsrSnapshot {
row_ptr,
col_idx,
edge_type,
edge_id,
edge_weight,
has_negative_weight: flags & FLAG_NEGATIVE_WEIGHT != 0,
in_row_ptr,
in_col_idx,
in_edge_type,
in_edge_id,
dense_to_id,
id_to_dense,
}))
};
read(&mut r).ok().flatten()
}
const COL_MAGIC: &[u8; 8] = b"ISSNCOL2";
#[derive(serde::Serialize)]
struct ColumnsPayloadRef<'a> {
dense_to_id: &'a Vec<u64>,
cols: Vec<(&'a String, &'a crate::columns::PropColumn)>,
}
#[derive(serde::Deserialize)]
struct ColumnsPayload {
dense_to_id: Vec<u64>,
cols: Vec<(String, crate::columns::PropColumn)>,
}
fn cache_file_gen(path: &Path, db_id: [u8; 16]) -> Option<u64> {
let mut r = BufReader::new(File::open(path).ok()?);
let mut header = [0u8; 32];
r.read_exact(&mut header).ok()?;
if &header[..8] != COL_MAGIC || header[8..24] != db_id {
return None;
}
let mut gen_bytes = [0u8; 8];
gen_bytes.copy_from_slice(&header[24..]);
Some(u64::from_le_bytes(gen_bytes))
}
pub(crate) fn save_columns<S: crate::columns::ColumnSource<Id = u64>>(
storage: &crate::storage::Storage,
cols: &crate::columns::PropColumns<S>,
commit_gen: u64,
) -> Result<(), Error> {
let dir = storage.env.path();
let path = dir.join(S::CACHE_FILE);
if cache_file_gen(&path, storage.db_id) == Some(commit_gen) {
return Ok(());
}
let tmp = dir.join(format!("{}.tmp", S::CACHE_FILE));
let (dense_to_id, col_list) = cols.cache_file_parts();
let payload = ColumnsPayloadRef {
dense_to_id,
cols: col_list,
};
let write = || -> Result<(), Error> {
let mut w = SumWriter {
inner: BufWriter::new(File::create(&tmp).map_err(Error::Io)?),
sum: Fnv::new(),
};
w.put(COL_MAGIC).map_err(Error::Io)?;
w.put(&storage.db_id).map_err(Error::Io)?;
w.put_u64(commit_gen).map_err(Error::Io)?;
rmp_serde::encode::write(&mut w, &payload)?;
let sum = w.sum.0;
w.inner.write_all(&sum.to_le_bytes()).map_err(Error::Io)?;
w.inner.flush().map_err(Error::Io)?;
Ok(())
};
write()?;
std::fs::rename(&tmp, path).map_err(Error::Io)?;
Ok(())
}
pub(crate) fn load_columns<S: crate::columns::ColumnSource<Id = u64>>(
storage: &crate::storage::Storage,
) -> Option<crate::columns::PropColumns<S>> {
let expected_gen = {
let rtxn = storage.env.read_txn().ok()?;
crate::storage::ids::commit_gen(storage, &rtxn).ok()?
};
let path = storage.env.path().join(S::CACHE_FILE);
let mut r = SumReader {
inner: BufReader::new(File::open(path).ok()?),
sum: Fnv::new(),
};
let magic = r.get::<8>().ok()?;
if &magic != COL_MAGIC {
return None;
}
if r.get::<16>().ok()? != storage.db_id {
return None;
}
if r.get_u64().ok()? != expected_gen {
return None;
}
let payload: ColumnsPayload = rmp_serde::decode::from_read(&mut r).ok()?;
let expected_sum = r.sum.0;
let mut sum_buf = [0u8; 8];
r.inner.read_exact(&mut sum_buf).ok()?;
if u64::from_le_bytes(sum_buf) != expected_sum {
return None;
}
crate::columns::PropColumns::from_cache_file(payload.dense_to_id, payload.cols)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use tempfile::TempDir;
use super::*;
use crate::Graph;
fn open_tmp() -> (TempDir, Graph) {
let dir = TempDir::new().unwrap();
let g = Graph::open(dir.path(), 1).unwrap();
(dir, g)
}
#[test]
fn rebuild_persists_and_a_reopen_loads() {
let dir = TempDir::new().unwrap();
let (a, b, c);
{
let g = Graph::open(dir.path(), 1).unwrap();
a = g.add_node("N", &json!({})).unwrap();
b = g.add_node("N", &json!({})).unwrap();
c = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
g.add_edge(b, c, "R", &json!({})).unwrap();
g.rebuild_csr().unwrap();
}
assert!(csr_path(dir.path()).exists(), "rebuild_csr must persist");
let g = Graph::open(dir.path(), 1).unwrap();
let path = g.shortest_path(a, c).unwrap().expect("a -> b -> c");
assert_eq!(path, vec![a, b, c]);
}
#[test]
fn a_stale_cache_file_is_refused() {
let dir = TempDir::new().unwrap();
let (a, b, c);
{
let g = Graph::open(dir.path(), 1).unwrap();
a = g.add_node("N", &json!({})).unwrap();
b = g.add_node("N", &json!({})).unwrap();
c = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
g.rebuild_csr().unwrap();
g.add_edge(b, c, "R", &json!({})).unwrap();
}
let g = Graph::open(dir.path(), 1).unwrap();
let path = g.shortest_path(a, c).unwrap();
assert_eq!(
path,
Some(vec![a, b, c]),
"the post-save edge must be visible, so the stale cache file must not serve"
);
}
#[test]
fn a_corrupt_cache_file_is_refused() {
let dir = TempDir::new().unwrap();
let (a, b);
{
let g = Graph::open(dir.path(), 1).unwrap();
a = g.add_node("N", &json!({})).unwrap();
b = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
g.rebuild_csr().unwrap();
}
let p = csr_path(dir.path());
let mut bytes = std::fs::read(&p).unwrap();
let mid = bytes.len() / 2;
bytes[mid] ^= 0xff;
std::fs::write(&p, bytes).unwrap();
let g = Graph::open(dir.path(), 1).unwrap();
let path = g.shortest_path(a, b).unwrap();
assert_eq!(path, Some(vec![a, b]));
}
#[test]
fn materialize_persists_columns_and_a_reopen_loads_them() {
use crate::columns::{ColumnSource, NodeSource};
let dir = TempDir::new().unwrap();
let ids: Vec<u64>;
{
let g = Graph::open(dir.path(), 1).unwrap();
ids = vec![
g.add_node(
"N",
&json!({ "i": 42, "f": 1.5, "b": true, "s": "x", "m": 1 }),
)
.unwrap(),
g.add_node("N", &json!({ "s": "y", "m": "one" })).unwrap(),
g.add_node("N", &json!({ "i": 7, "m": [1, 2] })).unwrap(),
];
g.materialize_property_columns().unwrap();
assert!(dir.path().join(NodeSource::CACHE_FILE).exists());
}
let g = Graph::open(dir.path(), 1).unwrap();
assert!(
load_columns::<NodeSource>(&g.storage).is_some(),
"a fresh cache file must load"
);
let vals = g
.node_prop_json_column(&ids, "m")
.expect("bulk gather through the loaded columns");
assert_eq!(vals, vec![json!(1), json!("one"), json!([1, 2])]);
let vals = g.node_prop_json_column(&ids, "s").unwrap();
assert_eq!(vals, vec![json!("x"), json!("y"), serde_json::Value::Null]);
g.update_node(ids[2], &json!({ "s": "x" })).unwrap();
assert_eq!(g.node_prop_json(ids[2], "s").unwrap(), Some(json!("x")));
g.add_node("N", &json!({ "i": 1 })).unwrap();
assert!(
load_columns::<NodeSource>(&g.storage).is_none(),
"a stale columns cache file must be refused"
);
}
#[test]
fn materialize_persists_edge_columns_and_a_reopen_loads_them() {
use crate::columns::{ColumnSource, EdgeSource};
let dir = TempDir::new().unwrap();
let e;
{
let g = Graph::open(dir.path(), 1).unwrap();
let a = g.add_node("N", &json!({})).unwrap();
let b = g.add_node("N", &json!({})).unwrap();
e = g.add_edge(a, b, "R", &json!({ "w": 2.5 })).unwrap();
g.materialize_edge_property_columns().unwrap();
assert!(dir.path().join(EdgeSource::CACHE_FILE).exists());
}
let g = Graph::open(dir.path(), 1).unwrap();
assert!(
load_columns::<EdgeSource>(&g.storage).is_some(),
"a fresh edge columns cache file must load"
);
assert_eq!(
g.edge_prop_json_column(&[e], "w").unwrap(),
vec![json!(2.5)]
);
g.add_node("N", &json!({})).unwrap();
assert!(load_columns::<EdgeSource>(&g.storage).is_none());
}
#[test]
fn materialize_at_an_unchanged_generation_skips_the_rewrite() {
use crate::columns::{ColumnSource, NodeSource};
let dir = TempDir::new().unwrap();
let g = Graph::open(dir.path(), 1).unwrap();
g.add_node("N", &json!({ "i": 1 })).unwrap();
g.materialize_property_columns().unwrap();
let path = dir.path().join(NodeSource::CACHE_FILE);
let before = std::fs::metadata(&path).unwrap().modified().unwrap();
g.materialize_property_columns().unwrap();
let after = std::fs::metadata(&path).unwrap().modified().unwrap();
assert_eq!(before, after, "an unchanged generation must skip the save");
}
#[test]
fn a_corrupt_length_field_is_refused_without_allocating() {
let byte_offset_of_n = MAGIC.len() + 16 + 8 + 8;
for corrupt_len in [u64::MAX, 1u64 << 40] {
let dir = TempDir::new().unwrap();
let (a, b);
{
let g = Graph::open(dir.path(), 1).unwrap();
a = g.add_node("N", &json!({})).unwrap();
b = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
g.rebuild_csr().unwrap();
}
let p = csr_path(dir.path());
let mut bytes = std::fs::read(&p).unwrap();
bytes[byte_offset_of_n..byte_offset_of_n + 8]
.copy_from_slice(&corrupt_len.to_le_bytes());
std::fs::write(&p, bytes).unwrap();
let g = Graph::open(dir.path(), 1).unwrap();
let path = g.shortest_path(a, b).unwrap();
assert_eq!(path, Some(vec![a, b]), "length {corrupt_len}");
}
}
#[test]
fn a_cache_file_from_another_database_is_refused() {
let dir1 = TempDir::new().unwrap();
let dir2 = TempDir::new().unwrap();
let (a, b, c);
{
let g = Graph::open(dir1.path(), 1).unwrap();
a = g.add_node("N", &json!({})).unwrap();
b = g.add_node("N", &json!({})).unwrap();
c = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
g.update_node(a, &json!({ "x": 1 })).unwrap();
g.rebuild_csr().unwrap();
}
{
let g = Graph::open(dir2.path(), 1).unwrap();
let a2 = g.add_node("N", &json!({})).unwrap();
let b2 = g.add_node("N", &json!({})).unwrap();
let c2 = g.add_node("N", &json!({})).unwrap();
g.add_edge(a2, b2, "R", &json!({})).unwrap();
g.add_edge(b2, c2, "R", &json!({})).unwrap();
assert_eq!((a2, b2, c2), (a, b, c));
}
std::fs::remove_file(dir1.path().join("data.mdb")).unwrap();
let _ = std::fs::remove_file(dir1.path().join("lock.mdb"));
std::fs::copy(dir2.path().join("data.mdb"), dir1.path().join("data.mdb")).unwrap();
let g = Graph::open(dir1.path(), 1).unwrap();
assert_eq!(
g.shortest_path(a, c).unwrap(),
Some(vec![a, b, c]),
"the foreign cache file must not serve the old database's adjacency"
);
}
#[test]
fn rebuild_csr_rebuilds_from_storage_not_the_cache_file() {
let dir = TempDir::new().unwrap();
let g = Graph::open(dir.path(), 1).unwrap();
let a = g.add_node("N", &json!({})).unwrap();
let b = g.add_node("N", &json!({})).unwrap();
let c = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({})).unwrap();
let stale_snap = CsrSnapshot::build(&g.storage).unwrap();
g.add_edge(b, c, "R", &json!({})).unwrap();
let persisted_gen = {
let rtxn = g.storage.env.read_txn().unwrap();
crate::storage::ids::commit_gen(&g.storage, &rtxn).unwrap()
};
save_csr(dir.path(), &stale_snap, g.storage.db_id, persisted_gen).unwrap();
g.rebuild_csr().unwrap();
assert_eq!(
g.shortest_path(a, c).unwrap(),
Some(vec![a, b, c]),
"rebuild_csr must not answer from the cache file"
);
drop(g);
let g = Graph::open(dir.path(), 1).unwrap();
assert_eq!(g.shortest_path(a, c).unwrap(), Some(vec![a, b, c]));
}
#[test]
fn arrays_round_trip_exactly() {
let (_dir, g) = open_tmp();
let a = g.add_node("N", &json!({})).unwrap();
let b = g.add_node("N", &json!({})).unwrap();
let c = g.add_node("N", &json!({})).unwrap();
g.add_edge(a, b, "R", &json!({ "weight": 2.5 })).unwrap();
g.add_edge(b, c, "S", &json!({ "weight": -1.0 })).unwrap();
g.add_edge(a, c, "R", &json!({})).unwrap();
let snap = CsrSnapshot::build_weighted(&g.storage).unwrap();
let out = TempDir::new().unwrap();
let db_id = g.storage.db_id;
save_csr(out.path(), &snap, db_id, 7).unwrap();
assert!(
load_csr(out.path(), db_id, 8, false).is_none(),
"wrong generation"
);
assert!(
load_csr(out.path(), [0xAB; 16], 7, false).is_none(),
"wrong database identity"
);
let loaded = load_csr(out.path(), db_id, 7, true).expect("fresh and weighted");
assert_eq!(loaded.row_ptr, snap.row_ptr);
assert_eq!(loaded.col_idx, snap.col_idx);
assert_eq!(loaded.edge_type, snap.edge_type);
assert_eq!(loaded.edge_id, snap.edge_id);
assert_eq!(loaded.edge_weight, snap.edge_weight);
assert_eq!(loaded.has_negative_weight, snap.has_negative_weight);
assert!(loaded.has_negative_weight);
assert_eq!(loaded.in_row_ptr, snap.in_row_ptr);
assert_eq!(loaded.in_col_idx, snap.in_col_idx);
assert_eq!(loaded.in_edge_type, snap.in_edge_type);
assert_eq!(loaded.in_edge_id, snap.in_edge_id);
assert_eq!(loaded.dense_to_id, snap.dense_to_id);
assert_eq!(loaded.id_to_dense, snap.id_to_dense);
assert!(load_csr(out.path(), db_id, 7, false).is_some());
}
}