#[cfg(test)]
mod tests;
use std::num::NonZeroU64;
use std::path::PathBuf;
use replay::WalReplayer;
use storage::{SegmentStorage, SegmentWriter};
use thiserror::Error;
use crate::paths::DbPaths;
use crate::serialization::SerializationError;
use crate::types::{BlobHash, CheckpointReason, CheckpointState, TypesError, WalOp};
use crate::{KeyBytes, calculate_blob_hash};
mod replay;
mod storage;
const INITIAL_SEGMENT_ID: u64 = 0;
const FIRST_OP_VERSION: NonZeroU64 = NonZeroU64::new(1).unwrap();
#[derive(Debug)]
pub(crate) struct WalAppendInfo {
pub version: NonZeroU64,
#[allow(unused)]
pub op_hash: BlobHash,
}
#[derive(Error, Debug)]
pub enum WalError {
#[error("Failed to parse segment ID from checkpoint metadata")]
ParseCheckpointMetaSegmentId(#[source] std::num::ParseIntError),
#[error("WAL IO error during {operation:?} (path: {path:?})")]
Io {
operation: WalIoOperation,
path: Option<PathBuf>,
#[source]
source: std::io::Error,
},
#[error("Failed to write WAL entry data (version {op_version}, segment {segment_id})")]
WriteWalEntryDataIO {
op_version: NonZeroU64,
segment_id: u64,
#[source]
source: std::io::Error,
},
#[error("Invalid operation version: {version} (must be > 0)")]
InvalidOpVersion { version: u64 },
#[error("WAL replay IO error during {step:?} for segment {segment_id} (path: {path:?})")]
ReplayIo {
step: WalReplayIoStep,
segment_id: u64,
path: PathBuf,
#[source]
source: std::io::Error,
},
#[error(
"WAL corruption: checksum mismatch for WAL entry (version {version}, segment {segment_id})"
)]
ReplayChecksumMismatch { version: u64, segment_id: u64, expected: BlobHash, actual: BlobHash },
#[error(
"WAL corruption: failed to deserialize WalOpRaw (version {version}, segment {segment_id})"
)]
ReplayDeserializeWalOpRaw {
version: NonZeroU64,
segment_id: u64,
#[source]
source: SerializationError,
},
#[error(
"WAL replay: failed to convert WalOpRaw to WalOp<K> (version {version}, segment {segment_id})"
)]
ReplayConvertWalOp {
version: NonZeroU64,
segment_id: u64,
#[source]
source: TypesError,
},
}
#[derive(Debug, Clone, Copy)]
pub enum WalIoOperation {
CreateDbDir,
ReadCheckpointMeta,
OpenSegmentWrite,
WriteSentinel,
CreateInitialFile,
SyncInitialFile,
FlushWriter,
SyncData,
ReadDbDirDiscovery,
ReadEntryDiscovery,
RemoveStaleSegment,
}
#[derive(Debug, Clone, Copy)]
pub enum WalReplayIoStep {
OpenSegment,
ReadHeader,
ReadOpData,
}
pub(crate) struct WalManager {
num_ops_per_wal: NonZeroU64,
next_op_version: NonZeroU64,
storage: SegmentStorage,
active_writer: Option<SegmentWriter>,
}
impl WalManager {
pub(crate) fn new(paths: DbPaths, num_ops_per_wal: NonZeroU64) -> Result<Self, WalError> {
std::fs::create_dir_all(paths.db_root_path()).map_err(|e| WalError::Io {
operation: WalIoOperation::CreateDbDir,
path: None,
source: e,
})?;
let storage = SegmentStorage::new(paths.clone());
Ok(WalManager {
num_ops_per_wal,
next_op_version: FIRST_OP_VERSION,
storage,
active_writer: None,
})
}
pub(crate) fn compute_checkpoint_target(
&self,
reason: CheckpointReason,
last_checkpointed_version: CheckpointState,
) -> Option<NonZeroU64> {
let should_checkpoint = match reason {
CheckpointReason::InitialSetup => last_checkpointed_version.is_none(),
CheckpointReason::AfterReplay | CheckpointReason::Explicit => true,
CheckpointReason::SegmentRollover => match last_checkpointed_version {
None => self.has_written_ops(),
Some(version) => self.has_new_ops_since(version),
},
};
if should_checkpoint { self.last_written_op_version() } else { None }
}
pub(crate) fn commit_checkpoint(
&mut self,
version: NonZeroU64,
last_checkpointed_version: CheckpointState,
) -> Result<(), WalError> {
if let Some(last) = last_checkpointed_version
&& version <= last
{
tracing::debug!(
last_checkpointed_version = last,
commit_version = version,
"Skipping checkpoint commit (idempotent/no-op).",
);
return Ok(());
}
let last_checkpointed_segment = self.segment_id_for_op_version(version.get());
let _ = self.storage.prune_stale_segments(last_checkpointed_segment);
tracing::info!(checkpoint_version = version, "Checkpoint committed");
Ok(())
}
fn has_written_ops(&self) -> bool {
self.next_op_version > FIRST_OP_VERSION
}
fn has_new_ops_since(&self, checkpoint_version: NonZeroU64) -> bool {
self.next_op_version > (checkpoint_version.saturating_add(1))
}
fn last_written_op_version(&self) -> Option<NonZeroU64> {
NonZeroU64::new(self.next_op_version.get().saturating_sub(1))
}
pub(crate) fn get_next_op_version(&self) -> NonZeroU64 {
self.next_op_version
}
pub(crate) fn allocate_next_op_version(&mut self) -> NonZeroU64 {
let current_version = self.next_op_version;
self.next_op_version = self.next_op_version.saturating_add(1);
current_version
}
pub(crate) fn replay_and_prepare<K>(
&mut self,
last_checkpointed_version: CheckpointState,
apply_op_fn: impl FnMut(WalOp<K>),
) -> Result<(), WalError>
where
K: KeyBytes + Clone + Eq + Ord + std::fmt::Debug + 'static,
{
let replayer = WalReplayer::new(&self.storage, last_checkpointed_version);
let highest_op_version = replayer.replay(apply_op_fn)?;
self.next_op_version = match highest_op_version {
Some(v) => v.saturating_add(1),
None => FIRST_OP_VERSION,
};
let next_op_version = self.get_next_op_version();
let target_segment_id = self.segment_id_for_op_version(next_op_version.get());
self.storage.ensure_segment_file_exists(target_segment_id, next_op_version.get())
}
pub(crate) fn append_op(&mut self, op_data: &[u8]) -> Result<WalAppendInfo, WalError> {
let version = self.allocate_next_op_version();
let target_segment_id = self.segment_id_for_op_version(version.get());
let must_rollover =
self.active_writer.as_ref().is_none_or(|w| w.segment_id() != target_segment_id);
if must_rollover {
if let Some(old_writer) = self.active_writer.take() {
old_writer.seal()?;
}
self.active_writer = Some(self.storage.open_writer(target_segment_id)?);
}
let writer = self.active_writer.as_mut().unwrap();
let op_hash = calculate_blob_hash(op_data);
writer.write_entry(version, op_hash, op_data)?;
Ok(WalAppendInfo { version, op_hash })
}
pub(crate) fn get_segment_id_for_previous_op(&self) -> u64 {
self.last_written_op_version()
.map_or(INITIAL_SEGMENT_ID, |v| self.segment_id_for_op_version(v.get()))
}
pub(crate) fn segment_id_for_op_version(&self, op_version: u64) -> u64 {
assert_ne!(op_version, 0);
(op_version.saturating_sub(1)) / self.num_ops_per_wal.get()
}
}
impl Drop for WalManager {
fn drop(&mut self) {
if let Some(writer) = self.active_writer.take()
&& let Err(e) = writer.close()
{
tracing::error!("Error closing WAL segment during drop: {:?}", e);
}
}
}