use std::collections::HashSet;
use prikk_error::{PrikkError, Result};
use prikk_object::{BlockKind, BlockPayload, ObjectId, ObjectType, RefStatePayload};
use crate::layout::RepositoryLayout;
use crate::object_store::{ObjectReadSnapshot, ObjectReader};
use crate::refs::RefStore;
use crate::rollback_verify::verify_rollback_patch_envelope;
pub const DEFAULT_HISTORY_LIMIT: usize = 20;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RefHistory {
pub ref_name: String,
pub entries: Vec<HistoryEntry>,
}
impl RefHistory {
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct HistoryEntry {
pub ref_state_id: ObjectId,
pub block_id: ObjectId,
pub update_seq: u64,
pub previous_ref_state_id: Option<ObjectId>,
pub block_kind: BlockKind,
pub parent_count: usize,
pub patch_count: usize,
pub required_attestation_count: usize,
pub rollback_patch_count: usize,
pub is_rollback_block: bool,
}
pub fn load_ref_history(
layout: &RepositoryLayout,
ref_name: &str,
limit: usize,
) -> Result<RefHistory> {
let ref_store = RefStore::new(layout.clone());
let object_store = ObjectReadSnapshot::open(layout)?;
let mut current = ref_store.read_current_ref_state_id(ref_name)?;
let mut entries = Vec::new();
let mut seen = HashSet::new();
while let Some(ref_state_id) = current {
if entries.len() >= limit {
break;
}
if !seen.insert(ref_state_id) {
return Err(PrikkError::Integrity(format!(
"RefState chain for {ref_name} contains a cycle at {ref_state_id}"
)));
}
let ref_state = read_ref_state(&object_store, ref_state_id, ref_name)?;
let block = read_block(&object_store, ref_state.target_object_id)?;
let rollback_patch_count =
count_rollback_patches(&object_store, ref_state.target_object_id, &block.patch_ids)?;
entries.push(HistoryEntry {
ref_state_id,
block_id: ref_state.target_object_id,
update_seq: ref_state.update_seq,
previous_ref_state_id: ref_state.previous_ref_state_id,
block_kind: block.kind,
parent_count: block.parent_block_ids.len(),
patch_count: block.patch_ids.len(),
required_attestation_count: ref_state.required_attestation_ids.len(),
rollback_patch_count,
is_rollback_block: rollback_patch_count != 0,
});
current = ref_state.previous_ref_state_id;
}
Ok(RefHistory {
ref_name: ref_name.to_string(),
entries,
})
}
pub fn load_received_ref_history(
layout: &RepositoryLayout,
received_ref_name: &str,
limit: usize,
) -> Result<RefHistory> {
let Some(pointer) = crate::received::read_received_pointer(layout, received_ref_name)? else {
return Ok(RefHistory {
ref_name: received_ref_name.to_string(),
entries: Vec::new(),
});
};
let Some(origin_ref_name) = received_ref_name.strip_prefix("remotes/") else {
return Err(PrikkError::InvalidName(format!(
"{received_ref_name} is not a received ref"
)));
};
let object_store = ObjectReadSnapshot::open(layout)?;
let mut current = Some(pointer.ref_state_id);
let mut entries = Vec::new();
let mut seen = HashSet::new();
while let Some(ref_state_id) = current {
if entries.len() >= limit {
break;
}
if !seen.insert(ref_state_id) {
return Err(PrikkError::Integrity(format!(
"RefState chain for {received_ref_name} contains a cycle at {ref_state_id}"
)));
}
let ref_state = read_ref_state(&object_store, ref_state_id, origin_ref_name)?;
let block = read_block(&object_store, ref_state.target_object_id)?;
let rollback_patch_count =
count_rollback_patches(&object_store, ref_state.target_object_id, &block.patch_ids)?;
entries.push(HistoryEntry {
ref_state_id,
block_id: ref_state.target_object_id,
update_seq: ref_state.update_seq,
previous_ref_state_id: ref_state.previous_ref_state_id,
block_kind: block.kind,
parent_count: block.parent_block_ids.len(),
patch_count: block.patch_ids.len(),
required_attestation_count: ref_state.required_attestation_ids.len(),
rollback_patch_count,
is_rollback_block: rollback_patch_count != 0,
});
current = ref_state.previous_ref_state_id;
}
Ok(RefHistory {
ref_name: received_ref_name.to_string(),
entries,
})
}
fn read_ref_state(
object_store: &impl ObjectReader,
ref_state_id: ObjectId,
ref_name: &str,
) -> Result<RefStatePayload> {
let Some(envelope) = object_store.read_object(ref_state_id)? else {
return Err(PrikkError::Integrity(format!(
"history RefState {ref_state_id} is missing"
)));
};
if envelope.object_type != ObjectType::RefState {
return Err(PrikkError::Integrity(format!(
"history object {ref_state_id} is {}, expected RefState",
envelope.object_type
)));
}
let payload =
RefStatePayload::decode_canonical(&envelope.canonical_payload, envelope.schema_version)?;
if payload.ref_name != ref_name {
return Err(PrikkError::Integrity(format!(
"history RefState {ref_state_id} name mismatch: expected {ref_name}, got {}",
payload.ref_name
)));
}
Ok(payload)
}
fn count_rollback_patches(
object_store: &impl ObjectReader,
block_id: ObjectId,
patch_ids: &[ObjectId],
) -> Result<usize> {
let mut count = 0_usize;
for patch_id in patch_ids {
let Some(envelope) = object_store.read_typed(*patch_id, ObjectType::Patch)? else {
return Err(PrikkError::Integrity(format!(
"history Block {block_id} references missing Patch {patch_id}"
)));
};
let context = format!("history Block {block_id} Patch {patch_id}");
if verify_rollback_patch_envelope(&envelope, &context)? {
count = count.checked_add(1).ok_or_else(|| {
PrikkError::Integrity("history rollback patch count overflow".to_string())
})?;
}
}
Ok(count)
}
fn read_block(object_store: &impl ObjectReader, block_id: ObjectId) -> Result<BlockPayload> {
let Some(envelope) = object_store.read_object(block_id)? else {
return Err(PrikkError::Integrity(format!(
"history Block {block_id} is missing"
)));
};
if envelope.object_type != ObjectType::Block {
return Err(PrikkError::Integrity(format!(
"history object {block_id} is {}, expected Block",
envelope.object_type
)));
}
BlockPayload::decode_canonical(&envelope.canonical_payload)
}
#[cfg(test)]
mod tests;