use crate::error::{GraphError, Result};
use crate::graph::Id;
use std::fs::{File, OpenOptions};
use std::io::{BufReader, BufWriter, Read, Seek, SeekFrom, Write};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU64, Ordering};
const WAL_MAGIC: &[u8; 4] = b"GWAL";
const WAL_VERSION: u32 = 1;
const HEADER_SIZE: usize = 4 + 4 + 8 + 8;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum SyncMode {
Immediate,
Batched(usize),
OsManaged,
}
impl Default for SyncMode {
fn default() -> Self {
SyncMode::Batched(100)
}
}
#[derive(Debug, Clone)]
pub struct WalConfig {
pub sync_mode: SyncMode,
pub max_size_bytes: u64,
pub use_fsync: bool,
pub checkpoint_interval: Option<u64>,
}
impl Default for WalConfig {
fn default() -> Self {
WalConfig {
sync_mode: SyncMode::default(),
max_size_bytes: 64 * 1024 * 1024, use_fsync: true,
checkpoint_interval: Some(10_000),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum WalOperation {
CreateNode {
id: Id,
data: Vec<u8>,
},
UpdateNode {
id: Id,
data: Vec<u8>,
},
DeleteNode {
id: Id,
},
CreateRelationship {
id: Id,
from_id: Id,
to_id: Id,
rel_type: String,
data: Vec<u8>,
},
UpdateRelationship {
id: Id,
data: Vec<u8>,
},
DeleteRelationship {
id: Id,
},
BeginTransaction {
tx_id: u64,
},
CommitTransaction {
tx_id: u64,
},
RollbackTransaction {
tx_id: u64,
},
Checkpoint {
sequence: u64,
timestamp: u64,
},
}
#[derive(Debug, Clone)]
pub struct WalEntry {
pub sequence: u64,
pub timestamp: u64,
pub operation: WalOperation,
pub checksum: u32,
}
impl WalEntry {
pub fn new(sequence: u64, operation: WalOperation) -> Self {
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
let mut entry = WalEntry {
sequence,
timestamp,
operation,
checksum: 0,
};
entry.checksum = entry.calculate_checksum();
entry
}
fn calculate_checksum(&self) -> u32 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
self.sequence.hash(&mut hasher);
self.timestamp.hash(&mut hasher);
format!("{:?}", self.operation).hash(&mut hasher);
hasher.finish() as u32
}
pub fn is_valid(&self) -> bool {
self.checksum == self.calculate_checksum()
}
pub fn serialize(&self) -> Vec<u8> {
let op_data = self.serialize_operation();
let total_len = 8 + 8 + op_data.len() + 4;
let mut data = Vec::with_capacity(4 + total_len);
data.extend_from_slice(&(total_len as u32).to_le_bytes());
data.extend_from_slice(&self.sequence.to_le_bytes());
data.extend_from_slice(&self.timestamp.to_le_bytes());
data.extend_from_slice(&op_data);
data.extend_from_slice(&self.checksum.to_le_bytes());
data
}
fn serialize_operation(&self) -> Vec<u8> {
let mut data = Vec::new();
match &self.operation {
WalOperation::CreateNode {
id,
data: node_data,
} => {
data.push(0x01);
data.extend_from_slice(&id.to_le_bytes());
data.extend_from_slice(&(node_data.len() as u32).to_le_bytes());
data.extend_from_slice(node_data);
}
WalOperation::UpdateNode {
id,
data: node_data,
} => {
data.push(0x02);
data.extend_from_slice(&id.to_le_bytes());
data.extend_from_slice(&(node_data.len() as u32).to_le_bytes());
data.extend_from_slice(node_data);
}
WalOperation::DeleteNode { id } => {
data.push(0x03);
data.extend_from_slice(&id.to_le_bytes());
}
WalOperation::CreateRelationship {
id,
from_id,
to_id,
rel_type,
data: rel_data,
} => {
data.push(0x04);
data.extend_from_slice(&id.to_le_bytes());
data.extend_from_slice(&from_id.to_le_bytes());
data.extend_from_slice(&to_id.to_le_bytes());
let rel_type_bytes = rel_type.as_bytes();
data.extend_from_slice(&(rel_type_bytes.len() as u16).to_le_bytes());
data.extend_from_slice(rel_type_bytes);
data.extend_from_slice(&(rel_data.len() as u32).to_le_bytes());
data.extend_from_slice(rel_data);
}
WalOperation::UpdateRelationship { id, data: rel_data } => {
data.push(0x05);
data.extend_from_slice(&id.to_le_bytes());
data.extend_from_slice(&(rel_data.len() as u32).to_le_bytes());
data.extend_from_slice(rel_data);
}
WalOperation::DeleteRelationship { id } => {
data.push(0x06);
data.extend_from_slice(&id.to_le_bytes());
}
WalOperation::BeginTransaction { tx_id } => {
data.push(0x10);
data.extend_from_slice(&tx_id.to_le_bytes());
}
WalOperation::CommitTransaction { tx_id } => {
data.push(0x11);
data.extend_from_slice(&tx_id.to_le_bytes());
}
WalOperation::RollbackTransaction { tx_id } => {
data.push(0x12);
data.extend_from_slice(&tx_id.to_le_bytes());
}
WalOperation::Checkpoint {
sequence,
timestamp,
} => {
data.push(0x20);
data.extend_from_slice(&sequence.to_le_bytes());
data.extend_from_slice(×tamp.to_le_bytes());
}
}
data
}
pub fn deserialize(data: &[u8]) -> Result<(Self, usize)> {
if data.len() < 4 {
return Err(GraphError::Storage("WAL entry too short".to_string()));
}
let len = u32::from_le_bytes([data[0], data[1], data[2], data[3]]) as usize;
if data.len() < 4 + len {
return Err(GraphError::Storage("WAL entry truncated".to_string()));
}
let entry_data = &data[4..4 + len];
let sequence = u64::from_le_bytes([
entry_data[0],
entry_data[1],
entry_data[2],
entry_data[3],
entry_data[4],
entry_data[5],
entry_data[6],
entry_data[7],
]);
let timestamp = u64::from_le_bytes([
entry_data[8],
entry_data[9],
entry_data[10],
entry_data[11],
entry_data[12],
entry_data[13],
entry_data[14],
entry_data[15],
]);
let (operation, op_len) = Self::deserialize_operation(&entry_data[16..])?;
let checksum_start = 16 + op_len;
let checksum = u32::from_le_bytes([
entry_data[checksum_start],
entry_data[checksum_start + 1],
entry_data[checksum_start + 2],
entry_data[checksum_start + 3],
]);
let entry = WalEntry {
sequence,
timestamp,
operation,
checksum,
};
if !entry.is_valid() {
return Err(GraphError::Storage(
"WAL entry checksum mismatch".to_string(),
));
}
Ok((entry, 4 + len))
}
fn deserialize_operation(data: &[u8]) -> Result<(WalOperation, usize)> {
if data.is_empty() {
return Err(GraphError::Storage("Empty operation data".to_string()));
}
let op_type = data[0];
let mut offset = 1;
let operation = match op_type {
0x01 => {
let id = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
let data_len = u32::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
]) as usize;
offset += 4;
let node_data = data[offset..offset + data_len].to_vec();
offset += data_len;
WalOperation::CreateNode {
id,
data: node_data,
}
}
0x03 => {
let id = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
WalOperation::DeleteNode { id }
}
0x10 => {
let tx_id = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
WalOperation::BeginTransaction { tx_id }
}
0x11 => {
let tx_id = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
WalOperation::CommitTransaction { tx_id }
}
0x20 => {
let seq = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
let ts = u64::from_le_bytes([
data[offset],
data[offset + 1],
data[offset + 2],
data[offset + 3],
data[offset + 4],
data[offset + 5],
data[offset + 6],
data[offset + 7],
]);
offset += 8;
WalOperation::Checkpoint {
sequence: seq,
timestamp: ts,
}
}
_ => {
return Err(GraphError::Storage(format!(
"Unknown WAL operation type: {op_type}"
)));
}
};
Ok((operation, offset))
}
}
pub struct WriteAheadLog {
path: PathBuf,
writer: Option<BufWriter<File>>,
config: WalConfig,
sequence: AtomicU64,
unflushed_count: AtomicU64,
last_checkpoint: AtomicU64,
}
impl WriteAheadLog {
pub fn open<P: AsRef<Path>>(path: P) -> Result<Self> {
Self::open_with_config(path, WalConfig::default())
}
pub fn open_with_config<P: AsRef<Path>>(path: P, config: WalConfig) -> Result<Self> {
let path = path.as_ref().to_path_buf();
let file = OpenOptions::new()
.read(true)
.create(true)
.append(true)
.open(&path)
.map_err(|e| GraphError::Io(e.to_string()))?;
let mut wal = WriteAheadLog {
path,
writer: Some(BufWriter::new(file)),
config,
sequence: AtomicU64::new(0),
unflushed_count: AtomicU64::new(0),
last_checkpoint: AtomicU64::new(0),
};
wal.init_header()?;
Ok(wal)
}
fn init_header(&mut self) -> Result<()> {
let writer = self
.writer
.as_mut()
.ok_or_else(|| GraphError::Storage("WAL writer not available".to_string()))?;
let file = writer.get_ref();
let file_len = file
.metadata()
.map_err(|e| GraphError::Io(e.to_string()))?
.len();
if file_len == 0 {
writer
.write_all(WAL_MAGIC)
.map_err(|e| GraphError::Io(e.to_string()))?;
writer
.write_all(&WAL_VERSION.to_le_bytes())
.map_err(|e| GraphError::Io(e.to_string()))?;
writer
.write_all(&0u64.to_le_bytes()) .map_err(|e| GraphError::Io(e.to_string()))?;
writer
.write_all(&0u64.to_le_bytes()) .map_err(|e| GraphError::Io(e.to_string()))?;
writer.flush().map_err(|e| GraphError::Io(e.to_string()))?;
} else {
let recovered_seq = self.scan_for_sequence()?;
self.sequence.store(recovered_seq, Ordering::SeqCst);
}
Ok(())
}
fn scan_for_sequence(&self) -> Result<u64> {
let file = File::open(&self.path).map_err(|e| GraphError::Io(e.to_string()))?;
let mut reader = BufReader::new(file);
reader
.seek(SeekFrom::Start(HEADER_SIZE as u64))
.map_err(|e| GraphError::Io(e.to_string()))?;
let mut last_seq = 0u64;
let mut buffer = vec![0u8; 4];
while reader.read_exact(&mut buffer).is_ok() {
let len = u32::from_le_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
let mut entry_data = vec![0u8; len];
if reader.read_exact(&mut entry_data).is_err() {
break; }
if len >= 8 {
last_seq = u64::from_le_bytes([
entry_data[0],
entry_data[1],
entry_data[2],
entry_data[3],
entry_data[4],
entry_data[5],
entry_data[6],
entry_data[7],
]);
}
}
Ok(last_seq)
}
pub fn append(&mut self, operation: WalOperation) -> Result<u64> {
let seq = self.sequence.fetch_add(1, Ordering::SeqCst) + 1;
let entry = WalEntry::new(seq, operation);
let writer = self
.writer
.as_mut()
.ok_or_else(|| GraphError::Storage("WAL writer not available".to_string()))?;
let data = entry.serialize();
writer
.write_all(&data)
.map_err(|e| GraphError::Io(e.to_string()))?;
let count = self.unflushed_count.fetch_add(1, Ordering::SeqCst) + 1;
match self.config.sync_mode {
SyncMode::Immediate => {
self.sync()?;
}
SyncMode::Batched(batch_size) if count as usize >= batch_size => {
self.sync()?;
}
_ => {}
}
Ok(seq)
}
pub fn sync(&mut self) -> Result<()> {
let writer = self
.writer
.as_mut()
.ok_or_else(|| GraphError::Storage("WAL writer not available".to_string()))?;
writer.flush().map_err(|e| GraphError::Io(e.to_string()))?;
if self.config.use_fsync {
writer
.get_ref()
.sync_all()
.map_err(|e| GraphError::Io(e.to_string()))?;
} else {
writer
.get_ref()
.sync_data()
.map_err(|e| GraphError::Io(e.to_string()))?;
}
self.unflushed_count.store(0, Ordering::SeqCst);
Ok(())
}
pub fn recover(&self) -> Result<Vec<WalEntry>> {
let file = File::open(&self.path).map_err(|e| GraphError::Io(e.to_string()))?;
let mut reader = BufReader::new(file);
reader
.seek(SeekFrom::Start(HEADER_SIZE as u64))
.map_err(|e| GraphError::Io(e.to_string()))?;
let mut entries = Vec::new();
let mut buffer = Vec::new();
reader
.read_to_end(&mut buffer)
.map_err(|e| GraphError::Io(e.to_string()))?;
let mut offset = 0;
while offset < buffer.len() {
match WalEntry::deserialize(&buffer[offset..]) {
Ok((entry, consumed)) => {
entries.push(entry);
offset += consumed;
}
Err(_) => {
break;
}
}
}
Ok(entries)
}
pub fn checkpoint(&mut self) -> Result<u64> {
let seq = self.sequence.load(Ordering::SeqCst);
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
self.append(WalOperation::Checkpoint {
sequence: seq,
timestamp,
})?;
self.last_checkpoint.store(seq, Ordering::SeqCst);
self.sync()?;
Ok(seq)
}
pub fn current_sequence(&self) -> u64 {
self.sequence.load(Ordering::SeqCst)
}
pub fn path(&self) -> &Path {
&self.path
}
pub fn truncate_before(&mut self, _sequence: u64) -> Result<()> {
Ok(())
}
}
impl Drop for WriteAheadLog {
fn drop(&mut self) {
let _ = self.sync();
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[test]
fn test_entry_serialization() {
let entry = WalEntry::new(
1,
WalOperation::CreateNode {
id: 42,
data: vec![1, 2, 3, 4],
},
);
let serialized = entry.serialize();
let (deserialized, _) = WalEntry::deserialize(&serialized).unwrap();
assert_eq!(deserialized.sequence, entry.sequence);
assert!(deserialized.is_valid());
}
#[test]
fn test_append_and_recover() {
let temp_file = NamedTempFile::new().unwrap();
let path = temp_file.path();
{
let mut wal = WriteAheadLog::open(path).unwrap();
wal.append(WalOperation::CreateNode {
id: 1,
data: vec![1, 2, 3],
})
.unwrap();
wal.append(WalOperation::CreateNode {
id: 2,
data: vec![4, 5, 6],
})
.unwrap();
wal.sync().unwrap();
}
{
let wal = WriteAheadLog::open(path).unwrap();
let entries = wal.recover().unwrap();
assert_eq!(entries.len(), 2);
assert_eq!(entries[0].sequence, 1);
assert_eq!(entries[1].sequence, 2);
}
}
#[test]
fn test_transaction_markers() {
let temp_file = NamedTempFile::new().unwrap();
let path = temp_file.path();
let mut wal = WriteAheadLog::open(path).unwrap();
wal.append(WalOperation::BeginTransaction { tx_id: 100 })
.unwrap();
wal.append(WalOperation::CreateNode {
id: 1,
data: vec![],
})
.unwrap();
wal.append(WalOperation::CommitTransaction { tx_id: 100 })
.unwrap();
wal.sync().unwrap();
let entries = wal.recover().unwrap();
assert_eq!(entries.len(), 3);
match &entries[0].operation {
WalOperation::BeginTransaction { tx_id } => assert_eq!(*tx_id, 100),
_ => panic!("Expected BeginTransaction"),
}
match &entries[2].operation {
WalOperation::CommitTransaction { tx_id } => assert_eq!(*tx_id, 100),
_ => panic!("Expected CommitTransaction"),
}
}
#[test]
fn test_checkpoint() {
let temp_file = NamedTempFile::new().unwrap();
let path = temp_file.path();
let mut wal = WriteAheadLog::open(path).unwrap();
wal.append(WalOperation::CreateNode {
id: 1,
data: vec![],
})
.unwrap();
let checkpoint_seq = wal.checkpoint().unwrap();
assert!(checkpoint_seq > 0);
let entries = wal.recover().unwrap();
let checkpoint_entry = entries
.iter()
.find(|e| matches!(e.operation, WalOperation::Checkpoint { .. }));
assert!(checkpoint_entry.is_some());
}
#[test]
fn test_sequence_recovery() {
let temp_file = NamedTempFile::new().unwrap();
let path = temp_file.path();
{
let mut wal = WriteAheadLog::open(path).unwrap();
for i in 0..5 {
wal.append(WalOperation::CreateNode {
id: i,
data: vec![],
})
.unwrap();
}
wal.sync().unwrap();
}
{
let mut wal = WriteAheadLog::open(path).unwrap();
let seq = wal
.append(WalOperation::CreateNode {
id: 5,
data: vec![],
})
.unwrap();
assert_eq!(seq, 6); }
}
}