use std::{
sync::{
Arc,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
thread::{available_parallelism, sleep, spawn, yield_now},
time::Duration,
};
use aok::{OK, Void};
use compio::runtime::Runtime;
use log::info;
use tempfile::tempdir;
use waof::{Error, RECORD_HEADER_LEN, RecordHeader, WalConfig, WalLog};
use wdev::SegmentedDevice;
use super::support::{
ShortWriteDevice, WalFixture, make_pattern_payload, make_payload, reopen_single_file,
};
#[test]
fn test_enqueue_and_wait_for_commit() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let fixture = WalFixture::single_file("enqueue_commit.log", 64 * 1024)?;
let wal = fixture.wal;
let mut expected_tail = 0u64;
for i in 0..100 {
let payload = format!("commit_test_record_{i}").into_bytes();
let record_len = (RECORD_HEADER_LEN + payload.len()) as u64;
let addr = wal.enqueue(&payload)?;
assert_eq!(addr, expected_tail);
expected_tail += record_len;
}
assert_eq!(wal.tail_address(), expected_tail);
let committed = wal.commit().await?;
assert_eq!(committed, expected_tail);
assert_eq!(wal.committed_until_address(), expected_tail);
let fast_committed = wal.wait_for_commit(expected_tail).await?;
assert_eq!(fast_committed, expected_tail);
let mut iter = wal.scan_committed();
let records = iter.collect_all().await?;
assert_eq!(records.len(), 100);
info!("EnqueueAndWaitForCommit 异步写入与提交测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_try_enqueue_buffer_full_and_commit_cycle() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let fixture = WalFixture::single_file("try_enqueue.log", 16 * 1024)?;
let wal = fixture.wal;
let payload = make_payload(1000, 0xAB);
let mut first_batch_count = 0;
let mut hit_buffer_full = false;
for _ in 0..100 {
match wal.enqueue(&payload) {
Ok(_) => first_batch_count += 1,
Err(Error::BufferFull { .. }) => {
hit_buffer_full = true;
break;
}
Err(e) => return Err(e.into()),
}
}
assert!(hit_buffer_full, "缓冲区应当在容量耗尽时返回 BufferFull");
wal.commit().await?;
for _ in 0..10 {
wal.enqueue(&payload)?;
}
wal.commit().await?;
let mut iter = wal.scan(0, wal.tail_address());
let all_records = iter.collect_all().await?;
assert_eq!(all_records.len(), first_batch_count + 10);
for rec in all_records {
assert_eq!(rec.payload, payload);
}
info!("TryEnqueue 缓冲区满与 commit 淘汰循环覆写测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_enqueue_payload_limits_and_config_guards() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let buf_size = 8 * 1024;
let fixture = WalFixture::single_file("limits_guard.log", buf_size)?;
let wal = fixture.wal;
let max_allowed_payload_len = buf_size - RECORD_HEADER_LEN;
let too_large = make_payload(max_allowed_payload_len + 1, 1);
assert!(matches!(
wal.enqueue(&too_large),
Err(Error::RecordTooLarge { .. })
));
let c0 = wal.commit().await?;
assert_eq!(c0, 0);
let addr = wal.enqueue(&[])?;
assert_eq!(addr, 0);
wal.commit().await?;
let mut iter = wal.scan_committed();
let rec = iter.next().await?.expect("应有空 payload 记录");
assert!(rec.payload.is_empty());
assert_eq!(rec.header.entry_len, 0);
let dir = tempdir()?;
let db_path = dir.path().join("zero_inflight.log");
let device = Arc::new(SegmentedDevice::single_file(&db_path)?);
let config = WalConfig {
inflight_slots: 0,
..Default::default()
};
let zero_wal = WalLog::new(device, config)?;
assert_eq!(zero_wal.inflight_slots.len(), 1);
let zero_addr = zero_wal.enqueue(b"safe fallback with min 1 slot")?;
assert_eq!(zero_addr, 0);
info!("Enqueue 边界限制与配置守卫测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_commit_record_bounded_growth_concurrent_enqueue() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let fixture = WalFixture::segmented("concurrent_growth.log", 64 * 1024, 128 * 1024)?;
let wal = fixture.wal;
let num_threads = 4;
let records_per_thread = 200;
let mut handles = Vec::new();
for t_id in 0..num_threads {
let wal_clone = Arc::clone(&wal);
let handle = spawn(move || {
for i in 0..records_per_thread {
let mut data = [0u8; 64];
data[0] = t_id as u8;
data[1..5].copy_from_slice(&(i as u32).to_le_bytes());
wal_clone.enqueue(&data).expect("并发 enqueue 应当成功");
}
});
handles.push(handle);
}
for _ in 0..5 {
let _ = wal.commit().await;
yield_now();
}
for h in handles {
h.join().expect("线程应当正常结束");
}
let final_tail = wal.commit().await?;
assert_eq!(final_tail, wal.tail_address());
let mut iter = wal.scan(0, final_tail);
let records = iter.collect_all().await?;
assert_eq!(records.len(), num_threads * records_per_thread);
let mut thread_counts = [0usize; 4];
for rec in records {
let t_id = rec.payload[0] as usize;
assert!(t_id < num_threads);
thread_counts[t_id] += 1;
}
for (t_id, &count) in thread_counts.iter().enumerate() {
assert_eq!(count, records_per_thread, "线程 {t_id} 记录数应完整");
}
info!("CommitRecordBoundedGrowth 高并发写入与并发 commit 测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_fast_commit_concurrent_waiters_barrier() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let fixture = WalFixture::single_file("fast_commit_barrier.log", 2 * 1024 * 1024)?;
let wal = fixture.wal;
let thread_count = (available_parallelism().map_or(4, |n| n.get()) * 2).clamp(4, 8);
let ops_per_thread = 200 / thread_count;
let done = Arc::new(AtomicUsize::new(0));
let watchdog_on = Arc::new(AtomicBool::new(false));
{
let wal = Arc::clone(&wal);
let done = Arc::clone(&done);
let watchdog_on = Arc::clone(&watchdog_on);
spawn(move || {
while !watchdog_on.load(Ordering::Relaxed) {
sleep(Duration::from_secs(5));
let lock_free = wal.commit_lock.try_lock().is_some();
println!(
"watchdog: committed={} flushed={} tail={} safe_tail={} lock_free={} done={}",
wal.committed_until_address.load(Ordering::Relaxed),
wal.flushed_until_address.load(Ordering::Relaxed),
wal.tail_address.load(Ordering::Relaxed),
wal.safe_tail_address(),
lock_free,
done.load(Ordering::Relaxed),
);
}
});
}
let mut handles = Vec::with_capacity(thread_count);
for t_id in 0..thread_count {
let wal_clone = Arc::clone(&wal);
let done = Arc::clone(&done);
let handle = spawn(move || {
let rt = Runtime::new().unwrap();
rt.block_on(async move {
for i in 0..ops_per_thread {
let payload = format!("barrier-thread-{t_id}-item-{i}").into_bytes();
let addr = wal_clone.enqueue(&payload).unwrap();
let target_addr = addr + (RECORD_HEADER_LEN + payload.len()) as u64;
let committed = wal_clone.wait_for_commit(target_addr).await.unwrap();
assert!(committed >= target_addr);
done.fetch_add(1, Ordering::Relaxed);
}
});
});
handles.push(handle);
}
for h in handles {
h.join().unwrap();
}
watchdog_on.store(true, Ordering::Relaxed);
let final_tail = wal.commit().await?;
assert_eq!(final_tail, wal.tail_address());
let mut iter = wal.scan_committed();
let records = iter.collect_all().await?;
assert_eq!(records.len(), thread_count * ops_per_thread);
for rec in records {
rec.header.verify(&rec.payload)?;
}
info!("FastCommitConcurrentWaiters 多线程快速提交屏障测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_commit_short_write_guard() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let dir = tempdir()?;
let db_path = dir.path().join("short_write.log");
let device = Arc::new(ShortWriteDevice::single_file(&db_path)?);
let config = WalConfig::new(64 * 1024);
let wal = WalLog::new(device, config)?;
wal.enqueue(b"record-one")?;
wal.enqueue(b"record-two")?;
assert!(matches!(wal.commit().await, Err(Error::ShortWrite { .. })));
assert_eq!(wal.flushed_until_address(), 0);
assert_eq!(wal.committed_until_address(), 0);
assert!(matches!(wal.commit().await, Err(Error::ShortWrite { .. })));
assert_eq!(wal.flushed_until_address(), 0);
assert_eq!(wal.committed_until_address(), 0);
info!("Commit 短写入守卫与位点不回滚测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}
#[test]
fn test_enqueue_raw_frame_fidelity() -> Void {
let rt = Runtime::new()?;
rt.block_on(async {
let fixture = WalFixture::single_file("raw_frame.log", 64 * 1024)?;
let wal = fixture.wal;
let mut frames: Vec<Vec<u8>> = Vec::new();
for i in 0..5u32 {
let payload = make_pattern_payload(i as usize, 32 + i as usize * 10);
wal.enqueue(&payload)?;
let raw_payload = make_pattern_payload(100 + i as usize, 20 + i as usize * 7);
let mut frame = RecordHeader::for_payload(&raw_payload).to_bytes().to_vec();
frame.extend_from_slice(&raw_payload);
wal.enqueue_raw(&frame)?;
frames.push(frame);
}
let mut iter = wal.scan(0, wal.tail_address());
let records = iter.collect_all().await?;
assert_eq!(records.len(), 10);
for (i, rec) in records.iter().enumerate() {
if i % 2 == 0 {
assert_eq!(rec.payload, make_pattern_payload(i / 2, 32 + (i / 2) * 10));
} else {
let mut full = rec.header.to_bytes().to_vec();
full.extend_from_slice(&rec.payload);
assert_eq!(full, frames[(i - 1) / 2]);
}
}
wal.commit().await?;
let tail = wal.tail_address();
let wal2 = reopen_single_file(fixture.dir.path(), "raw_frame.log", 64 * 1024).await?;
assert_eq!(wal2.tail_address(), tail);
let mut iter2 = wal2.scan_committed();
assert_eq!(iter2.collect_all().await?.len(), 10);
let replica_fixture = WalFixture::single_file("raw_replica.log", 64 * 1024)?;
let replica = replica_fixture.wal;
let mut expect_addr = 0u64;
for frame in &frames {
let addr = replica.enqueue_raw(frame)?;
assert_eq!(addr, expect_addr, "从节点重放地址应与主机一致");
expect_addr += frame.len() as u64;
}
let mut iter3 = replica.scan(0, replica.tail_address());
let replayed = iter3.collect_all().await?;
for (rec, frame) in replayed.iter().zip(&frames) {
let mut full = rec.header.to_bytes().to_vec();
full.extend_from_slice(&rec.payload);
assert_eq!(&full, frame);
}
assert!(matches!(
replica.enqueue_raw(&[0u8; 4]),
Err(Error::InvalidRecordHeader)
));
let mut torn = RecordHeader::for_payload(b"abc").to_bytes().to_vec();
torn.extend_from_slice(b"abcd");
assert!(matches!(
replica.enqueue_raw(&torn),
Err(Error::InvalidRecordHeader)
));
assert!(matches!(
replica.enqueue_raw(&[0u8; 16]),
Err(Error::InvalidRecordHeader)
));
let oversized_len = 64 * 1024 + 1;
let mut oversized =
RecordHeader::new(oversized_len as u32 - RECORD_HEADER_LEN as u32, 0xDEAD_BEEF)
.to_bytes()
.to_vec();
oversized.resize(oversized_len, 0xEE);
assert!(matches!(
replica.enqueue_raw(&oversized),
Err(Error::RecordTooLarge { .. })
));
info!("EnqueueRaw 帧保真、重放地址一致与错误路径测试通过");
aok::Result::<()>::Ok(())
})?;
OK
}