use crate::error::{GraphError, Result};
use crate::graph::{Id, Node, Relationship};
use memmap2::{MmapMut, MmapOptions};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::fs::OpenOptions;
use std::path::Path;
use std::sync::Arc;
#[allow(unsafe_code)]
pub struct FileHeader {
pub magic: [u8; 4],
pub version: u32,
pub node_count: u64,
pub rel_count: u64,
}
impl Default for FileHeader {
fn default() -> Self {
Self::new()
}
}
impl FileHeader {
const MAGIC: &'static [u8; 4] = b"GRPH";
const VERSION: u32 = 1;
const SIZE: usize = 4 + 4 + 8 + 8;
pub fn new() -> Self {
FileHeader {
magic: *Self::MAGIC,
version: Self::VERSION,
node_count: 0,
rel_count: 0,
}
}
pub fn is_valid(&self) -> bool {
self.magic == *Self::MAGIC && self.version == Self::VERSION
}
pub fn serialize(&self) -> Vec<u8> {
let mut data = Vec::with_capacity(Self::SIZE);
data.extend_from_slice(&self.magic);
data.extend_from_slice(&self.version.to_le_bytes());
data.extend_from_slice(&self.node_count.to_le_bytes());
data.extend_from_slice(&self.rel_count.to_le_bytes());
data
}
pub fn deserialize(data: &[u8]) -> Result<Self> {
if data.len() < Self::SIZE {
return Err(GraphError::Storage("Invalid header size".to_string()));
}
let mut magic = [0u8; 4];
magic.copy_from_slice(&data[0..4]);
let version = u32::from_le_bytes([data[4], data[5], data[6], data[7]]);
let node_count = u64::from_le_bytes([
data[8], data[9], data[10], data[11], data[12], data[13], data[14], data[15],
]);
let rel_count = u64::from_le_bytes([
data[16], data[17], data[18], data[19], data[20], data[21], data[22], data[23],
]);
Ok(FileHeader {
magic,
version,
node_count,
rel_count,
})
}
}
pub struct MmapStorage {
mmap: Arc<RwLock<MmapMut>>,
file_path: std::path::PathBuf,
node_cache: Arc<RwLock<HashMap<Id, Node>>>,
relationship_cache: Arc<RwLock<HashMap<Id, Relationship>>>,
node_relationships: Arc<RwLock<HashMap<Id, Vec<Id>>>>,
}
impl MmapStorage {
pub fn create<P: AsRef<Path>>(path: P) -> Result<Self> {
let file_path = path.as_ref().to_path_buf();
let initial_size = 1024 * 1024; let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(true)
.open(&file_path)?;
file.set_len(initial_size)?;
#[allow(unsafe_code)]
let mut mmap = unsafe { MmapOptions::new().map_mut(&file)? };
let header = FileHeader::new();
let header_data = header.serialize();
mmap[0..FileHeader::SIZE].copy_from_slice(&header_data);
Ok(MmapStorage {
mmap: Arc::new(RwLock::new(mmap)),
file_path,
node_cache: Arc::new(RwLock::new(HashMap::new())),
relationship_cache: Arc::new(RwLock::new(HashMap::new())),
node_relationships: Arc::new(RwLock::new(HashMap::new())),
})
}
pub fn open<P: AsRef<Path>>(path: P) -> Result<Self> {
let file_path = path.as_ref().to_path_buf();
if !file_path.exists() {
return Self::create(path);
}
let file = OpenOptions::new().read(true).write(true).open(&file_path)?;
#[allow(unsafe_code)]
let mmap = unsafe { MmapOptions::new().map_mut(&file)? };
if mmap.len() < FileHeader::SIZE {
return Err(GraphError::Storage(
"File too small to contain header".to_string(),
));
}
let header = FileHeader::deserialize(&mmap[0..FileHeader::SIZE])?;
if !header.is_valid() {
return Err(GraphError::Storage("Invalid file header".to_string()));
}
let storage = MmapStorage {
mmap: Arc::new(RwLock::new(mmap)),
file_path,
node_cache: Arc::new(RwLock::new(HashMap::new())),
relationship_cache: Arc::new(RwLock::new(HashMap::new())),
node_relationships: Arc::new(RwLock::new(HashMap::new())),
};
storage.load_caches()?;
Ok(storage)
}
fn load_caches(&self) -> Result<()> {
let mmap = self.mmap.read();
let header = FileHeader::deserialize(&mmap[0..FileHeader::SIZE])?;
let mut offset = FileHeader::SIZE;
for _ in 0..header.node_count {
if offset + 8 > mmap.len() {
break; }
let data_len = u64::from_le_bytes([
mmap[offset],
mmap[offset + 1],
mmap[offset + 2],
mmap[offset + 3],
mmap[offset + 4],
mmap[offset + 5],
mmap[offset + 6],
mmap[offset + 7],
]) as usize;
offset += 8;
if offset + data_len > mmap.len() {
break; }
let node_data = &mmap[offset..offset + data_len];
if let Ok(json_str) = std::str::from_utf8(node_data) {
if let Ok(node) = serde_json::from_str::<Node>(json_str) {
self.node_cache.write().insert(node.id, node.clone());
self.node_relationships.write().entry(node.id).or_default();
}
}
offset += data_len;
}
for _ in 0..header.rel_count {
if offset + 8 > mmap.len() {
break;
}
let data_len = u64::from_le_bytes([
mmap[offset],
mmap[offset + 1],
mmap[offset + 2],
mmap[offset + 3],
mmap[offset + 4],
mmap[offset + 5],
mmap[offset + 6],
mmap[offset + 7],
]) as usize;
offset += 8;
if offset + data_len > mmap.len() {
break;
}
let rel_data = &mmap[offset..offset + data_len];
if let Ok(json_str) = std::str::from_utf8(rel_data) {
if let Ok(rel) = serde_json::from_str::<Relationship>(json_str) {
let from_id = rel.from_id;
let to_id = rel.to_id;
let rel_id = rel.id;
self.relationship_cache.write().insert(rel_id, rel);
self.node_relationships
.write()
.entry(from_id)
.or_default()
.push(rel_id);
if from_id != to_id {
self.node_relationships
.write()
.entry(to_id)
.or_default()
.push(rel_id);
}
}
}
offset += data_len;
}
Ok(())
}
pub fn store_node(&self, node: Node) -> Result<()> {
self.node_cache.write().insert(node.id, node.clone());
self.node_relationships.write().entry(node.id).or_default();
Ok(())
}
pub fn get_node(&self, id: Id) -> Result<Option<Node>> {
Ok(self.node_cache.read().get(&id).cloned())
}
pub fn store_relationship(&self, relationship: Relationship) -> Result<()> {
let from_id = relationship.from_id;
let to_id = relationship.to_id;
let rel_id = relationship.id;
if !self.node_cache.read().contains_key(&from_id) {
return Err(GraphError::NotFound(format!("Node {from_id} not found")));
}
if !self.node_cache.read().contains_key(&to_id) {
return Err(GraphError::NotFound(format!("Node {to_id} not found")));
}
self.relationship_cache.write().insert(rel_id, relationship);
self.node_relationships
.write()
.entry(from_id)
.or_default()
.push(rel_id);
if from_id != to_id {
self.node_relationships
.write()
.entry(to_id)
.or_default()
.push(rel_id);
}
Ok(())
}
pub fn get_relationship(&self, id: Id) -> Result<Option<Relationship>> {
Ok(self.relationship_cache.read().get(&id).cloned())
}
pub fn get_relationships_for_node(&self, node_id: Id) -> Result<Vec<Relationship>> {
let empty_vec = Vec::new();
let relationship_ids = self
.node_relationships
.read()
.get(&node_id)
.unwrap_or(&empty_vec)
.clone();
let mut relationships = Vec::new();
let rel_cache = self.relationship_cache.read();
for rel_id in relationship_ids {
if let Some(rel) = rel_cache.get(&rel_id) {
relationships.push(rel.clone());
}
}
Ok(relationships)
}
pub fn node_count(&self) -> usize {
self.node_cache.read().len()
}
pub fn relationship_count(&self) -> usize {
self.relationship_cache.read().len()
}
pub fn delete_node(&self, id: Id) -> Result<Option<Node>> {
self.node_relationships.write().remove(&id);
Ok(self.node_cache.write().remove(&id))
}
pub fn delete_relationship(&self, id: Id) -> Result<Option<Relationship>> {
if let Some(rel) = self.relationship_cache.write().remove(&id) {
let mut node_rels = self.node_relationships.write();
if let Some(rels) = node_rels.get_mut(&rel.from_id) {
rels.retain(|&r| r != id);
}
if rel.from_id != rel.to_id {
if let Some(rels) = node_rels.get_mut(&rel.to_id) {
rels.retain(|&r| r != id);
}
}
Ok(Some(rel))
} else {
Ok(None)
}
}
pub fn node_has_relationships(&self, node_id: Id) -> bool {
self.node_relationships
.read()
.get(&node_id)
.map(|rels| !rels.is_empty())
.unwrap_or(false)
}
pub fn flush(&self) -> Result<()> {
let mut mmap = self.mmap.write();
let mut offset = FileHeader::SIZE;
let nodes = self.node_cache.read();
let relationships = self.relationship_cache.read();
let mut total_size = FileHeader::SIZE;
for node in nodes.values() {
if let Ok(json_str) = serde_json::to_string(node) {
total_size += 8 + json_str.len(); }
}
for rel in relationships.values() {
if let Ok(json_str) = serde_json::to_string(rel) {
total_size += 8 + json_str.len(); }
}
if total_size > mmap.len() {
drop(mmap);
let file = OpenOptions::new()
.read(true)
.write(true)
.open(&self.file_path)?;
file.set_len(total_size as u64 * 2)?;
#[allow(unsafe_code)]
let new_mmap = unsafe { MmapOptions::new().map_mut(&file)? };
*self.mmap.write() = new_mmap;
mmap = self.mmap.write();
}
for node in nodes.values() {
if let Ok(json_str) = serde_json::to_string(node) {
let serialized = json_str.as_bytes();
if offset + 8 + serialized.len() > mmap.len() {
return Err(GraphError::Storage("Not enough space in mmap".to_string()));
}
let len_bytes = (serialized.len() as u64).to_le_bytes();
mmap[offset..offset + 8].copy_from_slice(&len_bytes);
offset += 8;
mmap[offset..offset + serialized.len()].copy_from_slice(serialized);
offset += serialized.len();
}
}
for rel in relationships.values() {
if let Ok(json_str) = serde_json::to_string(rel) {
let serialized = json_str.as_bytes();
if offset + 8 + serialized.len() > mmap.len() {
return Err(GraphError::Storage("Not enough space in mmap".to_string()));
}
let len_bytes = (serialized.len() as u64).to_le_bytes();
mmap[offset..offset + 8].copy_from_slice(&len_bytes);
offset += 8;
mmap[offset..offset + serialized.len()].copy_from_slice(serialized);
offset += serialized.len();
}
}
let mut header = FileHeader::new();
header.node_count = nodes.len() as u64;
header.rel_count = relationships.len() as u64;
let header_data = header.serialize();
mmap[0..FileHeader::SIZE].copy_from_slice(&header_data);
mmap.flush()?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use tempfile::NamedTempFile;
#[test]
fn test_mmap_storage_creation() {
let temp_file = NamedTempFile::new().unwrap();
let storage = MmapStorage::create(temp_file.path()).unwrap();
assert_eq!(storage.node_count(), 0);
assert_eq!(storage.relationship_count(), 0);
}
#[test]
fn test_mmap_node_operations() {
let temp_file = NamedTempFile::new().unwrap();
let storage = MmapStorage::create(temp_file.path()).unwrap();
let node = Node::new(1, HashMap::new());
storage.store_node(node.clone()).unwrap();
let retrieved = storage.get_node(1).unwrap();
assert_eq!(retrieved, Some(node));
assert_eq!(storage.node_count(), 1);
}
#[test]
fn test_mmap_persistence() {
let temp_file = NamedTempFile::new().unwrap();
let path = temp_file.path().to_path_buf();
{
let storage = MmapStorage::create(&path).unwrap();
let mut properties = HashMap::new();
properties.insert("name".to_string(), serde_json::json!("Alice"));
properties.insert("age".to_string(), serde_json::json!(30));
let node = Node::new(1, properties);
storage.store_node(node).unwrap();
assert_eq!(storage.node_count(), 1);
storage.flush().unwrap();
let retrieved = storage.get_node(1).unwrap();
assert!(retrieved.is_some());
let retrieved_node = retrieved.unwrap();
assert_eq!(retrieved_node.id, 1);
assert_eq!(
retrieved_node.get_property("name"),
Some(&serde_json::json!("Alice"))
);
}
{
let storage = MmapStorage::open(&path).unwrap();
assert_eq!(storage.node_count(), 1, "Data should persist after reopen");
let retrieved = storage.get_node(1).unwrap();
assert!(retrieved.is_some());
let retrieved_node = retrieved.unwrap();
assert_eq!(retrieved_node.id, 1);
assert_eq!(
retrieved_node.get_property("name"),
Some(&serde_json::json!("Alice"))
);
assert_eq!(
retrieved_node.get_property("age"),
Some(&serde_json::json!(30))
);
}
}
}