use crate::Result;
use crate::btree::BTree;
use crate::catalog::codec::{
Catalog, CatalogRowKind, RekeyIntent, RekeySegmentProgress, RekeySegmentProgressState,
SegmentMeta,
};
use crate::crypto::DerivedKey;
use crate::errors::PagedbError;
use crate::segment::reader::SegmentReader;
use crate::segment::writer::SegmentWriter;
use crate::txn::write::SegmentSideEffect;
use crate::vfs::Vfs;
#[cfg(test)]
use super::super::core::RekeyTestFault;
use super::super::core::{Db, WriterState};
use super::main::RekeyCatalogCommit;
const REKEY_ROW_BATCH: usize = 256;
struct SegmentEntry {
key: Vec<u8>,
meta: SegmentMeta,
}
impl<V: Vfs + Clone> Db<V> {
pub(super) async fn rekey_segments_pending(
&self,
state: &WriterState,
intent: &RekeyIntent,
) -> Result<bool> {
let mut cursor: Vec<u8> = vec![CatalogRowKind::Segment as u8];
loop {
let batch = self.segment_batch_from(state, &cursor).await?;
let Some(last) = batch.last() else {
return Ok(false);
};
cursor.clear();
cursor.extend_from_slice(&last.key);
cursor.push(0);
let exhausted = batch.len() < REKEY_ROW_BATCH;
if batch
.iter()
.any(|entry| Self::segment_needs_rekey(&entry.meta, intent))
{
return Ok(true);
}
if exhausted {
return Ok(false);
}
}
}
pub(super) async fn migrate_rekey_segments(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
) -> Result<()> {
let mut cursor: Vec<u8> = vec![CatalogRowKind::Segment as u8];
loop {
let batch = self.segment_batch_from(state, &cursor).await?;
let Some(last) = batch.last() else {
break;
};
cursor.clear();
cursor.extend_from_slice(&last.key);
cursor.push(0);
let exhausted = batch.len() < REKEY_ROW_BATCH;
for entry in batch {
self.migrate_rekey_segment_entry(state, intent, target_hk, entry)
.await?;
}
if exhausted {
break;
}
}
self.drain_orphaned_rekey_progress(state, intent, target_hk)
.await
}
async fn drain_orphaned_rekey_progress(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
) -> Result<()> {
let prefix = [CatalogRowKind::RekeySegmentProgress as u8];
let mut cursor: Vec<u8> = prefix.to_vec();
loop {
if state.catalog_root_page_id == 0 {
return Ok(());
}
let tree = self.rekey_catalog_tree(state);
let batch = tree
.collect_prefix_batch_from(&prefix, &cursor, REKEY_ROW_BATCH)
.await?;
let Some((last_key, _)) = batch.last() else {
return Ok(());
};
cursor.clear();
cursor.extend_from_slice(last_key);
cursor.push(0);
let exhausted = batch.len() < REKEY_ROW_BATCH;
for (key, value) in &batch {
if key.len() != 17 {
return Err(PagedbError::rekey_state_invalid(
"rekey.segment_progress.key",
));
}
let mut source_id = [0u8; 16];
source_id.copy_from_slice(&key[1..17]);
let progress = Catalog::decode_rekey_segment_progress(value)?;
match self.segment_entry_by_id(state, source_id).await? {
Some(entry) => {
self.migrate_rekey_segment_entry(state, intent, target_hk, entry)
.await?;
}
None => {
self.finish_orphaned_rekey_progress(
state, intent, target_hk, source_id, progress,
)
.await?;
}
}
}
if exhausted {
return Ok(());
}
}
}
fn segment_needs_rekey(meta: &SegmentMeta, intent: &RekeyIntent) -> bool {
meta.cipher_id == intent.source_cipher_id && meta.mk_epoch != intent.target_mk_epoch
}
async fn migrate_rekey_segment_entry(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
source: SegmentEntry,
) -> Result<()> {
let source_id = source.meta.segment_id;
let progress = if let Some(progress) = self.rekey_progress_row(state, source_id).await? {
progress
} else {
if !Self::segment_needs_rekey(&source.meta, intent) {
return Ok(());
}
self.seal_rekey_replacement(state, intent, target_hk, &source)
.await?
};
let replacement = SegmentReader::open_rekey_replacement(
self.pager.clone(),
&source.meta,
progress.replacement_segment_id,
intent.target_mk_epoch,
intent.target_cipher_id,
self.mmap_bytes_in_use.clone(),
u64::try_from(self.options.mmap_view_scratch_bytes).unwrap_or(u64::MAX),
)
.await?;
let replacement_meta = replacement.meta().clone();
self.replace_rekey_segment(
state,
intent,
target_hk,
&source,
source_id,
&replacement_meta,
)
.await?;
#[cfg(test)]
self.interrupt_rekey_if_requested(RekeyTestFault::CatalogSwapEffects)?;
self.delete_rekey_progress(state, source_id, intent.target_mk_epoch, target_hk)
.await
}
async fn seal_rekey_replacement(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
source: &SegmentEntry,
) -> Result<RekeySegmentProgress> {
let replacement_id = crate::crypto::random::segment_id()?;
let limit = u64::try_from(self.options.mmap_view_scratch_bytes).unwrap_or(u64::MAX);
let reader = SegmentReader::open_internal(
self.pager.clone(),
source.meta.clone(),
self.mmap_bytes_in_use.clone(),
limit,
)
.await?;
let footer = reader.authenticated_footer();
let mut writer = SegmentWriter::create_rekey_internal(
self.pager.clone(),
&source.meta,
replacement_id,
footer.fields.index_start_page,
footer.fields.index_page_count,
)
.await?;
writer.set_manifest(&footer.manifest)?;
for page_id in 1..source.meta.page_count.saturating_sub(1) {
let (kind, body) = reader.read_authenticated_page(page_id).await?;
let copied_page_id = writer.append_rekey_page(kind, &body).await?;
if copied_page_id != page_id {
return Err(PagedbError::rekey_state_invalid("rekey.page_id_ordering"));
}
}
let replacement = writer.seal().await?;
#[cfg(test)]
self.interrupt_rekey_if_requested(RekeyTestFault::SegmentSeal)?;
drop(reader);
if replacement.page_count != source.meta.page_count
|| replacement.format_version != source.meta.format_version
|| replacement.segment_kind != source.meta.segment_kind
|| replacement.evictable != source.meta.evictable
{
return Err(PagedbError::rekey_state_invalid(
"rekey.replacement_metadata",
));
}
let progress = RekeySegmentProgress {
replacement_segment_id: replacement_id,
state: RekeySegmentProgressState::Sealed,
};
self.write_rekey_progress(
state,
source.meta.segment_id,
progress,
intent.target_mk_epoch,
target_hk,
)
.await?;
Ok(progress)
}
async fn finish_orphaned_rekey_progress(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
source_id: [u8; 16],
progress: RekeySegmentProgress,
) -> Result<()> {
let replacement = self
.segment_entry_by_id(state, progress.replacement_segment_id)
.await?
.ok_or(PagedbError::RekeyReplacementMissing {
replacement_segment_id: progress.replacement_segment_id,
})?;
if replacement.meta.mk_epoch != intent.target_mk_epoch
|| replacement.meta.cipher_id != intent.target_cipher_id
{
return Err(PagedbError::RekeyReplacementMissing {
replacement_segment_id: progress.replacement_segment_id,
});
}
let effects = [SegmentSideEffect::Promote {
segment_id: progress.replacement_segment_id,
}];
self.reconcile_segment_effects(&effects, state.latest_commit_id)
.await?;
let limit = u64::try_from(self.options.mmap_view_scratch_bytes).unwrap_or(u64::MAX);
SegmentReader::open_internal(
self.pager.clone(),
replacement.meta.clone(),
self.mmap_bytes_in_use.clone(),
limit,
)
.await
.map_err(|_| PagedbError::RekeyReplacementMissing {
replacement_segment_id: progress.replacement_segment_id,
})?;
self.delete_rekey_progress(state, source_id, intent.target_mk_epoch, target_hk)
.await
}
async fn replace_rekey_segment(
&self,
state: &mut WriterState,
intent: &RekeyIntent,
target_hk: &DerivedKey,
source: &SegmentEntry,
source_id: [u8; 16],
replacement: &SegmentMeta,
) -> Result<()> {
let mut tree = BTree::open(
self.pager.clone(),
self.realm_id,
state.catalog_root_page_id,
state.next_page_id,
self.page_size,
);
tree.put(&source.key, &Catalog::encode_segment_meta(replacement))
.await?;
tree.flush().await?;
let freed_pages = tree.drain_freed();
let effects = [
SegmentSideEffect::Promote {
segment_id: replacement.segment_id,
},
SegmentSideEffect::Tombstone {
segment_id: source_id,
tombstone_commit_id: None,
},
];
self.commit_rekey_catalog_root(
state,
RekeyCatalogCommit {
catalog_root_page_id: tree.root_page_id(),
next_page_id: tree.next_page_id(),
freed_pages: &freed_pages,
effects: &effects,
},
intent.target_mk_epoch,
target_hk,
)
.await
.map(|_| ())
}
async fn write_rekey_progress(
&self,
state: &mut WriterState,
source_id: [u8; 16],
progress: RekeySegmentProgress,
header_epoch: u64,
header_hk: &DerivedKey,
) -> Result<()> {
let mut tree = self.rekey_catalog_tree(state);
tree.put(
&Catalog::rekey_segment_progress_key(source_id),
&Catalog::encode_rekey_segment_progress(progress),
)
.await?;
tree.flush().await?;
let freed_pages = tree.drain_freed();
self.commit_rekey_catalog_root(
state,
RekeyCatalogCommit {
catalog_root_page_id: tree.root_page_id(),
next_page_id: tree.next_page_id(),
freed_pages: &freed_pages,
effects: &[],
},
header_epoch,
header_hk,
)
.await?;
#[cfg(test)]
self.interrupt_rekey_if_requested(RekeyTestFault::ProgressRowCommit)?;
Ok(())
}
async fn delete_rekey_progress(
&self,
state: &mut WriterState,
source_id: [u8; 16],
header_epoch: u64,
header_hk: &DerivedKey,
) -> Result<()> {
let mut tree = self.rekey_catalog_tree(state);
let _ = tree
.delete(&Catalog::rekey_segment_progress_key(source_id))
.await?;
tree.flush().await?;
let freed_pages = tree.drain_freed();
self.commit_rekey_catalog_root(
state,
RekeyCatalogCommit {
catalog_root_page_id: tree.root_page_id(),
next_page_id: tree.next_page_id(),
freed_pages: &freed_pages,
effects: &[],
},
header_epoch,
header_hk,
)
.await?;
#[cfg(test)]
self.interrupt_rekey_if_requested(RekeyTestFault::ProgressDeletion)?;
Ok(())
}
fn rekey_catalog_tree(&self, state: &WriterState) -> BTree<V> {
BTree::open(
self.pager.clone(),
self.realm_id,
state.catalog_root_page_id,
state.next_page_id,
self.page_size,
)
}
async fn rekey_progress_row(
&self,
state: &WriterState,
source_id: [u8; 16],
) -> Result<Option<RekeySegmentProgress>> {
if state.catalog_root_page_id == 0 {
return Ok(None);
}
let tree = self.rekey_catalog_tree(state);
let key = Catalog::rekey_segment_progress_key(source_id);
let Some(bytes) = tree.get(&key).await? else {
return Ok(None);
};
Ok(Some(Catalog::decode_rekey_segment_progress(&bytes)?))
}
async fn segment_batch_from(
&self,
state: &WriterState,
cursor: &[u8],
) -> Result<Vec<SegmentEntry>> {
if state.catalog_root_page_id == 0 {
return Ok(Vec::new());
}
let tree = self.rekey_catalog_tree(state);
let prefix = [CatalogRowKind::Segment as u8];
tree.collect_prefix_batch_from(&prefix, cursor, REKEY_ROW_BATCH)
.await?
.into_iter()
.map(|(key, value)| {
Ok(SegmentEntry {
key: key.to_vec(),
meta: Catalog::decode_segment_meta(&value)?,
})
})
.collect()
}
async fn segment_entry_by_id(
&self,
state: &WriterState,
segment_id: [u8; 16],
) -> Result<Option<SegmentEntry>> {
let mut cursor: Vec<u8> = vec![CatalogRowKind::Segment as u8];
loop {
let batch = self.segment_batch_from(state, &cursor).await?;
let Some(last) = batch.last() else {
return Ok(None);
};
cursor.clear();
cursor.extend_from_slice(&last.key);
cursor.push(0);
let exhausted = batch.len() < REKEY_ROW_BATCH;
if let Some(entry) = batch
.into_iter()
.find(|entry| entry.meta.segment_id == segment_id)
{
return Ok(Some(entry));
}
if exhausted {
return Ok(None);
}
}
}
}
#[cfg(test)]
mod tests {
use crate::pager::format::segment_footer::FORMAT_VERSION;
use crate::segment::types::SegmentPageKind;
use crate::vfs::memory::MemVfs;
use crate::{RealmId, SegmentKind};
use super::*;
const PAGE: usize = 4096;
const REALM: RealmId = RealmId::new([0x61; 16]);
const KEK: [u8; 32] = [0x62; 32];
#[tokio::test(flavor = "current_thread")]
async fn rekey_preserves_manifest_footer_version_and_extent_index() {
let db = Db::open_internal(MemVfs::new(), KEK, PAGE, REALM)
.await
.unwrap();
let mut plain_writer = db
.create_segment(REALM, SegmentKind::Unspecified)
.await
.unwrap();
plain_writer
.append_page(SegmentPageKind::Data, b"plain-data")
.await
.unwrap();
plain_writer.set_manifest(b"plain-manifest").unwrap();
let plain_meta = plain_writer.seal().await.unwrap();
let mut txn = db.begin_write().await.unwrap();
txn.link_segment("plain", &plain_meta).await.unwrap();
txn.commit().await.unwrap();
let mut indexed_writer = db
.create_segment(REALM, SegmentKind::Unspecified)
.await
.unwrap();
let extent = indexed_writer
.append_extent(&[b"extent-a", b"extent-b"])
.await
.unwrap();
indexed_writer.set_manifest(b"indexed-manifest").unwrap();
let indexed_meta = indexed_writer.seal().await.unwrap();
let mut txn = db.begin_write().await.unwrap();
txn.link_segment("indexed", &indexed_meta).await.unwrap();
txn.commit().await.unwrap();
db.rekey_db(KEK, 1).await.unwrap();
let plain = db.open_segment(REALM, "plain").await.unwrap();
assert_eq!(plain.meta().format_version, FORMAT_VERSION);
assert_eq!(plain.index_page_count(), 0);
assert_eq!(
plain.authenticated_footer().manifest.as_slice(),
b"plain-manifest"
);
let indexed = db.open_segment(REALM, "indexed").await.unwrap();
assert_eq!(indexed.meta().format_version, FORMAT_VERSION);
assert_eq!(
indexed.authenticated_footer().manifest.as_slice(),
b"indexed-manifest"
);
let pages = indexed.find_extent(extent.start_page_id).await.unwrap();
assert!(pages[0].starts_with(b"extent-a"));
assert!(pages[1].starts_with(b"extent-b"));
}
}