mirror of
https://github.com/Chia-Network/chia-blockchain.git
synced 2026-09-28 01:56:38 -04:00
Adding type annotations to merkle set (#1082)
* Adding type annotations to merkle set * executing black * removing not used imports
This commit is contained in:
+115
-60
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user