use std::sync::{Arc, atomic::Ordering};
use wdev::{Device, Error as DeviceError};
use super::{
disk_window::DiskWindow,
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_win: DiskWindow,
overwritten_skips: 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_win: DiskWindow::new(),
overwritten_skips: 0,
}
}
#[inline]
pub fn current_address(&self) -> u64 {
self.cur_address
}
#[inline]
pub fn is_ended(&self) -> bool {
self.cur_address >= self.end_address
}
#[inline]
pub fn overwritten_skips(&self) -> u64 {
self.overwritten_skips
}
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) {
self.overwritten_skips += 1;
log::warn!(
"扫描在地址 {record_addr:#x} 遇到被环形覆写的未提交记录(tail={tail_now:#x}, \
容量={cap}),迭代提前终止,跳过 {skipped} 条",
skipped = self.overwritten_skips,
);
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);
}
if !self.disk_win.covers(record_addr, RECORD_HEADER_LEN) {
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_win.replace(record_addr, buf);
}
let rel_off = (record_addr - self.disk_win.offset()) as usize;
let Some(header) = RecordHeader::decode_opt(&self.disk_win.slice()[rel_off..]) else {
return Ok(None);
};
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 <= self.disk_win.slice().len() {
let p = unsafe {
self
.disk_win
.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_win.replace(record_addr, 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)
}
}