merkleberg 0.2.1

Merkle mountain range library in Rust
Documentation
use super::new_blake2b;
use crate::MMR;
use crate::leaf_index_to_pos;
use crate::merge::Merge;
use crate::mmr::InclusionProof;
use crate::mmr_store::MMRStoreReadOps as _;
use crate::util::MemStore;
use bytes::{Bytes, BytesMut};
use std::convert::Infallible;
use std::fmt;

#[derive(Clone)]
struct Header {
  number: u64,
  parent_hash: Bytes,
  difficulty: u64,
  chain_root: Bytes,
}

impl Header {
  fn default() -> Self {
    Header {
      number: 0,
      parent_hash: vec![0; 32].into(),
      difficulty: 0,
      chain_root: vec![0; 32].into(),
    }
  }

  fn hash(&self) -> Bytes {
    let mut hasher = new_blake2b();
    let mut hash = [0u8; 32];
    hasher.update(&self.number.to_le_bytes());
    hasher.update(&self.parent_hash);
    hasher.update(&self.difficulty.to_le_bytes());
    hasher.update(&self.chain_root);
    hasher.finalize(&mut hash);
    hash.to_vec().into()
  }
}

#[derive(Eq, PartialEq, Clone, Default)]
struct HashWithTD {
  hash: Bytes,
  td: u64,
}

impl HashWithTD {
  fn serialize(&self) -> Bytes {
    let mut data = BytesMut::from(self.hash.as_ref());
    data.extend(&self.td.to_le_bytes());
    data.into()
  }

  fn deserialize(mut data: Bytes) -> Self {
    assert_eq!(data.len(), 40);
    let mut td_bytes = [0u8; 8];
    td_bytes.copy_from_slice(&data[32..]);
    let td = u64::from_le_bytes(td_bytes);
    data.truncate(32);
    HashWithTD { hash: data, td }
  }
}

impl fmt::Debug for HashWithTD {
  fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
    write!(
      f,
      "HashWithTD {{ hash: {}, td: {} }}",
      faster_hex::hex_string(&self.hash),
      self.td
    )
  }
}

struct MergeHashWithTD;

impl Merge for MergeHashWithTD {
  type Item = HashWithTD;
  type Error = Infallible;

  fn leaf_hash(data: &[u8]) -> Result<Self::Item, Self::Error> {
    let mut hasher = new_blake2b();
    let mut hash = [0u8; 32];
    hasher.update(&[0x00]);
    hasher.update(data);
    hasher.finalize(&mut hash);

    let td = if data.len() >= 40 {
      let mut td_bytes = [0u8; 8];
      td_bytes.copy_from_slice(&data[32..40]);
      u64::from_le_bytes(td_bytes)
    } else {
      0
    };

    Ok(HashWithTD {
      hash: hash.to_vec().into(),
      td,
    })
  }

  fn merge_pos(
    _pos: u64,
    left: &Self::Item,
    right: &Self::Item,
  ) -> Result<Self::Item, Self::Error> {
    let mut hasher = new_blake2b();
    let mut hash = [0u8; 32];
    hasher.update(&left.serialize());
    hasher.update(&right.serialize());
    hasher.finalize(&mut hash);
    let td = left.td + right.td;
    Ok(HashWithTD {
      hash: hash.to_vec().into(),
      td,
    })
  }
}

struct Prover {
  headers: Vec<(Header, u64)>,
  positions: Vec<u64>,
  store: MemStore<HashWithTD>,
}

impl Prover {
  fn new() -> Prover {
    let store = MemStore::default();
    Prover {
      headers: Vec::new(),
      positions: Vec::new(),
      store,
    }
  }

  async fn gen_blocks(&mut self, count: u64) {
    let mut mmr = MMR::<MergeHashWithTD, _>::new(
      self.positions.len() as u64,
      self.store.clone(),
    );
    let mut previous = if let Some(pos) = self.positions.last() {
      mmr.store().get_elem(*pos).await.unwrap().unwrap()
    } else {
      let genesis = Header::default();
      let previous = HashWithTD {
        hash: genesis.hash(),
        td: genesis.difficulty,
      };
      self.headers.push((genesis, previous.td));
      let pos = mmr.push(&previous.serialize()).await.unwrap();
      self.positions.push(pos);
      previous
    };
    let last_number = self.headers.last().unwrap().0.number;
    for i in (last_number + 1)..=(last_number + count) {
      let block = Header {
        number: i,
        parent_hash: previous.hash.clone(),
        difficulty: i,
        chain_root: mmr.get_root().await.unwrap().serialize(),
      };
      previous = HashWithTD {
        hash: block.hash(),
        td: block.difficulty,
      };
      let pos = mmr.push(&previous.serialize()).await.unwrap();
      self.positions.push(pos);
      self.headers.push((block, previous.td));
    }
    mmr.commit().await.unwrap();
  }

  fn get_header(&self, number: u64) -> (Header, u64) {
    self.headers[number as usize].clone()
  }

  async fn gen_proof(
    &self,
    number: u64,
    later_number: u64,
  ) -> InclusionProof<MergeHashWithTD> {
    assert!(number < later_number);
    let pos = self.positions[number as usize];
    let later_pos = self.positions[later_number as usize];
    let mut mmr = MMR::<MergeHashWithTD, MemStore<HashWithTD>>::new(
      later_pos,
      self.store.clone(),
    );
    assert_eq!(
      mmr.get_root().await.unwrap().serialize(),
      self.headers[later_number as usize].0.chain_root
    );
    mmr.gen_proof(vec![pos]).await.unwrap()
  }

  fn get_pos(&self, number: u64) -> u64 {
    self.positions[number as usize]
  }
}

#[tokio::test]
async fn test_insert_header() {
  let mut prover = Prover::new();
  prover.gen_blocks(30).await;
  let h1 = 11;
  let h2 = 19;

  let prove_elem = {
    let (header, td) = prover.get_header(h1);
    MergeHashWithTD::leaf_hash(
      &HashWithTD {
        hash: header.hash(),
        td,
      }
      .serialize(),
    )
    .unwrap()
  };
  let root = {
    let (later_header, _later_td) = prover.get_header(h2);
    HashWithTD::deserialize(later_header.chain_root)
  };
  let proof = prover.gen_proof(h1, h2).await;
  let pos = leaf_index_to_pos(h1);
  assert_eq!(pos, prover.get_pos(h1));
  assert_eq!(
    prove_elem,
    prover.store.get_elem(pos).await.unwrap().unwrap()
  );
  let result = proof.verify(&root, vec![(pos, prove_elem)]).unwrap();
  assert!(result);
}