use std::{
hint::spin_loop,
ops::Deref,
sync::{
Arc,
atomic::{AtomicU64, Ordering, fence},
},
thread::yield_now,
};
use event_listener::Event;
use wdev::{Device, Error as DeviceError};
use super::{
config::WalConfig,
error::{Error, Result},
header::{RECORD_HEADER_LEN, RecordHeader},
iterator::WalScanIterator,
ring_buffer::RingBuffer,
};
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,
}
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(),
}),
})
}
pub async fn open(device: Arc<D>, config: WalConfig) -> Result<Self> {
let log = Self::new(Arc::clone(&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_buf: Option<wram::AlignedBuf> = None;
let mut disk_buf_offset = 0u64;
loop {
let need_fetch = match &disk_buf {
None => true,
Some(buf) => {
let buf_end = disk_buf_offset + buf.len() as u64;
cur < disk_buf_offset || cur + (RECORD_HEADER_LEN as u64) > buf_end
}
};
if need_fetch {
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_buf_offset = cur;
disk_buf = Some(buf);
}
let buf = unsafe { disk_buf.as_ref().unwrap_unchecked() };
let rel_off = (cur - disk_buf_offset) as usize;
let header = unsafe {
let chunk = &*(buf.as_slice().as_ptr().add(rel_off) as *const [u8; RECORD_HEADER_LEN]);
RecordHeader::from_bytes(chunk)
};
let entry_len = header.payload_len();
if header.is_zero() || entry_len > self.config.buffer_size {
break;
}
match self
.verify_candidate(cur, header, buf.as_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 off + RECORD_HEADER_LEN <= slice.len() {
let hdr = unsafe {
let chunk = &*(slice.as_ptr().add(off) as *const [u8; RECORD_HEADER_LEN]);
RecordHeader::from_bytes(chunk)
};
if hdr.is_zero() || hdr.payload_len() > cap {
off += 1;
continue;
}
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 payload_len = payload.len();
if payload_len > u32::MAX as usize {
return Err(Error::PayloadTooLarge(payload_len));
}
let record_len = RECORD_HEADER_LEN as u64 + payload_len as u64;
if record_len > self.config.buffer_size as u64 {
return Err(Error::PayloadTooLarge(payload_len));
}
let header = RecordHeader::for_payload(payload);
let header_bytes = header.to_bytes();
let (slot_idx, current_tail) = self.acquire_inflight_slot();
let reserved_addr = self.reserve_address(record_len, current_tail, slot_idx)?;
self
.ring_buffer
.write_record(reserved_addr, &header_bytes, payload);
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<()> {
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
}
}