Adding type annotations to merkle set (#1082)

* Adding type annotations to merkle set

* executing black

* removing not used imports
This commit is contained in:
Jesús Espino
2021-02-28 19:24:03 -08:00
committed by GitHub
parent d1ba029695
commit db333fb14e
+115 -60
View File
@@ -1,5 +1,8 @@
from abc import ABCMeta, abstractmethod
from hashlib import sha256
from typing import Dict, List
from typing import Dict, List, Any, Tuple
from src.types.blockchain_format.sized_bytes import bytes32
"""
A simple, confidence-inspiring Merkle Set standard
@@ -51,14 +54,14 @@ def init_prehashed():
init_prehashed()
def hashdown(mystr):
def hashdown(mystr: bytes):
assert len(mystr) == 66
h = prehashed[bytes(mystr[0:1] + mystr[33:34])].copy()
h.update(mystr[1:33] + mystr[34:])
return h.digest()[:32]
def compress_root(mystr):
def compress_root(mystr: bytes):
assert len(mystr) == 33
if mystr[0:1] == MIDDLE:
return mystr[1:]
@@ -68,93 +71,136 @@ def compress_root(mystr):
return sha256(mystr).digest()[:32]
def get_bit(mybytes, pos):
def get_bit(mybytes: bytes, pos: int):
assert len(mybytes) == 32
return (mybytes[pos // 8] >> (7 - (pos % 8))) & 1
class Node(metaclass=ABCMeta):
hash: bytes
@abstractmethod
def get_hash(self) -> bytes:
pass
@abstractmethod
def is_empty(self) -> bool:
pass
@abstractmethod
def is_terminal(self) -> bool:
pass
@abstractmethod
def is_double(self):
pass
@abstractmethod
def add(self, toadd: bytes, depth: int) -> "Node":
pass
@abstractmethod
def remove(self, toremove: bytes, depth: int):
pass
@abstractmethod
def is_included(self, tocheck: bytes, depth: int, p: List[bytes]):
pass
@abstractmethod
def other_included(self, tocheck: bytes, depth: int, p: List[bytes], collapse: bool):
pass
@abstractmethod
def _audit(self, hashes: List[bytes], bits: List[int]):
pass
class MerkleSet:
def __init__(self, root=None):
self.root = root
root: Node
def __init__(self, root: Node = None):
if root is None:
self.root = _empty
else:
self.root = root
def get_root(self):
def get_root(self) -> Node:
return compress_root(self.root.get_hash())
def add_already_hashed(self, toadd):
def add_already_hashed(self, toadd: bytes):
self.root = self.root.add(toadd, 0)
def remove_already_hashed(self, toremove):
def remove_already_hashed(self, toremove: bytes):
self.root = self.root.remove(toremove, 0)
def is_included_already_hashed(self, tocheck):
def is_included_already_hashed(self, tocheck: bytes) -> Tuple[bool, bytes]:
proof: List = []
r = self.root.is_included(tocheck, 0, proof)
return r, b"".join(proof)
def _audit(self, hashes):
def _audit(self, hashes: List[bytes]):
newhashes: List = []
self.root._audit(newhashes, [])
assert newhashes == sorted(newhashes)
class EmptyNode:
class EmptyNode(Node):
def __init__(self):
self.hash = BLANK
def get_hash(self):
def get_hash(self) -> bytes:
return EMPTY + BLANK
def is_empty(self):
def is_empty(self) -> bool:
return True
def is_terminal(self):
def is_terminal(self) -> bool:
return False
def is_double(self):
raise SetError()
def add(self, toadd, depth):
def add(self, toadd: bytes, depth: int) -> Node:
return TerminalNode(toadd)
def remove(self, toremove, depth):
def remove(self, toremove: bytes, depth: int) -> Node:
return self
def is_included(self, tocheck, depth, p):
def is_included(self, tocheck: bytes, depth: int, p: List[bytes]) -> bool:
p.append(EMPTY)
return False
def other_included(self, tocheck, depth, p, collapse):
def other_included(self, tocheck: bytes, depth: int, p: List[bytes], collapse: bool):
p.append(EMPTY)
def _audit(self, hashes, bits):
def _audit(self, hashes: List[bytes], bits: List[int]):
pass
_empty = EmptyNode()
class TerminalNode:
def __init__(self, hash, bits=None):
class TerminalNode(Node):
def __init__(self, hash: bytes, bits: List[int] = None):
assert len(hash) == 32
self.hash = hash
if bits is not None:
self._audit([], bits)
def get_hash(self):
def get_hash(self) -> bytes:
return TERMINAL + self.hash
def is_empty(self):
def is_empty(self) -> bool:
return False
def is_terminal(self):
def is_terminal(self) -> bool:
return True
def is_double(self):
def is_double(self) -> bool:
raise SetError()
def add(self, toadd, depth):
def add(self, toadd: bytes, depth: int) -> Node:
if toadd == self.hash:
return self
if toadd > self.hash:
@@ -162,35 +208,35 @@ class TerminalNode:
else:
return self._make_middle([TerminalNode(toadd), self], depth)
def _make_middle(self, children, depth):
def _make_middle(self, children: Any, depth: int) -> Node:
cbits = [get_bit(child.hash, depth) for child in children]
if cbits[0] != cbits[1]:
return MiddleNode(children)
nextvals = [None, None]
nextvals: List[Node] = [_empty, _empty]
nextvals[cbits[0] ^ 1] = _empty # type: ignore
nextvals[cbits[0]] = self._make_middle(children, depth + 1)
return MiddleNode(nextvals)
def remove(self, toremove, depth):
def remove(self, toremove: bytes, depth: int) -> Node:
if toremove == self.hash:
return _empty
return self
def is_included(self, tocheck, depth, proof):
proof.append(TERMINAL + self.hash)
def is_included(self, tocheck: bytes, depth: int, p: List[bytes]) -> bool:
p.append(TERMINAL + self.hash)
return tocheck == self.hash
def other_included(self, tocheck, depth, p, collapse):
def other_included(self, tocheck: bytes, depth: int, p: List[bytes], collapse: bool):
p.append(TERMINAL + self.hash)
def _audit(self, hashes, bits):
def _audit(self, hashes: List[bytes], bits: List[int]):
hashes.append(self.hash)
for pos, v in enumerate(bits):
assert get_bit(self.hash, pos) == v
class MiddleNode:
def __init__(self, children):
class MiddleNode(Node):
def __init__(self, children: List[Node]):
self.children = children
if children[0].is_empty() and children[1].is_double():
self.hash = children[1].hash
@@ -205,23 +251,23 @@ class MiddleNode:
raise SetError
self.hash = hashdown(children[0].get_hash() + children[1].get_hash())
def get_hash(self):
def get_hash(self) -> bytes:
return MIDDLE + self.hash
def is_empty(self):
def is_empty(self) -> bool:
return False
def is_terminal(self):
def is_terminal(self) -> bool:
return False
def is_double(self):
def is_double(self) -> bool:
if self.children[0].is_empty():
return self.children[1].is_double()
if self.children[1].is_empty():
return self.children[0].is_double()
return self.children[0].is_terminal() and self.children[1].is_terminal()
def add(self, toadd, depth):
def add(self, toadd: bytes, depth: int) -> Node:
bit = get_bit(toadd, depth)
child = self.children[bit]
newchild = child.add(toadd, depth + 1)
@@ -231,7 +277,7 @@ class MiddleNode:
newvals[bit] = newchild
return MiddleNode(newvals)
def remove(self, toremove, depth):
def remove(self, toremove: bytes, depth: int) -> Node:
bit = get_bit(toremove, depth)
child = self.children[bit]
newchild = child.remove(toremove, depth + 1)
@@ -246,7 +292,7 @@ class MiddleNode:
newvals[bit] = newchild
return MiddleNode(newvals)
def is_included(self, tocheck, depth, p):
def is_included(self, tocheck: bytes, depth: int, p: List[bytes]) -> bool:
p.append(MIDDLE)
if get_bit(tocheck, depth) == 0:
r = self.children[0].is_included(tocheck, depth + 1, p)
@@ -256,61 +302,70 @@ class MiddleNode:
self.children[0].other_included(tocheck, depth + 1, p, not self.children[1].is_empty())
return self.children[1].is_included(tocheck, depth + 1, p)
def other_included(self, tocheck, depth, p, collapse):
def other_included(self, tocheck: bytes, depth: int, p: List[bytes], collapse: bool):
if collapse or not self.is_double():
p.append(TRUNCATED + self.hash)
else:
self.is_included(tocheck, depth, p)
def _audit(self, hashes, bits):
def _audit(self, hashes: List[bytes], bits: List[int]):
self.children[0]._audit(hashes, bits + [0])
self.children[1]._audit(hashes, bits + [1])
class TruncatedNode:
def __init__(self, hash):
class TruncatedNode(Node):
def __init__(self, hash: bytes):
self.hash = hash
def get_hash(self):
def get_hash(self) -> bytes:
return MIDDLE + self.hash
def is_empty(self):
def is_empty(self) -> bool:
return False
def is_terminal(self):
def is_terminal(self) -> bool:
return False
def is_double(self):
def is_double(self) -> bool:
return False
def is_included(self, tocheck, depth, p):
def add(self, toadd: bytes, depth: int) -> Node:
return self
def remove(self, toremove: bytes, depth: int) -> Node:
return self
def is_included(self, tocheck: bytes, depth: int, p: List[bytes]) -> bool:
raise SetError()
def other_included(self, tocheck, depth, p, collapse):
def other_included(self, tocheck: bytes, depth: int, p: List[bytes], collapse: bool):
p.append(TRUNCATED + self.hash)
def _audit(self, hashes: List[bytes], bits: List[int]):
pass
class SetError(Exception):
pass
def confirm_included(root, val, proof):
def confirm_included(root: Node, val: bytes, proof: bytes32) -> bool:
return confirm_not_included_already_hashed(root, sha256(val).digest(), proof)
def confirm_included_already_hashed(root, val, proof):
def confirm_included_already_hashed(root: Node, val: bytes, proof: bytes32) -> bool:
return _confirm(root, val, proof, True)
def confirm_not_included(root, val, proof):
def confirm_not_included(root: Node, val: bytes, proof: bytes32) -> bool:
return confirm_not_included_already_hashed(root, sha256(val).digest(), proof)
def confirm_not_included_already_hashed(root, val, proof):
def confirm_not_included_already_hashed(root: Node, val: bytes, proof: bytes32) -> bool:
return _confirm(root, val, proof, False)
def _confirm(root, val, proof, expected):
def _confirm(root: Node, val: bytes, proof: bytes32, expected: bool) -> bool:
try:
p = deserialize_proof(proof)
if p.get_root() != root:
@@ -321,7 +376,7 @@ def _confirm(root, val, proof, expected):
return False
def deserialize_proof(proof):
def deserialize_proof(proof: bytes32) -> MerkleSet:
try:
r, pos = _deserialize(proof, 0, [])
if pos != len(proof):
@@ -331,7 +386,7 @@ def deserialize_proof(proof):
raise SetError()
def _deserialize(proof, pos, bits):
def _deserialize(proof: bytes32, pos: int, bits: List[int]) -> Tuple[Node, int]:
t = proof[pos : pos + 1] # flake8: noqa
if t == EMPTY:
return _empty, pos + 1