#![cfg_attr(
all(feature = "std", feature = "digest"),
doc = concat!("```rust\n", include_str!("../examples/mmriver_basic.rs"), "\n```")
)]
use crate::Error;
use crate::borrow::Cow;
use crate::error::UserError;
use crate::helper::{
PeaksMMRIVERIter, inclusion_proof_path, index_height_mmriver,
};
use crate::merge::Merge;
use crate::mmr_store::{MMRBatch, MMRStoreReadOps, MMRStoreWriteOps};
use crate::vec;
use crate::vec::Vec;
use core::marker::PhantomData;
pub struct MMRIVER<M: Merge, S> {
mmr_size: u64,
batch: MMRBatch<M::Item, S>,
_merge: PhantomData<M>,
}
impl<M: Merge, S> MMRIVER<M, S> {
pub fn new(mmr_size: u64, store: S) -> Self {
MMRIVER {
mmr_size,
batch: MMRBatch::new(store),
_merge: PhantomData,
}
}
pub fn mmr_size(&self) -> u64 {
self.mmr_size
}
pub fn is_empty(&self) -> bool {
self.mmr_size == 0
}
pub fn batch(&self) -> &MMRBatch<M::Item, S> {
&self.batch
}
pub fn store(&self) -> &S {
self.batch.store()
}
}
impl<M: Merge, S: MMRStoreReadOps<M::Item>> MMRIVER<M, S>
where
M::Item: Clone + PartialEq + Send + Sync,
M::Error: Into<UserError>,
S::Error: Into<UserError>,
{
async fn find_elem<'b>(
&self,
pos: u64,
hashes: &'b [M::Item],
) -> Result<Cow<'b, M::Item>, Error> {
let pos_offset = pos.checked_sub(self.mmr_size);
if let Some(elem) = pos_offset.and_then(|i| hashes.get(i as usize)) {
return Ok(Cow::Borrowed(elem));
}
self
.batch
.get_elem(pos)
.await
.map_err(|e| Error::StoreError(e.into()))?
.ok_or(Error::InconsistentStore)
.map(Cow::Owned)
}
pub async fn push(&mut self, data: &[u8]) -> Result<u64, Error> {
let elem = M::leaf_hash(data).map_err(|e| Error::MergeError(e.into()))?;
let elem_pos = self.mmr_size;
let mut elems = vec![elem];
let mut i = self.mmr_size + 1;
let mut g = 0u8;
while index_height_mmriver(i) > g {
let left_pos = i - (2 << g);
let right_pos = i - 1;
let (left_elem, right_elem) = futures_util::future::join(
self.find_elem(left_pos, &elems),
self.find_elem(right_pos, &elems),
)
.await;
let left_elem = left_elem?;
let right_elem = right_elem?;
let parent_elem = M::merge_pos(i + 1, &left_elem, &right_elem)
.map_err(|e| Error::MergeError(e.into()))?;
elems.push(parent_elem);
i += 1;
g += 1;
}
self.batch.append(elem_pos, elems);
self.mmr_size = i;
Ok(elem_pos)
}
pub async fn get_accumulator(&self) -> Result<Vec<M::Item>, Error> {
if self.mmr_size == 0 {
return Err(Error::GetRootOnEmpty);
}
let elems = self
.batch
.get_elems(PeaksMMRIVERIter::new(self.mmr_size - 1).collect())
.await
.map_err(|e| Error::StoreError(e.into()))?;
let peaks: Vec<M::Item> = elems
.into_iter()
.map(|elem| elem.ok_or(Error::InconsistentStore))
.collect::<Result<Vec<_>, _>>()?;
Ok(peaks)
}
pub async fn get_root(&self) -> Result<M::Item, Error> {
let peaks = self.get_accumulator().await?;
self
.bag_rhs_peaks(peaks)
.map_err(|e| Error::MergeError(e.into()))?
.ok_or(Error::InconsistentStore)
}
fn bag_rhs_peaks(
&self,
mut rhs_peaks: Vec<M::Item>,
) -> Result<Option<M::Item>, M::Error> {
while rhs_peaks.len() > 1 {
let right_peak = rhs_peaks.pop().expect("pop");
let left_peak = rhs_peaks.pop().expect("pop");
rhs_peaks.push(M::merge_peaks(&right_peak, &left_peak)?);
}
Ok(rhs_peaks.pop())
}
pub async fn gen_consistency_proof(
&self,
mmr_size_from: u64,
) -> Result<ConsistencyProof<M>, Error> {
if mmr_size_from == 0 || mmr_size_from > self.mmr_size {
return Err(Error::GenProofForInvalidLeaves);
}
let ifrom = mmr_size_from - 1;
let ito = self.mmr_size - 1;
let proof_indices = PeaksMMRIVERIter::new(ifrom)
.map(|ipeak| inclusion_proof_path(ipeak, ito));
let all_elems = self
.batch
.get_elems(proof_indices.clone().flatten().collect())
.await
.map_err(|e| Error::StoreError(e.into()))?;
let mut proof_paths: Vec<Vec<M::Item>> =
Vec::with_capacity(proof_indices.len());
let mut offset = 0;
for path_indices in proof_indices {
let path_values: Vec<M::Item> = all_elems
[offset..offset + path_indices.len()]
.iter()
.cloned()
.map(|elem| elem.ok_or(Error::InconsistentStore))
.collect::<Result<Vec<_>, _>>()?;
proof_paths.push(path_values);
offset += path_indices.len();
}
Ok(ConsistencyProof::new(
mmr_size_from,
self.mmr_size,
proof_paths,
))
}
pub async fn gen_inclusion_proof(
&self,
i: u64,
) -> Result<InclusionProof<M>, Error> {
if i >= self.mmr_size {
return Err(Error::GenProofForInvalidLeaves);
}
let c = self.mmr_size - 1;
let path_indices = crate::helper::inclusion_proof_path(i, c);
let elems = self
.batch
.get_elems(path_indices)
.await
.map_err(|e| Error::StoreError(e.into()))?;
let path_values: Vec<M::Item> = elems
.into_iter()
.map(|elem| elem.ok_or(Error::InconsistentStore))
.collect::<Result<Vec<_>, _>>()?;
Ok(InclusionProof::new(i, path_values))
}
}
impl<M: Merge, S: MMRStoreWriteOps<M::Item>> MMRIVER<M, S>
where
M::Item: Send,
S::Error: Into<UserError>,
{
pub async fn commit(&mut self) -> Result<(), Error> {
self
.batch
.commit()
.await
.map_err(|e| Error::StoreError(e.into()))
}
}
#[derive(Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(
feature = "serde",
serde(bound(serialize = "M::Item: serde::Serialize"))
)]
#[cfg_attr(
feature = "serde",
serde(bound(deserialize = "M::Item: serde::Deserialize<'de>"))
)]
pub struct InclusionProof<M: Merge> {
index: u64,
proof: Vec<M::Item>,
#[cfg_attr(feature = "serde", serde(skip))]
_merge: PhantomData<M>,
}
impl<M: Merge> InclusionProof<M>
where
M::Item: Clone + PartialEq,
M::Error: Into<UserError>,
{
#[must_use]
pub fn new(index: u64, proof: Vec<M::Item>) -> Self {
InclusionProof {
index,
proof,
_merge: PhantomData,
}
}
#[must_use]
pub fn index(&self) -> u64 {
self.index
}
#[must_use]
pub fn proof(&self) -> &[M::Item] {
&self.proof
}
pub fn included_root(&self, nodehash: M::Item) -> Result<M::Item, M::Error> {
included_root::<M>(self.index, nodehash, &self.proof)
}
pub fn verify(
&self,
nodehash: M::Item,
accumulator: &[M::Item],
) -> Result<bool, Error> {
let root = self
.included_root(nodehash)
.map_err(|e| Error::MergeError(e.into()))?;
let peak_positions = PeaksMMRIVERIter::new(self.index);
if peak_positions.len() == 0 {
return Ok(false);
}
for peak in accumulator {
if *peak == root {
return Ok(true);
}
}
Ok(false)
}
}
#[derive(Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(
feature = "serde",
serde(bound(serialize = "M::Item: serde::Serialize"))
)]
#[cfg_attr(
feature = "serde",
serde(bound(deserialize = "M::Item: serde::Deserialize<'de>"))
)]
pub struct ConsistencyProof<M: Merge> {
mmr_size_from: u64,
mmr_size_to: u64,
proof_paths: Vec<Vec<M::Item>>,
#[cfg_attr(feature = "serde", serde(skip))]
_merge: PhantomData<M>,
}
impl<M: Merge> ConsistencyProof<M>
where
M::Item: Clone + PartialEq,
M::Error: Into<UserError>,
{
#[must_use]
pub fn new(
mmr_size_from: u64,
mmr_size_to: u64,
proof_paths: Vec<Vec<M::Item>>,
) -> Self {
ConsistencyProof {
mmr_size_from,
mmr_size_to,
proof_paths,
_merge: PhantomData,
}
}
#[must_use]
pub fn mmr_size_from(&self) -> u64 {
self.mmr_size_from
}
#[must_use]
pub fn mmr_size_to(&self) -> u64 {
self.mmr_size_to
}
#[must_use]
pub fn proof_paths(&self) -> &[Vec<M::Item>] {
&self.proof_paths
}
pub fn consistent_roots(
&self,
old_accumulator: &[M::Item],
) -> Result<Vec<M::Item>, Error> {
let from_peaks: Vec<u64> =
PeaksMMRIVERIter::new(self.mmr_size_from - 1).collect();
if from_peaks.len() != old_accumulator.len()
|| from_peaks.len() != self.proof_paths.len()
{
return Err(Error::CorruptedProof);
}
let mut roots: Vec<M::Item> = Vec::new();
for i in 0..from_peaks.len() {
let root = included_root::<M>(
from_peaks[i],
old_accumulator[i].clone(),
&self.proof_paths[i],
)
.map_err(|e| Error::MergeError(e.into()))?;
if roots.last().is_some_and(|r| *r == root) {
continue;
}
roots.push(root);
}
Ok(roots)
}
pub fn verify(
&self,
old_accumulator: &[M::Item],
new_accumulator: &[M::Item],
) -> Result<bool, Error> {
let proven = self.consistent_roots(old_accumulator)?;
let mut idx = 0;
for root in proven {
if idx >= new_accumulator.len() {
return Ok(false);
}
if new_accumulator[idx] == root {
continue;
}
idx += 1;
if idx >= new_accumulator.len() || new_accumulator[idx] != root {
return Ok(false);
}
}
Ok(true)
}
}
pub fn included_root<M: Merge>(
i: u64,
nodehash: M::Item,
proof: &[M::Item],
) -> Result<M::Item, M::Error>
where
M::Item: Clone,
{
let mut root = nodehash;
let mut current_i = i;
for (g, sibling) in (index_height_mmriver(i)..).zip(proof.iter()) {
if index_height_mmriver(current_i + 1) > g {
current_i += 1;
root = M::merge_pos(current_i + 1, sibling, &root)?;
} else {
current_i += 2 << g;
root = M::merge_pos(current_i + 1, &root, sibling)?;
}
}
Ok(root)
}