use std::path::Path;
use sparrowdb_common::{Error, Result};
use super::codec::{WalRecord, WAL_FORMAT_VERSION, WAL_FORMAT_VERSION_LEGACY};
use super::writer::segment_path;
#[derive(Debug, Clone)]
pub struct MigrationResult {
pub segments_inspected: usize,
pub segments_converted: usize,
pub segments_skipped: usize,
pub records_converted: usize,
}
pub fn migrate_wal(wal_dir: &Path) -> Result<MigrationResult> {
let segments = collect_segments(wal_dir)?;
let mut result = MigrationResult {
segments_inspected: segments.len(),
segments_converted: 0,
segments_skipped: 0,
records_converted: 0,
};
for seg_no in &segments {
let path = segment_path(wal_dir, *seg_no);
let data = std::fs::read(&path).map_err(Error::Io)?;
if data.is_empty() {
result.segments_skipped += 1;
continue;
}
let version = data[0];
match version {
WAL_FORMAT_VERSION => {
result.segments_skipped += 1;
}
WAL_FORMAT_VERSION_LEGACY => {
let (new_data, record_count) = convert_segment(&data)?;
result.records_converted += record_count;
result.segments_converted += 1;
let tmp_path = path.with_extension("wal.migrating");
std::fs::write(&tmp_path, &new_data).map_err(Error::Io)?;
std::fs::rename(&tmp_path, &path).map_err(Error::Io)?;
}
other => {
return Err(Error::Corruption(format!(
"WAL segment {} has unrecognised version byte {other}. \
Expected {WAL_FORMAT_VERSION} (current) or \
{WAL_FORMAT_VERSION_LEGACY} (legacy 0.1.2).",
seg_no
)));
}
}
}
Ok(result)
}
fn convert_segment(data: &[u8]) -> Result<(Vec<u8>, usize)> {
debug_assert_eq!(data[0], WAL_FORMAT_VERSION_LEGACY);
let mut out = Vec::with_capacity(data.len());
out.push(WAL_FORMAT_VERSION);
let mut offset = 1usize; let mut record_count = 0usize;
while offset < data.len() {
if data[offset..].iter().all(|&b| b == 0) {
out.extend_from_slice(&data[offset..]);
break;
}
let (record, consumed) =
WalRecord::decode_with_version(&data[offset..], WAL_FORMAT_VERSION_LEGACY)?;
let encoded = record.encode();
out.extend_from_slice(&encoded);
offset += consumed;
record_count += 1;
}
Ok((out, record_count))
}
fn collect_segments(wal_dir: &Path) -> Result<Vec<u64>> {
let mut segments = Vec::new();
let entries = std::fs::read_dir(wal_dir).map_err(Error::Io)?;
for entry in entries.flatten() {
let name = entry.file_name();
let name = name.to_string_lossy().to_string();
if name.starts_with("segment-") && name.ends_with(".wal") {
let num_str = &name["segment-".len()..name.len() - ".wal".len()];
if let Ok(n) = num_str.parse::<u64>() {
segments.push(n);
}
}
}
segments.sort();
Ok(segments)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::wal::codec::{WalPayload, WalRecordKind};
use sparrowdb_common::{Lsn, TxnId};
use tempfile::TempDir;
fn build_legacy_segment(records: &[WalRecord]) -> Vec<u8> {
let mut buf = Vec::new();
buf.push(WAL_FORMAT_VERSION_LEGACY);
for rec in records {
let payload_bytes = rec.payload.encode();
let payload_len = payload_bytes.len();
let length_val = (1u32 + 8 + 8) + payload_len as u32 + 4u32;
buf.extend_from_slice(&length_val.to_le_bytes());
buf.push(rec.kind.as_byte());
buf.extend_from_slice(&rec.lsn.0.to_le_bytes());
buf.extend_from_slice(&rec.txn_id.0.to_le_bytes());
buf.extend_from_slice(&payload_bytes);
let crc_start = buf.len() - (1 + 8 + 8 + payload_len);
let crc = crc32fast::hash(&buf[crc_start..]);
buf.extend_from_slice(&crc.to_le_bytes());
}
buf
}
fn sample_records() -> Vec<WalRecord> {
vec![
WalRecord {
lsn: Lsn(1),
txn_id: TxnId(100),
kind: WalRecordKind::Begin,
payload: WalPayload::Empty,
},
WalRecord {
lsn: Lsn(2),
txn_id: TxnId(100),
kind: WalRecordKind::NodeCreate,
payload: WalPayload::NodeCreate {
node_id: 42,
label_id: 1,
props: vec![("name".to_string(), b"Alice".to_vec())],
},
},
WalRecord {
lsn: Lsn(3),
txn_id: TxnId(100),
kind: WalRecordKind::Commit,
payload: WalPayload::Empty,
},
]
}
#[test]
fn test_legacy_segment_roundtrip_through_decode_with_version() {
let records = sample_records();
let seg_data = build_legacy_segment(&records);
let mut offset = 1; for expected in &records {
let (decoded, consumed) =
WalRecord::decode_with_version(&seg_data[offset..], WAL_FORMAT_VERSION_LEGACY)
.expect("should decode legacy record");
assert_eq!(decoded.lsn, expected.lsn);
assert_eq!(decoded.kind, expected.kind);
offset += consumed;
}
}
#[test]
fn test_migrate_converts_legacy_to_current() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = sample_records();
let seg_data = build_legacy_segment(&records);
let seg_path = segment_path(&wal_dir, 0);
std::fs::write(&seg_path, &seg_data).unwrap();
let result = migrate_wal(&wal_dir).unwrap();
assert_eq!(result.segments_inspected, 1);
assert_eq!(result.segments_converted, 1);
assert_eq!(result.segments_skipped, 0);
assert_eq!(result.records_converted, 3);
let migrated = std::fs::read(&seg_path).unwrap();
assert_eq!(migrated[0], WAL_FORMAT_VERSION);
let mut offset = 1;
for expected in &records {
let (decoded, consumed) =
WalRecord::decode_with_version(&migrated[offset..], WAL_FORMAT_VERSION)
.expect("should decode migrated record with CRC32C");
assert_eq!(decoded.lsn, expected.lsn);
assert_eq!(decoded.kind, expected.kind);
offset += consumed;
}
}
#[test]
fn test_migrate_skips_already_current() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = sample_records();
let mut v2_data = vec![WAL_FORMAT_VERSION];
for rec in &records {
v2_data.extend_from_slice(&rec.encode());
}
let seg_path = segment_path(&wal_dir, 0);
std::fs::write(&seg_path, &v2_data).unwrap();
let result = migrate_wal(&wal_dir).unwrap();
assert_eq!(result.segments_inspected, 1);
assert_eq!(result.segments_converted, 0);
assert_eq!(result.segments_skipped, 1);
assert_eq!(result.records_converted, 0);
let after = std::fs::read(&seg_path).unwrap();
assert_eq!(after, v2_data);
}
#[test]
fn test_migrate_idempotent() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = sample_records();
let seg_data = build_legacy_segment(&records);
let seg_path = segment_path(&wal_dir, 0);
std::fs::write(&seg_path, &seg_data).unwrap();
let r1 = migrate_wal(&wal_dir).unwrap();
assert_eq!(r1.segments_converted, 1);
let after_first = std::fs::read(&seg_path).unwrap();
let r2 = migrate_wal(&wal_dir).unwrap();
assert_eq!(r2.segments_converted, 0);
assert_eq!(r2.segments_skipped, 1);
let after_second = std::fs::read(&seg_path).unwrap();
assert_eq!(after_first, after_second);
}
#[test]
fn test_migrate_corrupt_segment_errors() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = sample_records();
let mut seg_data = build_legacy_segment(&records);
if seg_data.len() > 15 {
seg_data[15] ^= 0xFF;
}
let seg_path = segment_path(&wal_dir, 0);
std::fs::write(&seg_path, &seg_data).unwrap();
let err = migrate_wal(&wal_dir).unwrap_err();
match err {
Error::ChecksumMismatch | Error::Corruption(_) => {} other => panic!("expected checksum/corruption error, got: {other:?}"),
}
}
#[test]
fn test_migrate_mixed_segments() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = sample_records();
let legacy_data = build_legacy_segment(&records);
std::fs::write(segment_path(&wal_dir, 0), &legacy_data).unwrap();
let mut v2_data = vec![WAL_FORMAT_VERSION];
for rec in &records {
v2_data.extend_from_slice(&rec.encode());
}
std::fs::write(segment_path(&wal_dir, 1), &v2_data).unwrap();
let result = migrate_wal(&wal_dir).unwrap();
assert_eq!(result.segments_inspected, 2);
assert_eq!(result.segments_converted, 1);
assert_eq!(result.segments_skipped, 1);
}
#[test]
fn test_migrate_preserves_edge_create_records() {
let dir = TempDir::new().unwrap();
let wal_dir = dir.path().join("wal");
std::fs::create_dir_all(&wal_dir).unwrap();
let records = vec![
WalRecord {
lsn: Lsn(1),
txn_id: TxnId(200),
kind: WalRecordKind::Begin,
payload: WalPayload::Empty,
},
WalRecord {
lsn: Lsn(2),
txn_id: TxnId(200),
kind: WalRecordKind::EdgeCreate,
payload: WalPayload::EdgeCreate {
edge_id: 99,
src: 1,
dst: 2,
rel_type: "FOLLOWS".to_string(),
props: vec![("weight".to_string(), 42u64.to_le_bytes().to_vec())],
},
},
WalRecord {
lsn: Lsn(3),
txn_id: TxnId(200),
kind: WalRecordKind::Commit,
payload: WalPayload::Empty,
},
];
let seg_data = build_legacy_segment(&records);
std::fs::write(segment_path(&wal_dir, 0), &seg_data).unwrap();
let result = migrate_wal(&wal_dir).unwrap();
assert_eq!(result.records_converted, 3);
let migrated = std::fs::read(segment_path(&wal_dir, 0)).unwrap();
let mut offset = 1;
let (_, consumed) =
WalRecord::decode_with_version(&migrated[offset..], WAL_FORMAT_VERSION).unwrap();
offset += consumed;
let (edge_rec, _) =
WalRecord::decode_with_version(&migrated[offset..], WAL_FORMAT_VERSION).unwrap();
assert_eq!(edge_rec.kind, WalRecordKind::EdgeCreate);
match edge_rec.payload {
WalPayload::EdgeCreate {
edge_id,
src,
dst,
rel_type,
props,
} => {
assert_eq!(edge_id, 99);
assert_eq!(src, 1);
assert_eq!(dst, 2);
assert_eq!(rel_type, "FOLLOWS");
assert_eq!(props.len(), 1);
assert_eq!(props[0].0, "weight");
}
other => panic!("expected EdgeCreate, got: {other:?}"),
}
}
}