use std::collections::BTreeSet;
use crate::errors::PagedbError;
use crate::pager::Pager;
use crate::pager::format::data_page::ENVELOPE_OVERHEAD;
use crate::pager::format::page_kind::PageKind;
use crate::vfs::Vfs;
use crate::{RealmId, Result};
pub const OVERFLOW_HEADER_LEN: usize = 12;
#[must_use]
pub fn inline_value_threshold(page_size: usize) -> usize {
page_size / 4
}
const OVERFLOW_ROOT_PREFIX: usize = 4;
pub const OVERFLOW_ROOT_HEADER_LEN: usize = OVERFLOW_ROOT_PREFIX + OVERFLOW_HEADER_LEN;
#[must_use]
pub fn overflow_page_capacity(page_size: usize) -> usize {
page_size - ENVELOPE_OVERHEAD - OVERFLOW_HEADER_LEN
}
#[must_use]
pub fn overflow_root_capacity(page_size: usize) -> usize {
page_size - ENVELOPE_OVERHEAD - OVERFLOW_ROOT_HEADER_LEN
}
pub fn encode_overflow(body: &mut [u8], next: u64, data: &[u8]) -> Result<()> {
let page_size = body.len() + ENVELOPE_OVERHEAD;
let cap = overflow_page_capacity(page_size);
if data.len() > cap {
return Err(PagedbError::PayloadTooLarge);
}
for b in body.iter_mut() {
*b = 0;
}
body[0..8].copy_from_slice(&next.to_le_bytes());
let data_len = u32::try_from(data.len())
.map_err(|_| PagedbError::Io(std::io::Error::other("overflow data_len overflow")))?;
body[8..12].copy_from_slice(&data_len.to_le_bytes());
body[12..12 + data.len()].copy_from_slice(data);
Ok(())
}
pub fn decode_overflow(body: &[u8]) -> Result<(u64, &[u8])> {
if body.len() < OVERFLOW_HEADER_LEN {
return Err(PagedbError::overflow_body_malformed(
"chain_page.header_length",
));
}
let mut n = [0u8; 8];
n.copy_from_slice(&body[0..8]);
let next = u64::from_le_bytes(n);
let mut l = [0u8; 4];
l.copy_from_slice(&body[8..12]);
let data_len = u32::from_le_bytes(l) as usize;
if 12 + data_len > body.len() {
return Err(PagedbError::overflow_body_malformed(
"chain_page.data_length",
));
}
Ok((next, &body[12..12 + data_len]))
}
fn encode_overflow_root(body: &mut [u8], refcount: u32, next: u64, data: &[u8]) -> Result<()> {
let page_size = body.len() + ENVELOPE_OVERHEAD;
let cap = overflow_root_capacity(page_size);
if data.len() > cap {
return Err(PagedbError::PayloadTooLarge);
}
for b in body.iter_mut() {
*b = 0;
}
body[0..4].copy_from_slice(&refcount.to_le_bytes());
body[4..12].copy_from_slice(&next.to_le_bytes());
let data_len = u32::try_from(data.len())
.map_err(|_| PagedbError::Io(std::io::Error::other("overflow root data_len overflow")))?;
body[12..16].copy_from_slice(&data_len.to_le_bytes());
body[16..16 + data.len()].copy_from_slice(data);
Ok(())
}
fn decode_overflow_root(body: &[u8]) -> Result<(u32, u64, &[u8])> {
if body.len() < OVERFLOW_ROOT_HEADER_LEN {
return Err(PagedbError::overflow_body_malformed("root.header_length"));
}
let mut r = [0u8; 4];
r.copy_from_slice(&body[0..4]);
let refcount = u32::from_le_bytes(r);
let mut n = [0u8; 8];
n.copy_from_slice(&body[4..12]);
let next = u64::from_le_bytes(n);
let mut l = [0u8; 4];
l.copy_from_slice(&body[12..16]);
let data_len = u32::from_le_bytes(l) as usize;
if 16 + data_len > body.len() {
return Err(PagedbError::overflow_body_malformed("root.data_length"));
}
Ok((refcount, next, &body[16..16 + data_len]))
}
pub struct RootPageInfo {
pub refcount: u32,
pub next: u64,
pub root_data: Vec<u8>,
}
pub async fn read_root_page<V: Vfs>(
pager: &Pager<V>,
realm_id: RealmId,
root_page_id: u64,
) -> Result<RootPageInfo> {
let guard = pager
.read_main_page(root_page_id, realm_id, PageKind::OverflowRoot)
.await?;
let body = guard.body();
let (refcount, next, data) = decode_overflow_root(&body)?;
Ok(RootPageInfo {
refcount,
next,
root_data: data.to_vec(),
})
}
pub async fn write_chain<V: Vfs>(
pager: &Pager<V>,
realm_id: RealmId,
value: &[u8],
page_size: usize,
allocate_page: &mut (dyn FnMut() -> u64 + Send),
) -> Result<u64> {
let root_cap = overflow_root_capacity(page_size);
let chain_cap = overflow_page_capacity(page_size);
let mut offsets: Vec<usize> = Vec::new();
let mut o = 0usize;
loop {
offsets.push(o);
let cap = if offsets.len() == 1 {
root_cap
} else {
chain_cap
};
o += cap;
if o >= value.len() {
break;
}
}
let page_ids: Vec<u64> = offsets.iter().map(|_| allocate_page()).collect();
for (i, &start) in offsets.iter().enumerate() {
let is_root = i == 0;
let cap = if is_root { root_cap } else { chain_cap };
let end = (start + cap).min(value.len());
let next = if i + 1 < page_ids.len() {
page_ids[i + 1]
} else {
0
};
let chunk = &value[start..end];
let mut body = vec![0u8; page_size - ENVELOPE_OVERHEAD];
if is_root {
encode_overflow_root(&mut body, 1, next, chunk)?;
pager
.write_main_page(page_ids[i], realm_id, PageKind::OverflowRoot, &body)
.await?;
} else {
encode_overflow(&mut body, next, chunk)?;
pager
.write_main_page(page_ids[i], realm_id, PageKind::Overflow, &body)
.await?;
}
}
Ok(page_ids[0])
}
pub enum ReleaseResult {
Decremented,
Freed { freed_pages: Vec<u64> },
}
pub async fn release<V: Vfs>(
pager: &Pager<V>,
realm_id: RealmId,
root_page_id: u64,
new_page_id: u64,
) -> Result<ReleaseResult> {
let page_size = pager.page_size();
let info = read_root_page(pager, realm_id, root_page_id).await?;
if info.refcount <= 1 {
let mut freed = vec![root_page_id];
let mut seen = BTreeSet::from([root_page_id]);
let mut cur = info.next;
while cur != 0 {
if !seen.insert(cur) {
return Err(PagedbError::overflow_chain_cycle(root_page_id, cur));
}
let guard = pager
.read_main_page(cur, realm_id, PageKind::Overflow)
.await?;
let body = guard.body();
let (n, _) = decode_overflow(&body)?;
freed.push(cur);
cur = n;
}
return Ok(ReleaseResult::Freed { freed_pages: freed });
}
let new_refcount = info.refcount - 1;
let mut body = vec![0u8; page_size - ENVELOPE_OVERHEAD];
encode_overflow_root(&mut body, new_refcount, info.next, &info.root_data)?;
pager
.write_main_page(new_page_id, realm_id, PageKind::OverflowRoot, &body)
.await?;
Ok(ReleaseResult::Decremented)
}
pub async fn read_chain<V: Vfs>(
pager: &Pager<V>,
realm_id: RealmId,
root_page_id: u64,
total_len: u64,
) -> Result<Vec<u8>> {
let total_len = usize::try_from(total_len)
.ok()
.filter(|len| isize::try_from(*len).is_ok())
.ok_or_else(|| PagedbError::overflow_body_malformed("chain.total_length"))?;
let mut out: Vec<u8> = Vec::with_capacity(total_len.min(pager.page_size()));
let info = read_root_page(pager, realm_id, root_page_id).await?;
if info.root_data.len() > total_len {
return Err(PagedbError::overflow_body_malformed(
"chain.assembled_length",
));
}
out.extend_from_slice(&info.root_data);
let mut seen = BTreeSet::from([root_page_id]);
let mut next = info.next;
while next != 0 {
if !seen.insert(next) {
return Err(PagedbError::overflow_chain_cycle(root_page_id, next));
}
let guard = pager
.read_main_page(next, realm_id, PageKind::Overflow)
.await?;
let body = guard.body();
let (n, data) = decode_overflow(&body)?;
if out
.len()
.checked_add(data.len())
.is_none_or(|assembled| assembled > total_len)
{
return Err(PagedbError::overflow_body_malformed(
"chain.assembled_length",
));
}
out.extend_from_slice(data);
next = n;
}
if out.len() != total_len {
return Err(PagedbError::overflow_body_malformed(
"chain.assembled_length",
));
}
Ok(out)
}
pub async fn collect_chain<V: Vfs>(
pager: &Pager<V>,
realm_id: RealmId,
root_page_id: u64,
) -> Result<Vec<u64>> {
let mut out = vec![root_page_id];
let mut seen = BTreeSet::from([root_page_id]);
let info = read_root_page(pager, realm_id, root_page_id).await?;
let mut next = info.next;
while next != 0 {
if !seen.insert(next) {
return Err(PagedbError::overflow_chain_cycle(root_page_id, next));
}
let guard = pager
.read_main_page(next, realm_id, PageKind::Overflow)
.await?;
let body = guard.body();
let (n, _) = decode_overflow(&body)?;
out.push(next);
next = n;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use crate::crypto::CipherId;
use crate::crypto::kdf::derive_mk;
use crate::errors::CorruptionDetail;
use crate::pager::PagerConfig;
use crate::vfs::memory::MemVfs;
use super::*;
const TEST_PAGE_SIZE: usize = 4096;
const TEST_REALM: RealmId = RealmId::new([0xA4; 16]);
async fn test_pager() -> Arc<Pager<MemVfs>> {
let mk = derive_mk(&[0xA5; 32], &[0u8; 16], 0).unwrap();
let cfg = PagerConfig {
page_size: TEST_PAGE_SIZE,
buffer_pool_pages: 16,
segment_cache_pages: 16,
cipher_id: CipherId::Aes256Gcm,
mk_epoch: 0,
main_db_file_id: [0xB4; 16],
main_db_path: "/main.db".into(),
anchor_budget: 1_000_000,
dek_lru_capacity: 16,
observer_retry_count: 0,
metrics_enabled: true,
};
Arc::new(Pager::open(MemVfs::new(), mk, cfg).await.unwrap())
}
#[test]
fn round_trip_chain_page() {
let mut body = vec![0u8; 4096 - ENVELOPE_OVERHEAD];
encode_overflow(&mut body, 7, b"hello").unwrap();
let (n, d) = decode_overflow(&body).unwrap();
assert_eq!(n, 7);
assert_eq!(d, b"hello");
}
#[test]
fn round_trip_root() {
let mut body = vec![0u8; 4096 - ENVELOPE_OVERHEAD];
encode_overflow_root(&mut body, 3, 99, b"world").unwrap();
let (rc, n, d) = decode_overflow_root(&body).unwrap();
assert_eq!(rc, 3);
assert_eq!(n, 99);
assert_eq!(d, b"world");
}
#[test]
fn capacity_4k_page() {
assert_eq!(overflow_page_capacity(4096), 4044);
assert_eq!(overflow_root_capacity(4096), 4040);
}
async fn cyclic_chain(pager: &Pager<MemVfs>, root_page_id: u64, chain_page_id: u64) {
let mut root_body = vec![0u8; TEST_PAGE_SIZE - ENVELOPE_OVERHEAD];
encode_overflow_root(&mut root_body, 1, chain_page_id, b"").unwrap();
pager
.write_main_page(root_page_id, TEST_REALM, PageKind::OverflowRoot, &root_body)
.await
.unwrap();
let mut chain_body = vec![0u8; TEST_PAGE_SIZE - ENVELOPE_OVERHEAD];
encode_overflow(&mut chain_body, chain_page_id, b"").unwrap();
pager
.write_main_page(chain_page_id, TEST_REALM, PageKind::Overflow, &chain_body)
.await
.unwrap();
}
#[tokio::test(flavor = "current_thread")]
async fn read_chain_rejects_absurd_total_len_without_allocation_panic() {
let pager = test_pager().await;
let root_page_id = 42;
let mut body = vec![0u8; TEST_PAGE_SIZE - ENVELOPE_OVERHEAD];
encode_overflow_root(&mut body, 1, 0, b"small").unwrap();
pager
.write_main_page(root_page_id, TEST_REALM, PageKind::OverflowRoot, &body)
.await
.unwrap();
let error = read_chain(&pager, TEST_REALM, root_page_id, u64::MAX)
.await
.unwrap_err();
assert!(matches!(
error,
PagedbError::Corruption(CorruptionDetail::OverflowBodyMalformed { .. })
));
}
#[tokio::test(flavor = "current_thread")]
async fn read_chain_rejects_more_data_than_declared() {
let pager = test_pager().await;
let root_page_id = 43;
let mut body = vec![0u8; TEST_PAGE_SIZE - ENVELOPE_OVERHEAD];
encode_overflow_root(&mut body, 1, 0, b"too-long").unwrap();
pager
.write_main_page(root_page_id, TEST_REALM, PageKind::OverflowRoot, &body)
.await
.unwrap();
let error = read_chain(&pager, TEST_REALM, root_page_id, 1)
.await
.unwrap_err();
assert!(matches!(
error,
PagedbError::Corruption(CorruptionDetail::OverflowBodyMalformed { .. })
));
}
#[tokio::test(flavor = "current_thread")]
async fn read_chain_rejects_cycle_without_hanging() {
let pager = test_pager().await;
cyclic_chain(&pager, 44, 45).await;
let error = tokio::time::timeout(
Duration::from_secs(1),
read_chain(&pager, TEST_REALM, 44, 0),
)
.await
.expect("cycle detection should return before the timeout")
.unwrap_err();
assert!(matches!(
error,
PagedbError::Corruption(CorruptionDetail::OverflowChainCycle { .. })
));
}
#[tokio::test(flavor = "current_thread")]
async fn release_rejects_cycle_without_hanging() {
let pager = test_pager().await;
cyclic_chain(&pager, 46, 47).await;
let result =
tokio::time::timeout(Duration::from_secs(1), release(&pager, TEST_REALM, 46, 48))
.await
.expect("cycle detection should return before the timeout");
let Err(error) = result else {
panic!("overflow release cycles must not be accepted");
};
assert!(matches!(
error,
PagedbError::Corruption(CorruptionDetail::OverflowChainCycle { .. })
));
}
#[tokio::test(flavor = "current_thread")]
async fn collect_chain_rejects_cycle_without_hanging() {
let pager = test_pager().await;
cyclic_chain(&pager, 49, 50).await;
let error = tokio::time::timeout(
Duration::from_secs(1),
collect_chain(&pager, TEST_REALM, 49),
)
.await
.expect("cycle detection should return before the timeout")
.expect_err("overflow collect cycles must not be accepted");
assert!(matches!(
error,
PagedbError::Corruption(CorruptionDetail::OverflowChainCycle { .. })
));
}
}