use std::{
future::ready,
ops::Deref,
sync::{
Arc, OnceLock,
atomic::{AtomicI64, AtomicU64, Ordering, fence},
},
};
use async_lock::Mutex as AsyncLockMutex;
use crossfire::{MAsyncTx, mpsc::Array};
use wbase::{group_commit::GroupCommitPipeline, pool::AlignedBuf};
use wdev::{self, Device, Error as DeviceError};
use super::{
commit,
config::WalConfig,
disk_window::DiskWindow,
header::{RECORD_HEADER_LEN, WalFrameHeader},
iterator::WalScanIterator,
record::WalRecord,
ring_buffer::RingBuffer,
};
use crate::error::{Error, Result};
pub type ReplicationWakeTx = MAsyncTx<Array<()>>;
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_pipeline: GroupCommitPipeline,
pub replication_wake: OnceLock<ReplicationWakeTx>,
pub pending_cookie: AtomicI64,
pub last_commit_frame: AtomicU64,
pub recovered_cookie: AtomicI64,
pub recovered_committed_begin: AtomicU64,
pub recover_truncated_at: AtomicU64,
pub recover_dropped_bytes: AtomicU64,
}
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: AsyncLockMutex::new(()),
commit_pipeline: GroupCommitPipeline::new(),
replication_wake: OnceLock::new(),
pending_cookie: AtomicI64::new(commit::NO_COOKIE),
last_commit_frame: AtomicU64::new(begin_addr),
recovered_cookie: AtomicI64::new(commit::NO_COOKIE),
recovered_committed_begin: AtomicU64::new(begin_addr),
recover_truncated_at: AtomicU64::new(0),
recover_dropped_bytes: AtomicU64::new(0),
}),
})
}
pub fn set_replication_wake(&self, tx: ReplicationWakeTx) -> bool {
self.replication_wake.set(tx).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.recover_truncated_at.store(0, Ordering::Release);
self.recover_dropped_bytes.store(0, Ordering::Release);
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;
let mut last_commit: Option<(commit::CommitMeta, u64)> = None;
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) = WalFrameHeader::decode_opt(&disk_win.slice()[rel_off..]) else {
self
.note_recover_truncation(cur, "残缺帧头(不足 8 字节)")
.await;
break;
};
let entry_len = header.payload_len();
if header.is_zero() || entry_len > self.config.buffer_size {
let reason = if header.is_zero() {
"全零帧头(扇区填充或崩溃残缺尾部)"
} else {
"帧头负载长度超限(伪头)"
};
self.note_recover_truncation(cur, reason).await;
break;
}
match self
.verify_candidate(cur, header, disk_win.slice(), rel_off)
.await
{
Ok(true) => {
if entry_len == commit::COMMIT_FRAME_PAYLOAD_LEN {
let payload = self
.read_recover_payload(cur, &header, disk_win.slice(), rel_off)
.await?;
if commit::is_commit_frame(&payload)
&& let Some(meta) = commit::decode_payload(&payload)
{
last_commit = Some((meta, cur + commit::COMMIT_FRAME_TOTAL_LEN));
}
}
cur += (RECORD_HEADER_LEN + entry_len) as u64;
}
Ok(false) => {
self
.note_recover_truncation(cur, "帧负载 CRC 校验失败或负载未完整写入")
.await;
break;
}
Err(e) => return Err(e),
}
}
let committed = last_commit.as_ref().map_or(cur, |(_, end)| (*end).min(cur));
self
.committed_until_address
.store(committed, Ordering::Release);
if let Some((meta, _)) = last_commit {
self.recovered_cookie.store(meta.cookie, Ordering::Release);
self
.recovered_committed_begin
.store(meta.begin, Ordering::Release);
}
self.last_commit_frame.store(committed, Ordering::Release);
self.tail_address.store(cur, Ordering::Release);
self.flushed_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());
}
drop(guard);
Ok(cur)
}
async fn note_recover_truncation(&self, corrupt_addr: u64, reason: &str) {
let dropped = self.dropped_bytes_after(corrupt_addr);
if !self.has_nonzero_after(corrupt_addr, dropped).await {
return;
}
self
.recover_truncated_at
.store(corrupt_addr, Ordering::Release);
self.recover_dropped_bytes.store(dropped, Ordering::Release);
log::warn!(
"WAL 恢复保守截尾:{reason},损坏帧地址={corrupt_addr:#x},截断点={corrupt_addr:#x}(tail 收敛于此),丢弃其后 {dropped} 字节",
);
}
async fn has_nonzero_after(&self, addr: u64, len: u64) -> bool {
let probe_len = len.min(1024 * 1024) as usize;
if probe_len == 0 {
return false;
}
match self.device.read_range(addr, probe_len).await {
Ok(buf) => buf.as_slice().iter().any(|&b| b != 0),
Err(e) => {
log::warn!(
"WAL 恢复截断探测读失败,按常态放行:探测地址={addr:#x},探测长度={probe_len},错误={e}",
);
false
}
}
}
fn dropped_bytes_after(&self, addr: u64) -> u64 {
let dev = &*self.device;
let Some(seg_size) = dev.segment_size() else {
let size = dev.get_file_size(dev.start_segment()).unwrap_or(0);
return size.saturating_sub(addr);
};
if seg_size == 0 {
return 0;
}
let seg = addr / seg_size;
let mut dropped = dev
.get_file_size(seg as u32)
.unwrap_or(0)
.saturating_sub(addr % seg_size);
let end = dev.end_segment().unwrap_or(seg as u32);
for next in (seg + 1)..=(end as u64) {
dropped += dev.get_file_size(next as u32).unwrap_or(0);
}
dropped
}
async fn fetch_tail(
&self,
offset: u64,
requested_len: usize,
min_len: usize,
) -> Result<AlignedBuf> {
match self.device.read_range(offset, requested_len).await {
Ok(buf) => Ok(buf),
Err(DeviceError::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 = WalFrameHeader {
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: WalFrameHeader,
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()),
}
}
async fn read_recover_payload(
&self,
hdr_addr: u64,
hdr: &WalFrameHeader,
slice: &[u8],
off: usize,
) -> Result<Vec<u8>> {
let payload_end = off + RECORD_HEADER_LEN + hdr.payload_len();
if payload_end <= slice.len() {
return Ok(unsafe { slice.get_unchecked(off + RECORD_HEADER_LEN..payload_end) }.to_vec());
}
Ok(
self
.device
.read_range(hdr_addr + RECORD_HEADER_LEN as u64, hdr.payload_len())
.await?
.as_slice()
.to_vec(),
)
}
pub fn safe_initialize(&self, begin_address: u64, committed_until_address: u64) {
let end = committed_until_address.max(begin_address);
self.begin_address.store(begin_address, Ordering::Release);
self.tail_address.store(end, Ordering::Release);
self.flushed_until_address.store(end, Ordering::Release);
self.committed_until_address.store(end, Ordering::Release);
for slot in self.inflight_slots.iter() {
slot.store(u64::MAX, Ordering::Release);
}
}
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
.device
.truncate_begin_until(&self.begin_address, safe_until, ready(()))
.await?;
self
.last_commit_frame
.fetch_min(safe_until, Ordering::AcqRel);
drop(guard);
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)
}
pub fn scan_memory_records(
&self,
from: u64,
to: u64,
mut f: impl FnMut(&WalRecord) -> bool,
) -> bool {
let cap = self.ring_buffer.capacity() as u64;
let mem_base = self.tail_address();
let start = from.max(self.begin_address());
if start < mem_base.saturating_sub(cap) {
return false;
}
let scan_end = to.min(self.safe_tail_address());
let mut cur = start;
while cur + (RECORD_HEADER_LEN as u64) <= scan_end {
let header = self.ring_buffer.read_header(cur);
let entry_len = header.payload_len();
if header.is_zero() || entry_len > self.config().buffer_size {
break;
}
let next_addr = cur + (RECORD_HEADER_LEN as u64) + (entry_len as u64);
if next_addr > scan_end {
break;
}
let payload = self
.ring_buffer
.read_vec(cur + RECORD_HEADER_LEN as u64, entry_len);
if header.verify(&payload).is_err() {
break;
}
let rec = WalRecord {
address: cur,
next_address: next_addr,
header,
payload,
};
if !f(&rec) {
break;
}
cur = next_addr;
}
true
}
#[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);
self.last_commit_frame.store(begin_addr, Ordering::Release);
self
.pending_cookie
.store(commit::NO_COOKIE, Ordering::Release);
for slot in self.inflight_slots.iter() {
slot.store(u64::MAX, Ordering::Release);
}
self.device.sync_data().await?;
drop(guard);
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 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 set_pending_cookie(&self, cookie: i64) {
self.pending_cookie.store(cookie, Ordering::Release);
}
#[inline]
pub fn recovered_cookie(&self) -> i64 {
self.recovered_cookie.load(Ordering::Acquire)
}
#[inline]
pub fn recovered_committed_begin(&self) -> u64 {
self.recovered_committed_begin.load(Ordering::Acquire)
}
#[inline]
pub fn recover_truncation(&self) -> Option<(u64, u64)> {
let at = self.recover_truncated_at.load(Ordering::Acquire);
(at != 0).then(|| (at, self.recover_dropped_bytes.load(Ordering::Acquire)))
}
#[inline]
pub fn config(&self) -> &WalConfig {
&self.config
}
}