use std::{
hint::spin_loop,
ops::Deref,
sync::{
Arc, OnceLock,
atomic::{AtomicU64, Ordering, fence},
},
thread::yield_now,
};
use event_listener::Event;
use wdev::{Device, Error as DeviceError};
use super::{
config::WalConfig,
disk_window::DiskWindow,
error::{Error, Result},
header::{RECORD_HEADER_LEN, RecordHeader},
iterator::WalScanIterator,
ring_buffer::RingBuffer,
};
pub type ReplicationSinkFn = Arc<dyn Fn(&[u8]) + Send + Sync>;
pub(crate) const RECOVER_CHUNK_SIZE: usize = 64 * 1024;
pub struct WalLogInner<D: Device> {
pub begin_address: AtomicU64,
pub tail_address: AtomicU64,
pub flushed_until_address: AtomicU64,
pub committed_until_address: AtomicU64,
pub ring_buffer: RingBuffer,
pub device: Arc<D>,
pub config: WalConfig,
pub inflight_slots: Box<[AtomicU64]>,
pub commit_lock: async_lock::Mutex<()>,
pub commit_event: Event,
pub replication_sink: OnceLock<ReplicationSinkFn>,
}
impl<D: Device> WalLogInner<D> {
#[inline]
pub fn safe_tail_address(&self) -> u64 {
let tail = self.tail_address.load(Ordering::Acquire);
fence(Ordering::SeqCst);
self
.inflight_slots
.iter()
.fold(tail, |min, slot| min.min(slot.load(Ordering::Acquire)))
}
}
pub struct WalLog<D: Device> {
pub(crate) inner: Arc<WalLogInner<D>>,
}
impl<D: Device> Clone for WalLog<D> {
#[inline]
fn clone(&self) -> Self {
Self {
inner: Arc::clone(&self.inner),
}
}
}
impl<D: Device> Deref for WalLog<D> {
type Target = WalLogInner<D>;
#[inline]
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl<D: Device> WalLog<D> {
pub fn new(device: Arc<D>, config: WalConfig) -> Result<Self> {
let sector_size = device.sector_size();
let ring_buffer = RingBuffer::new(config.buffer_size, sector_size)?;
let slot_count = config.inflight_slots.max(1);
let slots = (0..slot_count)
.map(|_| AtomicU64::new(u64::MAX))
.collect::<Box<[_]>>();
let start_seg = device.start_segment() as u64;
let seg_size = device.segment_size().unwrap_or(0);
let begin_addr = start_seg * seg_size;
Ok(Self {
inner: Arc::new(WalLogInner {
begin_address: AtomicU64::new(begin_addr),
tail_address: AtomicU64::new(begin_addr),
flushed_until_address: AtomicU64::new(begin_addr),
committed_until_address: AtomicU64::new(begin_addr),
ring_buffer,
device,
config,
inflight_slots: slots,
commit_lock: async_lock::Mutex::new(()),
commit_event: Event::new(),
replication_sink: OnceLock::new(),
}),
})
}
pub fn set_replication_sink(&self, sink: ReplicationSinkFn) -> bool {
self.inner.replication_sink.set(sink).is_ok()
}
pub async fn open(device: Arc<D>, config: WalConfig) -> Result<Self> {
let log = Self::new(device, config)?;
log.recover().await?;
Ok(log)
}
pub async fn recover(&self) -> Result<u64> {
let _guard = self.commit_lock.lock().await;
self.device.recover()?;
let start_seg = self.device.start_segment() as u64;
let seg_size = self.device.segment_size().unwrap_or(0);
let begin_addr = start_seg * seg_size;
self.begin_address.store(begin_addr, Ordering::Release);
let mut cur = begin_addr;
if start_seg > 0
&& seg_size > 0
&& let Some(sync_addr) = self.frame_sync(cur).await?
{
cur = sync_addr;
}
let mut disk_win = DiskWindow::new();
loop {
if !disk_win.covers(cur, RECORD_HEADER_LEN) {
let buf = match self
.fetch_tail(cur, RECOVER_CHUNK_SIZE, RECORD_HEADER_LEN)
.await
{
Ok(buf) => buf,
Err(Error::Device(DeviceError::UnexpectedEof { .. })) => break,
Err(e) => return Err(e),
};
disk_win.replace(cur, buf);
}
let rel_off = (cur - disk_win.offset()) as usize;
let Some(header) = RecordHeader::decode_opt(&disk_win.slice()[rel_off..]) else {
break;
};
let entry_len = header.payload_len();
if header.is_zero() || entry_len > self.config.buffer_size {
break;
}
match self
.verify_candidate(cur, header, disk_win.slice(), rel_off)
.await
{
Ok(true) => cur += (RECORD_HEADER_LEN + entry_len) as u64,
Ok(false) => break,
Err(e) => return Err(e),
}
}
self.tail_address.store(cur, Ordering::Release);
self.flushed_until_address.store(cur, Ordering::Release);
self.committed_until_address.store(cur, Ordering::Release);
for slot in self.inflight_slots.iter() {
slot.store(u64::MAX, Ordering::Release);
}
let preload_start = cur
.saturating_sub(self.config.buffer_size as u64)
.max(self.begin_address.load(Ordering::Acquire));
if cur > preload_start {
let preload_len = (cur - preload_start) as usize;
let data = self.device.read_range(preload_start, preload_len).await?;
self.ring_buffer.write_bytes(preload_start, data.as_slice());
}
self.commit_event.notify(usize::MAX);
Ok(cur)
}
async fn fetch_tail(
&self,
offset: u64,
requested_len: usize,
min_len: usize,
) -> Result<wram::AlignedBuf> {
match self.device.read_range(offset, requested_len).await {
Ok(buf) => Ok(buf),
Err(wdev::Error::UnexpectedEof { actual, .. }) if actual >= min_len => {
Ok(self.device.read_range(offset, actual).await?)
}
Err(e) => Err(e.into()),
}
}
async fn frame_sync(&self, seg_start: u64) -> Result<Option<u64>> {
let cap = self.config.buffer_size;
let mut win_start = seg_start;
loop {
let probe = match self
.fetch_tail(win_start, RECOVER_CHUNK_SIZE, RECORD_HEADER_LEN)
.await
{
Ok(probe) => probe,
Err(Error::Device(DeviceError::UnexpectedEof { .. })) => break,
Err(e) => return Err(e),
};
let slice = probe.as_slice();
let mut off = 0;
while let Some((chunk, _)) = slice[off..].split_first_chunk::<RECORD_HEADER_LEN>() {
let packed = u64::from_le_bytes(*chunk);
let entry_len = packed as u32;
if packed == 0 || (entry_len as usize) > cap {
off += 1;
continue;
}
let hdr = RecordHeader {
entry_len,
crc32: (packed >> 32) as u32,
};
if self
.verify_candidate(win_start + off as u64, hdr, slice, off)
.await?
{
let sync_addr = win_start + off as u64;
self.begin_address.store(sync_addr, Ordering::Release);
return Ok(Some(sync_addr));
}
off += 1;
}
win_start += (slice.len() as u64)
.saturating_sub(RECORD_HEADER_LEN as u64)
.max(1);
}
Ok(None)
}
async fn verify_candidate(
&self,
hdr_addr: u64,
hdr: RecordHeader,
slice: &[u8],
off: usize,
) -> Result<bool> {
let payload_end = off + RECORD_HEADER_LEN + hdr.payload_len();
if payload_end <= slice.len() {
let payload = unsafe { slice.get_unchecked(off + RECORD_HEADER_LEN..payload_end) };
return Ok(hdr.verify(payload).is_ok());
}
match self
.device
.read_range(hdr_addr + RECORD_HEADER_LEN as u64, hdr.payload_len())
.await
{
Ok(payload) => Ok(hdr.verify(payload.as_slice()).is_ok()),
Err(DeviceError::UnexpectedEof { .. }) => Ok(false),
Err(e) => Err(e.into()),
}
}
pub fn enqueue(&self, payload: &[u8]) -> Result<u64> {
let record_len = RECORD_HEADER_LEN as u64 + payload.len() as u64;
self.check_record_len(record_len)?;
let header = RecordHeader::for_payload(payload);
self.enqueue_with(header.to_bytes(), payload)
}
pub fn enqueue_raw(&self, frame: &[u8]) -> Result<u64> {
let Some((header, payload)) = frame.split_first_chunk::<RECORD_HEADER_LEN>() else {
return Err(Error::InvalidRecordHeader);
};
let parsed = RecordHeader::from_bytes(header);
if parsed.is_zero() || parsed.entry_len as usize != payload.len() {
return Err(Error::InvalidRecordHeader);
}
self.check_record_len(frame.len() as u64)?;
self.enqueue_with(*header, payload)
}
fn check_record_len(&self, record_len: u64) -> Result<()> {
let header_limit = u32::MAX as u64 + RECORD_HEADER_LEN as u64;
if record_len > header_limit {
return Err(Error::RecordTooLarge {
len: record_len,
limit: header_limit,
});
}
if record_len > self.config.buffer_size as u64 {
return Err(Error::RecordTooLarge {
len: record_len,
limit: self.config.buffer_size as u64,
});
}
Ok(())
}
fn enqueue_with(&self, header_bytes: [u8; RECORD_HEADER_LEN], payload: &[u8]) -> Result<u64> {
let (slot_idx, current_tail) = self.acquire_inflight_slot();
let reserved_addr = self.reserve_address(
RECORD_HEADER_LEN as u64 + payload.len() as u64,
current_tail,
slot_idx,
)?;
self
.ring_buffer
.write_record(reserved_addr, &header_bytes, payload);
if let Some(sink) = self.replication_sink.get() {
let mut full_frame = Vec::with_capacity(RECORD_HEADER_LEN + payload.len());
full_frame.extend_from_slice(&header_bytes);
full_frame.extend_from_slice(payload);
sink(&full_frame);
}
unsafe { self.inflight_slots.get_unchecked(slot_idx) }.store(u64::MAX, Ordering::Release);
Ok(reserved_addr)
}
fn acquire_inflight_slot(&self) -> (usize, u64) {
let slots_len = self.inflight_slots.len();
let start = (wram::current_thread_id() as usize) % slots_len;
let mut spins = 0u32;
loop {
let current_tail = self.tail_address.load(Ordering::Acquire);
let mut idx = start;
for _ in 0..slots_len {
let slot = unsafe { self.inflight_slots.get_unchecked(idx) };
if slot
.compare_exchange_weak(u64::MAX, current_tail, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
return (idx, current_tail);
}
idx += 1;
if idx == slots_len {
idx = 0;
}
}
spins += 1;
if spins < 32 {
spin_loop();
} else {
yield_now();
}
}
}
fn reserve_address(
&self,
record_len: u64,
mut current_tail: u64,
slot_idx: usize,
) -> Result<u64> {
let sector_size = self.device.sector_size() as u64;
let buf_cap = self.config.buffer_size as u64;
loop {
let flushed = self.flushed_until_address.load(Ordering::Acquire);
let start_aligned = wram::align_down(flushed, sector_size);
let required_end = current_tail.saturating_add(record_len);
if required_end.saturating_sub(start_aligned) > buf_cap {
unsafe { self.inflight_slots.get_unchecked(slot_idx) }.store(u64::MAX, Ordering::Release);
return Err(Error::BufferFull {
available: buf_cap.saturating_sub(current_tail.saturating_sub(start_aligned)),
requested: record_len,
});
}
match self.tail_address.compare_exchange_weak(
current_tail,
required_end,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(reserved_addr) => return Ok(reserved_addr),
Err(actual) => {
current_tail = actual;
unsafe { self.inflight_slots.get_unchecked(slot_idx) }
.store(current_tail, Ordering::Release);
}
}
}
}
pub async fn commit(&self) -> Result<u64> {
let guard = self.commit_lock.lock().await;
self.commit_with_lock(guard).await
}
async fn commit_with_lock(&self, _guard: async_lock::MutexGuard<'_, ()>) -> Result<u64> {
let flushed = self.flushed_until_address.load(Ordering::Acquire);
let safe_tail = self.safe_tail_address();
if safe_tail <= flushed {
let committed = self.committed_until_address.load(Ordering::Acquire);
self.commit_event.notify(usize::MAX);
return Ok(committed);
}
let sector_size = self.device.sector_size() as u64;
let start_aligned = wram::align_down(flushed, sector_size);
let end_aligned = wram::align_up(safe_tail, sector_size);
let write_buf = self.ring_buffer.copy_range_with_padding(
start_aligned,
safe_tail,
end_aligned,
self.device.pool(),
)?;
let expected_len = write_buf.len();
let (res, _) = self.device.write_aligned(start_aligned, write_buf).await;
let written_len = res?;
if written_len != expected_len {
return Err(Error::ShortWrite {
expected: expected_len,
written: written_len,
});
}
self.device.sync_data().await?;
self
.flushed_until_address
.store(safe_tail, Ordering::Release);
self
.committed_until_address
.store(safe_tail, Ordering::Release);
self.commit_event.notify(usize::MAX);
Ok(safe_tail)
}
pub async fn wait_for_commit(&self, target_addr: u64) -> Result<u64> {
loop {
let committed = self.committed_until_address.load(Ordering::Acquire);
if committed >= target_addr {
return Ok(committed);
}
let listener = self.commit_event.listen();
let committed = self.committed_until_address.load(Ordering::Acquire);
if committed >= target_addr {
return Ok(committed);
}
if let Some(guard) = self.commit_lock.try_lock() {
self.commit_with_lock(guard).await?;
} else {
listener.await;
}
}
}
pub async fn enqueue_and_wait_for_commit(&self, payload: &[u8]) -> Result<u64> {
let addr = self.enqueue(payload)?;
let end_addr = addr + RECORD_HEADER_LEN as u64 + payload.len() as u64;
self.wait_for_commit(end_addr).await?;
Ok(addr)
}
pub async fn truncate(&self, until_address: u64) -> Result<()> {
let _guard = self.commit_lock.lock().await;
let committed = self.committed_until_address.load(Ordering::Acquire);
let safe_until = until_address.min(committed);
self.begin_address.fetch_max(safe_until, Ordering::SeqCst);
self.device.truncate_until_address(safe_until).await?;
Ok(())
}
pub fn scan(&self, from: u64, to: u64) -> WalScanIterator<D> {
let begin = self.begin_address.load(Ordering::Acquire);
let start_addr = from.max(begin);
WalScanIterator::new(Arc::clone(&self.inner), start_addr, to)
}
#[inline]
pub fn total_size(&self) -> u64 {
self.tail_address().saturating_sub(self.begin_address())
}
pub async fn reset(&self) -> Result<()> {
log::warn!(
"WAL reset:磁盘历史数据未清零,崩溃后恢复可能复活旧记录(begin={begin:#x})",
begin = self.begin_address.load(Ordering::Acquire),
);
let _guard = self.commit_lock.lock().await;
let start_seg = self.device.start_segment() as u64;
let seg_size = self.device.segment_size().unwrap_or(0);
let begin_addr = start_seg * seg_size;
self.begin_address.store(begin_addr, Ordering::Release);
self.tail_address.store(begin_addr, Ordering::Release);
self
.flushed_until_address
.store(begin_addr, Ordering::Release);
self
.committed_until_address
.store(begin_addr, Ordering::Release);
for slot in self.inflight_slots.iter() {
slot.store(u64::MAX, Ordering::Release);
}
self.device.sync_data().await?;
self.commit_event.notify(usize::MAX);
Ok(())
}
#[inline]
pub fn scan_committed(&self) -> WalScanIterator<D> {
self.scan(
self.begin_address.load(Ordering::Acquire),
self.committed_until_address.load(Ordering::Acquire),
)
}
#[inline]
pub fn scan_all(&self) -> WalScanIterator<D> {
self.scan(
self.begin_address.load(Ordering::Acquire),
self.tail_address.load(Ordering::Acquire),
)
}
#[inline]
pub fn begin_address(&self) -> u64 {
self.begin_address.load(Ordering::Acquire)
}
#[inline]
pub fn tail_address(&self) -> u64 {
self.tail_address.load(Ordering::Acquire)
}
#[inline]
pub fn flushed_until_address(&self) -> u64 {
self.flushed_until_address.load(Ordering::Acquire)
}
#[inline]
pub fn committed_until_address(&self) -> u64 {
self.committed_until_address.load(Ordering::Acquire)
}
#[inline]
pub fn device(&self) -> &Arc<D> {
&self.device
}
#[inline]
pub fn config(&self) -> &WalConfig {
&self.config
}
}