use std::{iter::once, sync::atomic::Ordering};
use wbase::{align::sector_bounds, backoff::Backoff, thread::current_thread_id};
use wdev::Device;
use super::{
header::{RECORD_HEADER_LEN, WalFrameHeader},
log::WalLog,
};
use crate::error::{Error, Result};
impl<D: Device> WalLog<D> {
#[inline]
pub fn enqueue(&self, payload: &[u8]) -> Result<u64> {
self.enqueue_parts(&[payload])
}
pub fn enqueue_parts(&self, parts: &[&[u8]]) -> Result<u64> {
let payload_len: u64 = parts.iter().map(|part| part.len() as u64).sum();
self.check_record_len(RECORD_HEADER_LEN as u64 + payload_len)?;
let header = WalFrameHeader::for_payload_parts(parts);
self.enqueue_reserved(
RECORD_HEADER_LEN as u64 + payload_len,
once((header, parts)),
)
}
pub fn enqueue_raw(&self, frame: &[u8]) -> Result<u64> {
let Some((header_bytes, payload)) = frame.split_first_chunk::<RECORD_HEADER_LEN>() else {
return Err(Error::InvalidRecordHeader);
};
let parsed = WalFrameHeader::from_bytes(header_bytes);
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_reserved(frame.len() as u64, once((parsed, &[payload][..])))
}
pub fn enqueue_frames(&self, frames: &[&[&[u8]]]) -> Result<u64> {
let mut total_len = 0u64;
for parts in frames {
let payload_len: u64 = parts.iter().map(|part| part.len() as u64).sum();
let frame_len = RECORD_HEADER_LEN as u64 + payload_len;
self.check_record_len(frame_len)?;
total_len += frame_len;
}
self.check_record_len(total_len)?;
self.enqueue_reserved(
total_len,
frames
.iter()
.map(|parts| (WalFrameHeader::for_payload_parts(parts), *parts)),
)
}
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_reserved<'p, I>(&self, total_len: u64, frames: I) -> Result<u64>
where
I: IntoIterator<Item = (WalFrameHeader, &'p [&'p [u8]])>,
{
let (slot_idx, current_tail) = self.acquire_inflight_slot();
let reserved_addr = self.reserve_address(total_len, current_tail, slot_idx)?;
let mut addr = reserved_addr;
for (header, parts) in frames {
let frame_len = RECORD_HEADER_LEN as u64 + header.payload_len() as u64;
self
.ring_buffer
.write_record_parts(addr, &header.to_bytes(), parts);
addr += frame_len;
}
debug_assert_eq!(
addr,
reserved_addr + total_len,
"帧序列实际写入长度须与预留总长一致"
);
if let Some(tx) = self.replication_wake.get() {
let _ = tx.try_send(());
}
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 = (current_thread_id() as usize) % slots_len;
let mut backoff = Backoff::new();
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;
}
}
backoff.stage().wait_busy();
backoff.advance();
}
}
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 required_end = current_tail.saturating_add(record_len);
let start_aligned = sector_bounds(flushed, required_end, sector_size).0;
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);
}
}
}
}
}