use std::sync::{Arc, atomic::Ordering};
use wdev::{Device, Error as DeviceError};
use super::{
error::Result,
header::{RECORD_HEADER_LEN, RecordHeader},
log::{RECOVER_CHUNK_SIZE, WalLogInner},
record::WalRecord,
};
pub struct WalScanIterator<D: Device> {
inner: Arc<WalLogInner<D>>,
cur_address: u64,
end_address: u64,
disk_buf: Option<wram::AlignedBuf>,
disk_buf_offset: u64,
}
impl<D: Device> WalScanIterator<D> {
pub(crate) fn new(inner: Arc<WalLogInner<D>>, start_address: u64, end_address: u64) -> Self {
Self {
inner,
cur_address: start_address,
end_address,
disk_buf: None,
disk_buf_offset: 0,
}
}
#[inline]
pub fn current_address(&self) -> u64 {
self.cur_address
}
#[inline]
pub fn end_address(&self) -> u64 {
self.end_address
}
#[inline]
pub fn set_end_address(&mut self, new_end: u64) {
self.end_address = new_end;
}
#[inline]
pub fn is_ended(&self) -> bool {
self.cur_address >= self.end_address
}
pub async fn next(&mut self) -> Result<Option<WalRecord>> {
if self.cur_address >= self.end_address {
return Ok(None);
}
let begin_addr = self.inner.begin_address.load(Ordering::Acquire);
if self.cur_address < begin_addr {
self.cur_address = begin_addr;
}
let (available_end, mem_base) =
if self.end_address <= self.inner.flushed_until_address.load(Ordering::Acquire) {
(
self.end_address,
self.inner.tail_address.load(Ordering::Acquire),
)
} else {
let safe_tail = self.inner.safe_tail_address();
(self.end_address.min(safe_tail), safe_tail)
};
if self.cur_address + (RECORD_HEADER_LEN as u64) > available_end {
return Ok(None);
}
let record_addr = self.cur_address;
let Some((header, payload)) = self
.read_record(record_addr, available_end, mem_base)
.await?
else {
return Ok(None);
};
self.cur_address = record_addr + (RECORD_HEADER_LEN as u64) + (header.payload_len() as u64);
Ok(Some(WalRecord {
address: record_addr,
next_address: self.cur_address,
header,
payload,
}))
}
async fn read_record(
&mut self,
record_addr: u64,
available_end: u64,
mem_base: u64,
) -> Result<Option<(RecordHeader, Vec<u8>)>> {
let cap = self.inner.ring_buffer.capacity() as u64;
if record_addr >= mem_base.saturating_sub(cap) {
let header = self.inner.ring_buffer.read_header(record_addr);
let entry_len = header.payload_len();
let next_addr = record_addr + (RECORD_HEADER_LEN as u64) + (entry_len as u64);
if entry_len > self.inner.config.buffer_size || next_addr > available_end {
if record_addr < self.inner.flushed_until_address.load(Ordering::Acquire) {
return self
.read_record_from_device(record_addr, available_end)
.await;
}
return Ok(None);
}
let payload = self
.inner
.ring_buffer
.read_vec(record_addr + RECORD_HEADER_LEN as u64, entry_len);
match header.verify(&payload) {
Ok(()) => return Ok(Some((header, payload))),
Err(e) => {
if record_addr >= self.inner.flushed_until_address.load(Ordering::Acquire) {
let tail_now = self.inner.tail_address.load(Ordering::Acquire);
if record_addr < tail_now.saturating_sub(cap) {
return Ok(None);
}
return Err(e);
}
}
}
}
self
.read_record_from_device(record_addr, available_end)
.await
}
async fn read_record_from_device(
&mut self,
record_addr: u64,
available_end: u64,
) -> Result<Option<(RecordHeader, Vec<u8>)>> {
let limit = available_end.min(self.inner.flushed_until_address.load(Ordering::Acquire));
if record_addr + (RECORD_HEADER_LEN as u64) > limit {
return Ok(None);
}
let need_fetch = match &self.disk_buf {
Some(buf) => {
let buf_end = self.disk_buf_offset + buf.len() as u64;
record_addr < self.disk_buf_offset || record_addr + (RECORD_HEADER_LEN as u64) > buf_end
}
None => true,
};
if need_fetch {
let fetch_len = ((limit - record_addr) as usize).clamp(RECORD_HEADER_LEN, RECOVER_CHUNK_SIZE);
let Some(buf) = self.fetch_window(record_addr, fetch_len).await? else {
return Ok(None);
};
self.disk_buf_offset = record_addr;
self.disk_buf = Some(buf);
}
let buf = unsafe { self.disk_buf.as_ref().unwrap_unchecked() };
let rel_off = (record_addr - self.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();
let next_addr = record_addr + (RECORD_HEADER_LEN as u64) + (entry_len as u64);
if next_addr > limit {
return Ok(None);
}
let total_rec_len = RECORD_HEADER_LEN + entry_len;
let payload = if rel_off + total_rec_len <= buf.len() {
let p = unsafe {
buf
.as_slice()
.get_unchecked(rel_off + RECORD_HEADER_LEN..rel_off + total_rec_len)
};
header.verify(p)?;
p.to_vec()
} else {
let fetch_len = ((limit - record_addr) as usize)
.clamp(total_rec_len, total_rec_len.max(RECOVER_CHUNK_SIZE));
let Some(full_buf) = self.fetch_window(record_addr, fetch_len).await? else {
return Ok(None);
};
let slice = full_buf.as_slice();
let Some(p) = slice.get(RECORD_HEADER_LEN..total_rec_len) else {
return Ok(None);
};
header.verify(p)?;
let payload_vec = p.to_vec();
self.disk_buf_offset = record_addr;
self.disk_buf = Some(full_buf);
payload_vec
};
Ok(Some((header, payload)))
}
async fn fetch_window(&self, offset: u64, len: usize) -> Result<Option<wram::AlignedBuf>> {
match self.inner.device.read_range(offset, len).await {
Ok(buf) => Ok(Some(buf)),
Err(DeviceError::SegmentNotFound(_))
if offset < self.inner.begin_address.load(Ordering::Acquire) =>
{
Ok(None)
}
Err(e) => Err(e.into()),
}
}
pub async fn collect_all(&mut self) -> Result<Vec<WalRecord>> {
let mut records = Vec::new();
while let Some(rec) = self.next().await? {
records.push(rec);
}
Ok(records)
}
}