use alloc::collections::BTreeMap;
use miden_crypto::Word;
use miden_crypto::merkle::smt::{LeafIndex, PartialSmt, SMT_DEPTH, SmtLeaf, SmtProof};
use miden_crypto::merkle::{InnerNodeInfo, MerkleError};
use crate::account::{StorageMap, StorageMapKey, StorageMapWitness};
use crate::utils::serde::{
ByteReader,
ByteWriter,
Deserializable,
DeserializationError,
Serializable,
};
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub struct PartialStorageMap {
partial_smt: PartialSmt,
entries: BTreeMap<StorageMapKey, Word>,
}
impl PartialStorageMap {
pub fn new(root: Word) -> Self {
PartialStorageMap {
partial_smt: PartialSmt::new(root),
entries: BTreeMap::new(),
}
}
pub fn with_witnesses(
witnesses: impl IntoIterator<Item = StorageMapWitness>,
) -> Result<Self, MerkleError> {
let mut map = BTreeMap::new();
let partial_smt = PartialSmt::from_proofs(witnesses.into_iter().map(|witness| {
map.extend(witness.entries());
SmtProof::from(witness)
}))?;
Ok(PartialStorageMap { partial_smt, entries: map })
}
pub fn new_full(storage_map: StorageMap) -> Self {
let partial_smt = PartialSmt::from(storage_map.smt);
let entries = storage_map.entries;
PartialStorageMap { partial_smt, entries }
}
pub fn new_minimal(storage_map: &StorageMap) -> Self {
Self::new(storage_map.root())
}
pub fn try_from_parts(
partial_smt: PartialSmt,
keys: impl IntoIterator<Item = StorageMapKey>,
) -> Result<Self, MerkleError> {
let mut entries = BTreeMap::new();
for key in keys {
if entries.contains_key(&key) {
return Err(MerkleError::DuplicateValuesForIndex(
key.hash().to_leaf_index().position(),
));
}
let value = partial_smt.get_value(&key.hash().as_word())?;
entries.insert(key, value);
}
Ok(Self { partial_smt, entries })
}
pub fn partial_smt(&self) -> &PartialSmt {
&self.partial_smt
}
pub fn root(&self) -> Word {
self.partial_smt.root()
}
pub fn get(&self, key: &StorageMapKey) -> Option<Word> {
let hash_word = key.hash().as_word();
self.partial_smt.get_value(&hash_word).ok()
}
pub fn open(&self, key: &StorageMapKey) -> Result<StorageMapWitness, MerkleError> {
let smt_proof = self.partial_smt.open(&key.hash().as_word())?;
let value = self.entries.get(key).copied().unwrap_or_default();
Ok(StorageMapWitness::new_unchecked(smt_proof, [(*key, value)]))
}
pub fn leaves(&self) -> impl Iterator<Item = (LeafIndex<SMT_DEPTH>, &SmtLeaf)> {
self.partial_smt.leaves()
}
pub fn entries(&self) -> impl Iterator<Item = (&StorageMapKey, &Word)> {
self.entries.iter()
}
pub fn inner_nodes(&self) -> impl Iterator<Item = InnerNodeInfo> + '_ {
self.partial_smt.inner_nodes()
}
pub fn add(&mut self, witness: StorageMapWitness) -> Result<(), MerkleError> {
self.entries.extend(witness.entries().map(|(key, value)| (*key, *value)));
self.partial_smt.add_proof(SmtProof::from(witness))
}
}
impl Serializable for PartialStorageMap {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
target.write(&self.partial_smt);
target.write_usize(self.entries.len());
target.write_many(self.entries.keys());
}
}
impl Deserializable for PartialStorageMap {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let partial_smt: PartialSmt = source.read()?;
let num_entries: usize = source.read()?;
let keys = source
.read_many_iter::<StorageMapKey>(num_entries)?
.collect::<Result<alloc::vec::Vec<_>, _>>()?;
Self::try_from_parts(partial_smt, keys).map_err(|err| {
DeserializationError::InvalidValue(format!(
"failed to construct partial storage map from supplied keys: {err}"
))
})
}
}
#[cfg(test)]
mod tests {
use alloc::vec::Vec;
use assert_matches::assert_matches;
use miden_crypto::merkle::MerkleError;
use miden_crypto::merkle::smt::PartialSmt;
use super::PartialStorageMap;
use crate::Word;
use crate::account::{StorageMap, StorageMapKey};
#[test]
fn try_from_parts_preserves_unrelated_partial_smt_material() -> anyhow::Result<()> {
let tracked_key = StorageMapKey::from_index(1);
let extra_key = StorageMapKey::from_index(2);
let tracked_value = Word::from([1_u32, 0, 0, 0]);
let extra_value = Word::from([2_u32, 0, 0, 0]);
let storage_map =
StorageMap::with_entries([(tracked_key, tracked_value), (extra_key, extra_value)])?;
let partial_smt = PartialSmt::from_proofs([
storage_map.open(&tracked_key).into(),
storage_map.open(&extra_key).into(),
])?;
let partial_map = PartialStorageMap::try_from_parts(partial_smt, [tracked_key])?;
assert_eq!(partial_map.entries().collect::<Vec<_>>(), [(&tracked_key, &tracked_value)]);
assert_eq!(partial_map.get(&extra_key), Some(extra_value));
Ok(())
}
#[test]
fn try_from_parts_rejects_duplicate_keys() -> anyhow::Result<()> {
let key = StorageMapKey::from_index(1);
let storage_map = StorageMap::with_entries([(key, Word::from([1_u32, 0, 0, 0]))])?;
let partial_smt = PartialSmt::from_proofs([storage_map.open(&key).into()])?;
let result = PartialStorageMap::try_from_parts(partial_smt, [key, key]);
assert_matches!(
result,
Err(MerkleError::DuplicateValuesForIndex(position))
if position == key.hash().to_leaf_index().position()
);
Ok(())
}
#[test]
fn try_from_parts_rejects_untracked_keys() {
let key = StorageMapKey::from_index(1);
let result = PartialStorageMap::try_from_parts(PartialSmt::new(Word::empty()), [key]);
assert_matches!(
result,
Err(MerkleError::UntrackedKey(hashed_key)) if hashed_key == key.hash().as_word()
);
}
}