use std::{
alloc::{Layout, alloc_zeroed, handle_alloc_error},
fs::remove_file,
mem::MaybeUninit,
ops::{Deref, DerefMut},
path::Path,
slice::from_raw_parts_mut,
sync::atomic::{AtomicU64, Ordering},
};
use compio::{
buf::{BufResult, IntoInner, IoBuf, IoBufMut, SetLen},
fs::{File, metadata, rename},
io::{AsyncReadAtExt, AsyncWriteAtExt},
};
use log::debug;
use wbase::crc::Crc32Hasher;
use windex::{ENTRIES_PER_BUCKET, HashBucket, HashBucketEntry, HashIndex};
use super::{
error::{Error, Result},
meta::{IndexMeta, index_filename, index_tmp_filename},
};
const INDEX_MAGIC: &[u8; 8] = b"WEDB_IDX";
const INDEX_MAGIC_U64: u64 = u64::from_le_bytes(*INDEX_MAGIC);
const INDEX_VERSION: u32 = 1;
const HEADER_SIZE: usize = 64;
const HEADER_CRC_OFFSET: u64 = 12;
const BUCKET_BYTES: usize = 64;
const BATCH_BUCKETS: usize = 512;
const BATCH_BYTES: usize = BATCH_BUCKETS * BUCKET_BYTES;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct IndexCkptHeader {
pub version: u32,
pub crc: u32,
pub token: u128,
pub num_buckets: u64,
pub overflow_count: u64,
pub entry_count: u64,
}
impl IndexCkptHeader {
pub const SIZE: usize = HEADER_SIZE;
#[inline(always)]
pub const fn encode(&self) -> [u8; HEADER_SIZE] {
let v = self.version.to_le_bytes();
let c = self.crc.to_le_bytes();
let t = self.token.to_le_bytes();
let nb = self.num_buckets.to_le_bytes();
let oc = self.overflow_count.to_le_bytes();
let ec = self.entry_count.to_le_bytes();
[
INDEX_MAGIC[0],
INDEX_MAGIC[1],
INDEX_MAGIC[2],
INDEX_MAGIC[3],
INDEX_MAGIC[4],
INDEX_MAGIC[5],
INDEX_MAGIC[6],
INDEX_MAGIC[7],
v[0],
v[1],
v[2],
v[3],
c[0],
c[1],
c[2],
c[3],
t[0],
t[1],
t[2],
t[3],
t[4],
t[5],
t[6],
t[7],
t[8],
t[9],
t[10],
t[11],
t[12],
t[13],
t[14],
t[15],
nb[0],
nb[1],
nb[2],
nb[3],
nb[4],
nb[5],
nb[6],
nb[7],
oc[0],
oc[1],
oc[2],
oc[3],
oc[4],
oc[5],
oc[6],
oc[7],
ec[0],
ec[1],
ec[2],
ec[3],
ec[4],
ec[5],
ec[6],
ec[7],
0,
0,
0,
0,
0,
0,
0,
0,
]
}
#[inline(always)]
pub const fn decode_opt(src: &[u8]) -> Option<Self> {
if let Some(chunk) = src.first_chunk::<HEADER_SIZE>() {
let magic = u64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6], chunk[7],
]);
if magic != INDEX_MAGIC_U64 {
return None;
}
let version = u32::from_le_bytes([chunk[8], chunk[9], chunk[10], chunk[11]]);
let crc = u32::from_le_bytes([chunk[12], chunk[13], chunk[14], chunk[15]]);
let token = u128::from_le_bytes([
chunk[16], chunk[17], chunk[18], chunk[19], chunk[20], chunk[21], chunk[22], chunk[23],
chunk[24], chunk[25], chunk[26], chunk[27], chunk[28], chunk[29], chunk[30], chunk[31],
]);
let num_buckets = u64::from_le_bytes([
chunk[32], chunk[33], chunk[34], chunk[35], chunk[36], chunk[37], chunk[38], chunk[39],
]);
let overflow_count = u64::from_le_bytes([
chunk[40], chunk[41], chunk[42], chunk[43], chunk[44], chunk[45], chunk[46], chunk[47],
]);
let entry_count = u64::from_le_bytes([
chunk[48], chunk[49], chunk[50], chunk[51], chunk[52], chunk[53], chunk[54], chunk[55],
]);
Some(Self {
version,
crc,
token,
num_buckets,
overflow_count,
entry_count,
})
} else {
None
}
}
}
#[repr(C, align(64))]
struct AlignedBatch([u8; BATCH_BYTES]);
impl AlignedBatch {
#[inline]
fn new_boxed() -> Box<Self> {
let layout = Layout::new::<Self>();
unsafe {
let ptr = alloc_zeroed(layout) as *mut Self;
if ptr.is_null() {
handle_alloc_error(layout);
}
Box::from_raw(ptr)
}
}
}
impl Deref for AlignedBatch {
type Target = [u8; BATCH_BYTES];
#[inline]
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl DerefMut for AlignedBatch {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl IoBuf for AlignedBatch {
#[inline]
fn as_init(&self) -> &[u8] {
&self.0
}
}
impl SetLen for AlignedBatch {
unsafe fn set_len(&mut self, len: usize) {
debug_assert!(len <= BATCH_BYTES);
}
}
impl IoBufMut for AlignedBatch {
fn as_uninit(&mut self) -> &mut [MaybeUninit<u8>] {
let ptr = self.0.as_mut_ptr();
unsafe { from_raw_parts_mut(ptr.cast(), BATCH_BYTES) }
}
}
#[inline(always)]
fn sanitize_data_slot(raw: u64, tail: Option<u64>) -> u64 {
const INVALID_FLAGS_MASK: u64 = HashBucketEntry::TENTATIVE_MASK | HashBucketEntry::READ_CACHE_BIT;
if raw == 0 || (raw & INVALID_FLAGS_MASK) != 0 {
return 0;
}
if tail.is_some_and(|t| (raw & HashBucketEntry::ADDRESS_MASK) >= t) {
return 0;
}
raw
}
#[inline(always)]
fn resolve_read_cache(rc_skip: &dyn Fn(u64) -> u64, raw: u64) -> u64 {
if raw != 0 && raw & HashBucketEntry::READ_CACHE_BIT != 0 {
match (rc_skip)(raw & HashBucketEntry::ADDRESS_MASK) {
0 => 0,
real => (raw & !HashBucketEntry::ADDRESS_MASK) | (real & HashBucketEntry::ADDRESS_MASK),
}
} else {
raw
}
}
#[inline(always)]
fn sanitize_overflow_slot(raw: u64) -> u64 {
raw & HashBucketEntry::ADDRESS_MASK
}
struct BatchWriter<'a> {
file: &'a mut File,
hasher: &'a mut Crc32Hasher,
buf: Option<Box<AlignedBatch>>,
cursor: usize,
pos: u64,
rc_skip: &'a dyn Fn(u64) -> u64,
}
impl<'a> BatchWriter<'a> {
fn new(
file: &'a mut File,
hasher: &'a mut Crc32Hasher,
start_pos: u64,
rc_skip: &'a dyn Fn(u64) -> u64,
) -> Self {
Self {
file,
hasher,
buf: Some(AlignedBatch::new_boxed()),
cursor: 0,
pos: start_pos,
rc_skip,
}
}
#[inline]
async fn write_bucket(&mut self, bucket: &HashBucket) -> Result<()> {
let mut vals = [0u64; ENTRIES_PER_BUCKET];
for (i, slot) in bucket.entries.iter().enumerate() {
vals[i] = if i == HashBucket::OVERFLOW_INDEX {
sanitize_overflow_slot(slot.load(Ordering::Acquire))
} else {
sanitize_data_slot(self.resolve_slot(slot), None)
};
}
let offset = self.cursor;
let batch = unsafe { self.buf.as_deref_mut().unwrap_unchecked() };
let target = &mut batch[offset..offset + BUCKET_BYTES];
let (chunks, _) = target.as_chunks_mut::<8>();
for (chunk, val) in chunks.iter_mut().zip(vals.iter()) {
*chunk = val.to_le_bytes();
}
self.cursor += BUCKET_BYTES;
if self.cursor == BATCH_BYTES {
self.flush_batch().await?;
}
Ok(())
}
#[inline]
fn resolve_slot(&self, slot: &AtomicU64) -> u64 {
let raw = slot.load(Ordering::Acquire);
let resolved = resolve_read_cache(self.rc_skip, raw);
if raw != 0 && raw & HashBucketEntry::READ_CACHE_BIT != 0 && resolved == 0 {
let fresh = slot.load(Ordering::Acquire);
if fresh != raw {
return resolve_read_cache(self.rc_skip, fresh);
}
}
resolved
}
#[inline]
async fn write_zero_bucket(&mut self) -> Result<()> {
let offset = self.cursor;
let batch = unsafe { self.buf.as_deref_mut().unwrap_unchecked() };
batch[offset..offset + BUCKET_BYTES].fill(0);
self.cursor += BUCKET_BYTES;
if self.cursor == BATCH_BYTES {
self.flush_batch().await?;
}
Ok(())
}
#[inline]
async fn flush_batch(&mut self) -> Result<()> {
if self.cursor > 0 {
self
.hasher
.update(&unsafe { self.buf.as_deref().unwrap_unchecked() }[..self.cursor]);
let BufResult(res, buf) = self
.file
.write_all_at(
unsafe { self.buf.take().unwrap_unchecked() }.slice(..self.cursor),
self.pos,
)
.await;
self.buf = Some(buf.into_inner());
res?;
self.pos += self.cursor as u64;
self.cursor = 0;
}
Ok(())
}
async fn finish(mut self) -> Result<()> {
self.flush_batch().await
}
}
struct BatchReader<'a> {
file: &'a mut File,
hasher: &'a mut Crc32Hasher,
buf: Option<Box<AlignedBatch>>,
valid_bytes: usize,
cursor: usize,
remaining_bytes: u64,
pos: u64,
}
impl<'a> BatchReader<'a> {
fn new(
file: &'a mut File,
hasher: &'a mut Crc32Hasher,
start_pos: u64,
total_data_bytes: u64,
) -> Self {
Self {
file,
hasher,
buf: Some(AlignedBatch::new_boxed()),
valid_bytes: 0,
cursor: 0,
remaining_bytes: total_data_bytes,
pos: start_pos,
}
}
async fn refill(&mut self) -> Result<()> {
let to_read = (self.remaining_bytes as usize).min(BATCH_BYTES);
if to_read == 0 {
return Err(Error::InvalidIndexCkpt("意外到达文件尾部".into()));
}
let BufResult(res, buf) = self
.file
.read_exact_at(
unsafe { self.buf.take().unwrap_unchecked() }.slice(..to_read),
self.pos,
)
.await;
self.buf = Some(buf.into_inner());
res?;
self.pos += to_read as u64;
self.remaining_bytes -= to_read as u64;
self
.hasher
.update(&unsafe { self.buf.as_deref().unwrap_unchecked() }[..to_read]);
self.valid_bytes = to_read;
self.cursor = 0;
Ok(())
}
#[inline]
async fn read_bucket_into(&mut self, bucket: &HashBucket, tail: Option<u64>) -> Result<()> {
if self.cursor == self.valid_bytes {
self.refill().await?;
}
let offset = self.cursor;
let src = &unsafe { self.buf.as_deref().unwrap_unchecked() }[offset..offset + BUCKET_BYTES];
let (chunks, _) = src.as_chunks::<8>();
for (slot, slot_bytes) in bucket.entries[..HashBucket::DATA_ENTRIES]
.iter()
.zip(chunks.iter())
{
let raw = u64::from_le_bytes(*slot_bytes);
slot.store(sanitize_data_slot(raw, tail), Ordering::Relaxed);
}
let overflow_raw = u64::from_le_bytes(chunks[HashBucket::OVERFLOW_INDEX]);
bucket.entries[HashBucket::OVERFLOW_INDEX]
.store(sanitize_overflow_slot(overflow_raw), Ordering::Relaxed);
self.cursor += BUCKET_BYTES;
Ok(())
}
}
pub async fn write_index_checkpoint(
index: &HashIndex,
entry_count: usize,
checkpoint_dir: impl AsRef<Path>,
token: u128,
rc_skip: &dyn Fn(u64) -> u64,
) -> Result<IndexMeta> {
let dir = checkpoint_dir.as_ref();
let res = write_index_checkpoint_inner(index, entry_count, dir, token, rc_skip).await;
if res.is_err() {
let _ = remove_file(dir.join(index_tmp_filename(token)));
}
res
}
async fn write_index_checkpoint_inner(
index: &HashIndex,
entry_count: usize,
dir: &Path,
token: u128,
rc_skip: &dyn Fn(u64) -> u64,
) -> Result<IndexMeta> {
let final_path = dir.join(index_filename(token));
let tmp_path = dir.join(index_tmp_filename(token));
let num_buckets = index.size as u64;
let overflow_count = index.overflow_pool.allocated_count();
let mut file = File::create(&tmp_path).await?;
let ckpt_hdr = IndexCkptHeader {
version: INDEX_VERSION,
crc: 0,
token,
num_buckets,
overflow_count,
entry_count: entry_count as u64,
};
file.write_all_at(ckpt_hdr.encode(), 0).await.0?;
let mut hasher = Crc32Hasher::new();
{
let mut batch_writer = BatchWriter::new(&mut file, &mut hasher, HEADER_SIZE as u64, rc_skip);
for bucket in index.buckets.iter() {
batch_writer.write_bucket(bucket).await?;
}
for id in 1..=overflow_count {
if let Some(bucket) = index.overflow_pool.get(id) {
batch_writer.write_bucket(bucket).await?;
} else {
batch_writer.write_zero_bucket().await?;
}
}
batch_writer.finish().await?;
}
let crc = hasher.finalize();
file
.write_all_at(crc.to_le_bytes(), HEADER_CRC_OFFSET)
.await
.0?;
file.sync_all().await?;
rename(&tmp_path, &final_path).await?;
super::manager::sync_checkpoint_dir(dir).await?;
debug!(
"成功刷写 Index Checkpoint: token={token:#x}, num_buckets={num_buckets}, overflow_count={overflow_count}, entry_count={entry_count}, crc={crc:#x}"
);
Ok(IndexMeta {
size: index.size,
overflow_count,
entry_count,
})
}
pub async fn read_index_checkpoint_truncated(
index_path: impl AsRef<Path>,
expected_token: u128,
tail: Option<u64>,
) -> Result<(HashIndex, IndexMeta)> {
let path = index_path.as_ref();
if metadata(path).await.is_err() {
return Err(Error::IndexCkptNotFound(path.to_path_buf()));
}
let mut file = File::open(path).await?;
let file_len = file.metadata().await?.len();
if file_len < HEADER_SIZE as u64 {
return Err(Error::InvalidIndexCkpt("文件长度不足头部大小".into()));
}
let BufResult(res, header) = file.read_exact_at([0u8; HEADER_SIZE], 0).await;
res?;
let Some(ckpt_hdr) = IndexCkptHeader::decode_opt(&header) else {
return Err(Error::InvalidIndexCkpt(
"文件魔数不匹配或头部长度不足".into(),
));
};
if ckpt_hdr.version != INDEX_VERSION {
let mut s = String::from("不支持的版本号: ");
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(ckpt_hdr.version));
return Err(Error::InvalidIndexCkpt(s));
}
let expected_crc = ckpt_hdr.crc;
let token = ckpt_hdr.token;
if token != expected_token {
return Err(Error::TokenMismatch {
expected: expected_token,
actual: token,
});
}
let num_buckets = ckpt_hdr.num_buckets;
let overflow_count = ckpt_hdr.overflow_count;
let entry_count = ckpt_hdr.entry_count as usize;
let total_buckets = num_buckets
.checked_add(overflow_count)
.ok_or_else(|| Error::InvalidIndexCkpt("桶总数算术溢出".into()))?;
let data_len = total_buckets
.checked_mul(BUCKET_BYTES as u64)
.ok_or_else(|| Error::InvalidIndexCkpt("数据长度算术溢出".into()))?;
let expected_len = (HEADER_SIZE as u64)
.checked_add(data_len)
.ok_or_else(|| Error::InvalidIndexCkpt("文件总长度算术溢出".into()))?;
if file_len != expected_len {
let mut s = String::from("文件长度异常: 期望 ");
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(expected_len));
s.push_str(" 字节,实际 ");
s.push_str(buf.format(file_len));
s.push_str(" 字节");
return Err(Error::InvalidIndexCkpt(s));
}
let new_index = HashIndex::new(num_buckets as usize)?;
let mut hasher = Crc32Hasher::new();
{
let mut batch_reader = BatchReader::new(&mut file, &mut hasher, HEADER_SIZE as u64, data_len);
for bucket in new_index.buckets.iter() {
batch_reader.read_bucket_into(bucket, tail).await?;
}
for j in 1..=overflow_count {
let id = new_index.overflow_pool.allocate()?;
if id != j {
let mut s = String::from("溢出桶分配序号错乱: 期望 ");
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(j));
s.push_str(",实际 ");
s.push_str(buf.format(id));
return Err(Error::InvalidIndexCkpt(s));
}
let overflow_bucket = new_index.overflow_pool.get(id).ok_or_else(|| {
let mut s = String::from("无法获取分配的溢出桶 ");
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(id));
Error::InvalidIndexCkpt(s)
})?;
batch_reader.read_bucket_into(overflow_bucket, tail).await?;
}
}
let actual_crc = hasher.finalize();
if actual_crc != expected_crc {
return Err(Error::ChecksumMismatch {
expected: expected_crc,
actual: actual_crc,
});
}
debug!(
"成功恢复 Index Checkpoint: token={token:#x}, num_buckets={num_buckets}, overflow_count={overflow_count}, entry_count={entry_count}, crc={actual_crc:#x}"
);
Ok((
new_index,
IndexMeta {
size: num_buckets as usize,
overflow_count,
entry_count,
},
))
}