diff --git a/chia/data_layer/data_layer_util.py b/chia/data_layer/data_layer_util.py index c34b80e730..ca7da7feed 100644 --- a/chia/data_layer/data_layer_util.py +++ b/chia/data_layer/data_layer_util.py @@ -299,6 +299,14 @@ class Root: status=Status(row["status"]), ) + def to_row(self) -> Dict[str, Any]: + return { + "tree_id": self.tree_id, + "node_hash": self.node_hash, + "generation": self.generation, + "status": self.status.value, + } + @classmethod def unmarshal(cls, marshalled: Dict[str, Any]) -> "Root": return cls( @@ -696,3 +704,9 @@ class PluginStatus: "downloaders": self.downloaders, } } + + +@dataclasses.dataclass(frozen=True) +class InsertResult: + node_hash: bytes32 + root: Root diff --git a/chia/data_layer/data_store.py b/chia/data_layer/data_store.py index 0fc7f5b6c4..697590ca21 100644 --- a/chia/data_layer/data_store.py +++ b/chia/data_layer/data_store.py @@ -12,6 +12,7 @@ import aiosqlite from chia.data_layer.data_layer_errors import KeyNotFoundError, NodeHashError, TreeGenerationIncrementingError from chia.data_layer.data_layer_util import ( DiffData, + InsertResult, InternalNode, Node, NodeType, @@ -168,7 +169,7 @@ class DataStore: node_hash: Optional[bytes32], status: Status, generation: Optional[int] = None, - ) -> None: + ) -> Root: # This should be replaced by an SQLite schema level check. # https://github.com/Chia-Network/chia-blockchain/pull/9284 tree_id = bytes32(tree_id) @@ -184,17 +185,19 @@ class DataStore: else: generation = existing_generation + 1 + new_root = Root( + tree_id=tree_id, + node_hash=None if node_hash is None else node_hash, + generation=generation, + status=status, + ) + await writer.execute( """ INSERT INTO root(tree_id, generation, node_hash, status) VALUES(:tree_id, :generation, :node_hash, :status) """, - { - "tree_id": tree_id, - "generation": generation, - "node_hash": None if node_hash is None else node_hash, - "status": status.value, - }, + new_root.to_row(), ) # `node_hash` is now a root, so it has no ancestor. @@ -213,6 +216,8 @@ class DataStore: values, ) + return new_root + async def _insert_node( self, node_hash: bytes32, @@ -609,12 +614,19 @@ class DataStore: return ancestors async def get_ancestors_optimized( - self, node_hash: bytes32, tree_id: bytes32, generation: Optional[int] = None + self, + node_hash: bytes32, + tree_id: bytes32, + generation: Optional[int] = None, + root_hash: Optional[bytes32] = None, ) -> List[InternalNode]: async with self.db_wrapper.reader(): nodes = [] - root = await self.get_tree_root(tree_id=tree_id, generation=generation) - if root.node_hash is None: + if root_hash is None: + root = await self.get_tree_root(tree_id=tree_id, generation=generation) + root_hash = root.node_hash + + if root_hash is None: return [] while True: @@ -625,7 +637,7 @@ class DataStore: node_hash = internal_node.hash if len(nodes) > 0: - if root.node_hash != nodes[-1].hash: + if root_hash != nodes[-1].hash: raise RuntimeError("Ancestors list didn't produce the root as top result.") return nodes @@ -719,13 +731,17 @@ class DataStore: return NodeType(raw_node_type["node_type"]) - async def get_terminal_node_for_seed(self, tree_id: bytes32, seed: bytes32) -> Optional[bytes32]: + async def get_terminal_node_for_seed( + self, tree_id: bytes32, seed: bytes32, root_hash: Optional[bytes32] = None + ) -> Optional[bytes32]: path = int.from_bytes(seed, byteorder="big") async with self.db_wrapper.reader(): - root = await self.get_tree_root(tree_id) - if root is None or root.node_hash is None: + if root_hash is None: + root = await self.get_tree_root(tree_id) + root_hash = root.node_hash + if root_hash is None: return None - node_hash = root.node_hash + node_hash = root_hash while True: node = await self.get_node(node_hash) assert node is not None @@ -751,15 +767,20 @@ class DataStore: hint_keys_values: Optional[Dict[bytes, bytes]] = None, use_optimized: bool = True, status: Status = Status.PENDING, - ) -> bytes32: + root: Optional[Root] = None, + ) -> InsertResult: async with self.db_wrapper.writer(): - was_empty = await self.table_is_empty(tree_id=tree_id) + if root is None: + root = await self.get_tree_root(tree_id=tree_id) + + was_empty = root.node_hash is None + if was_empty: reference_node_hash = None side = None else: seed = leaf_hash(key=key, value=value) - reference_node_hash = await self.get_terminal_node_for_seed(tree_id, seed) + reference_node_hash = await self.get_terminal_node_for_seed(tree_id, seed, root_hash=root.node_hash) side = self.get_side_for_seed(seed) return await self.insert( @@ -771,10 +792,11 @@ class DataStore: hint_keys_values=hint_keys_values, use_optimized=use_optimized, status=status, + root=root, ) - async def get_keys_values_dict(self, tree_id: bytes32) -> Dict[bytes, bytes]: - pairs = await self.get_keys_values(tree_id=tree_id) + async def get_keys_values_dict(self, tree_id: bytes32, root_hash: Optional[bytes32] = None) -> Dict[bytes, bytes]: + pairs = await self.get_keys_values(tree_id=tree_id, root_hash=root_hash) return {node.key: node.value for node in pairs} async def get_keys(self, tree_id: bytes32, root_hash: Optional[bytes32] = None) -> List[bytes]: @@ -812,10 +834,13 @@ class DataStore: hint_keys_values: Optional[Dict[bytes, bytes]] = None, use_optimized: bool = True, status: Status = Status.PENDING, - ) -> bytes32: + root: Optional[Root] = None, + ) -> InsertResult: async with self.db_wrapper.writer(): - was_empty = await self.table_is_empty(tree_id=tree_id) - root = await self.get_tree_root(tree_id=tree_id) + if root is None: + root = await self.get_tree_root(tree_id=tree_id) + + was_empty = root.node_hash is None if not was_empty: if hint_keys_values is None: @@ -842,7 +867,7 @@ class DataStore: if side is not None: raise Exception("Tree was empty so side must be unspecified, got: {side!r}") - await self._insert_root( + new_root = await self._insert_root( tree_id=tree_id, node_hash=new_terminal_node_hash, status=status, @@ -857,12 +882,20 @@ class DataStore: if use_optimized: ancestors: List[InternalNode] = await self.get_ancestors_optimized( - node_hash=reference_node_hash, tree_id=tree_id + node_hash=reference_node_hash, + tree_id=tree_id, + generation=root.generation, + root_hash=root.node_hash, ) else: - ancestors = await self.get_ancestors_optimized(node_hash=reference_node_hash, tree_id=tree_id) + ancestors = await self.get_ancestors_optimized( + node_hash=reference_node_hash, + tree_id=tree_id, + generation=root.generation, + root_hash=root.node_hash, + ) ancestors_2: List[InternalNode] = await self.get_ancestors( - node_hash=reference_node_hash, tree_id=tree_id + node_hash=reference_node_hash, tree_id=tree_id, root_hash=root.node_hash ) if ancestors != ancestors_2: raise RuntimeError("Ancestors optimized didn't produce the expected result.") @@ -902,18 +935,20 @@ class DataStore: new_hash = await self._insert_internal_node(left_hash=left, right_hash=right) insert_ancestors_cache.append((left, right, tree_id)) - await self._insert_root( + new_root = await self._insert_root( tree_id=tree_id, node_hash=new_hash, status=status, + generation=new_generation, ) + if status == Status.COMMITTED: for left_hash, right_hash, tree_id in insert_ancestors_cache: await self._insert_ancestor_table(left_hash, right_hash, tree_id, new_generation) if hint_keys_values is not None: hint_keys_values[bytes(key)] = value - return new_terminal_node_hash + return InsertResult(node_hash=new_terminal_node_hash, root=new_root) async def delete( self, @@ -922,23 +957,31 @@ class DataStore: hint_keys_values: Optional[Dict[bytes, bytes]] = None, use_optimized: bool = True, status: Status = Status.PENDING, - ) -> None: + root: Optional[Root] = None, + ) -> Optional[Root]: + root_hash = None if root is None else root.node_hash async with self.db_wrapper.writer(): if hint_keys_values is None: node = await self.get_node_by_key(key=key, tree_id=tree_id) else: if bytes(key) not in hint_keys_values: log.debug(f"Request to delete an unknown key ignored: {key.hex()}") - return + return root value = hint_keys_values[bytes(key)] node_hash = leaf_hash(key=key, value=value) node = TerminalNode(node_hash, key, value) del hint_keys_values[bytes(key)] if use_optimized: - ancestors: List[InternalNode] = await self.get_ancestors_optimized(node_hash=node.hash, tree_id=tree_id) + ancestors: List[InternalNode] = await self.get_ancestors_optimized( + node_hash=node.hash, tree_id=tree_id, root_hash=root_hash + ) else: - ancestors = await self.get_ancestors_optimized(node_hash=node.hash, tree_id=tree_id) - ancestors_2: List[InternalNode] = await self.get_ancestors(node_hash=node.hash, tree_id=tree_id) + ancestors = await self.get_ancestors_optimized( + node_hash=node.hash, tree_id=tree_id, root_hash=root_hash + ) + ancestors_2: List[InternalNode] = await self.get_ancestors( + node_hash=node.hash, tree_id=tree_id, root_hash=root_hash + ) if ancestors != ancestors_2: raise RuntimeError("Ancestors optimized didn't produce the expected result.") @@ -946,30 +989,29 @@ class DataStore: raise RuntimeError("Tree exceeded max height of 62.") if len(ancestors) == 0: # the only node is being deleted - await self._insert_root( + return await self._insert_root( tree_id=tree_id, node_hash=None, status=status, ) - return - parent = ancestors[0] other_hash = parent.other_child_hash(hash=node.hash) if len(ancestors) == 1: # the parent is the root so the other side will become the new root - await self._insert_root( + return await self._insert_root( tree_id=tree_id, node_hash=other_hash, status=status, ) - return - old_child_hash = parent.hash new_child_hash = other_hash - new_generation = await self.get_tree_generation(tree_id) + 1 + if root is None: + new_generation = await self.get_tree_generation(tree_id) + 1 + else: + new_generation = root.generation + 1 # update ancestors after inserting root, to keep table constraints. insert_ancestors_cache: List[Tuple[bytes32, bytes32, bytes32]] = [] # more parents to handle so let's traverse them @@ -987,16 +1029,17 @@ class DataStore: insert_ancestors_cache.append((left_hash, right_hash, tree_id)) old_child_hash = ancestor.hash - await self._insert_root( + new_root = await self._insert_root( tree_id=tree_id, node_hash=new_child_hash, status=status, + generation=new_generation, ) if status == Status.COMMITTED: for left_hash, right_hash, tree_id in insert_ancestors_cache: await self._insert_ancestor_table(left_hash, right_hash, tree_id, new_generation) - return + return new_root async def insert_batch( self, @@ -1005,8 +1048,14 @@ class DataStore: status: Status = Status.PENDING, ) -> Optional[bytes32]: async with self.db_wrapper.writer(): - hint_keys_values = await self.get_keys_values_dict(tree_id) old_root = await self.get_tree_root(tree_id) + root_hash = old_root.node_hash + if old_root.node_hash is None: + hint_keys_values = {} + else: + hint_keys_values = await self.get_keys_values_dict(tree_id, root_hash=root_hash) + + intermediate_root: Optional[Root] = old_root for change in changelist: if change["action"] == "insert": key = change["key"] @@ -1014,11 +1063,14 @@ class DataStore: reference_node_hash = change.get("reference_node_hash", None) side = change.get("side", None) if reference_node_hash is None and side is None: - await self.autoinsert(key, value, tree_id, hint_keys_values, True, Status.COMMITTED) + insert_result = await self.autoinsert( + key, value, tree_id, hint_keys_values, True, Status.COMMITTED, root=intermediate_root + ) + intermediate_root = insert_result.root else: if reference_node_hash is None or side is None: raise Exception("Provide both reference_node_hash and side or neither.") - await self.insert( + insert_result = await self.insert( key, value, tree_id, @@ -1027,10 +1079,14 @@ class DataStore: hint_keys_values, True, Status.COMMITTED, + root=intermediate_root, ) + intermediate_root = insert_result.root elif change["action"] == "delete": key = change["key"] - await self.delete(key, tree_id, hint_keys_values, True, Status.COMMITTED) + intermediate_root = await self.delete( + key, tree_id, hint_keys_values, True, Status.COMMITTED, root=intermediate_root + ) else: raise Exception(f"Operation in batch is not insert or delete: {change}") diff --git a/tests/core/data_layer/test_data_store.py b/tests/core/data_layer/test_data_store.py index c6af37b1c7..263b828d13 100644 --- a/tests/core/data_layer/test_data_store.py +++ b/tests/core/data_layer/test_data_store.py @@ -174,8 +174,8 @@ async def test_insert_over_empty(data_store: DataStore, tree_id: bytes32) -> Non key = b"\x01\x02" value = b"abc" - node_hash = await data_store.insert(key=key, value=value, tree_id=tree_id, reference_node_hash=None, side=None) - assert node_hash == leaf_hash(key=key, value=value) + insert_result = await data_store.insert(key=key, value=value, tree_id=tree_id, reference_node_hash=None, side=None) + assert insert_result.node_hash == leaf_hash(key=key, value=value) @pytest.mark.asyncio @@ -188,7 +188,7 @@ async def test_insert_increments_generation(data_store: DataStore, tree_id: byte node_hash = None for key, expected_generation in zip(keys, itertools.count(start=1)): - node_hash = await data_store.insert( + insert_result = await data_store.insert( key=key, value=value, tree_id=tree_id, @@ -196,6 +196,7 @@ async def test_insert_increments_generation(data_store: DataStore, tree_id: byte side=None if node_hash is None else Side.LEFT, status=Status.COMMITTED, ) + node_hash = insert_result.node_hash generation = await data_store.get_tree_generation(tree_id=tree_id) generations.append(generation) expected.append(expected_generation) @@ -344,7 +345,7 @@ async def test_get_ancestors_optimized(data_store: DataStore, tree_id: bytes32) node_count += 1 side = None if node_hash is None else data_store.get_side_for_seed(seed) - node_hash = await data_store.insert( + insert_result = await data_store.insert( key=key, value=value, tree_id=tree_id, @@ -353,6 +354,7 @@ async def test_get_ancestors_optimized(data_store: DataStore, tree_id: bytes32) use_optimized=False, status=Status.COMMITTED, ) + node_hash = insert_result.node_hash if node_hash is not None: generation = await data_store.get_tree_generation(tree_id=tree_id) current_ancestors = await data_store.get_ancestors(node_hash=node_hash, tree_id=tree_id) @@ -457,6 +459,46 @@ async def test_batch_update(data_store: DataStore, tree_id: bytes32, use_optimiz ancestors[node.right_hash] = node_hash +@pytest.mark.parametrize(argnames="side", argvalues=list(Side)) +@pytest.mark.asyncio +async def test_insert_batch_reference_and_side( + data_store: DataStore, + tree_id: bytes32, + side: Side, +) -> None: + insert_result = await data_store.autoinsert( + key=b"key1", + value=b"value1", + tree_id=tree_id, + status=Status.COMMITTED, + ) + + new_root_hash = await data_store.insert_batch( + tree_id=tree_id, + changelist=[ + { + "action": "insert", + "key": b"key2", + "value": b"value2", + "reference_node_hash": insert_result.node_hash, + "side": side, + }, + ], + ) + assert new_root_hash is not None, "batch insert failed or failed to update root" + + parent = await data_store.get_node(new_root_hash) + assert isinstance(parent, InternalNode) + if side == Side.LEFT: + child = await data_store.get_node(parent.left_hash) + assert parent.left_hash == child.hash + elif side == Side.RIGHT: + child = await data_store.get_node(parent.right_hash) + assert parent.right_hash == child.hash + else: # pragma: no cover + raise Exception("invalid side for test") + + @pytest.mark.asyncio async def test_ancestor_table_unique_inserts(data_store: DataStore, tree_id: bytes32) -> None: await add_0123_example(data_store=data_store, tree_id=tree_id) @@ -503,7 +545,7 @@ async def test_inserting_duplicate_key_fails( ) -> None: key = b"\x05" - first_hash = await data_store.insert( + insert_result = await data_store.insert( key=key, value=first_value, tree_id=tree_id, @@ -517,7 +559,7 @@ async def test_inserting_duplicate_key_fails( key=key, value=second_value, tree_id=tree_id, - reference_node_hash=first_hash, + reference_node_hash=insert_result.node_hash, side=Side.RIGHT, ) @@ -528,7 +570,7 @@ async def test_inserting_duplicate_key_fails( key=key, value=second_value, tree_id=tree_id, - reference_node_hash=first_hash, + reference_node_hash=insert_result.node_hash, side=Side.RIGHT, hint_keys_values=hint_keys_values, ) @@ -575,8 +617,8 @@ async def test_autoinsert_balances_from_scratch(data_store: DataStore, tree_id: for i in range(2000): key = (i + 100).to_bytes(4, byteorder="big") value = (i + 200).to_bytes(4, byteorder="big") - node_hash = await data_store.autoinsert(key, value, tree_id, hint_keys_values, status=Status.COMMITTED) - hashes.append(node_hash) + insert_result = await data_store.autoinsert(key, value, tree_id, hint_keys_values, status=Status.COMMITTED) + hashes.append(insert_result.node_hash) heights = {node_hash: len(await data_store.get_ancestors_optimized(node_hash, tree_id)) for node_hash in hashes} too_tall = {hash: height for hash, height in heights.items() if height > 14} @@ -595,10 +637,10 @@ async def test_autoinsert_balances_gaps(data_store: DataStore, tree_id: bytes32) key = (i + 100).to_bytes(4, byteorder="big") value = (i + 200).to_bytes(4, byteorder="big") if i == 0 or i > 10: - node_hash = await data_store.autoinsert(key, value, tree_id, hint_keys_values, status=Status.COMMITTED) + insert_result = await data_store.autoinsert(key, value, tree_id, hint_keys_values, status=Status.COMMITTED) else: reference_node_hash = await data_store.get_terminal_node_for_seed(tree_id, bytes32([0] * 32)) - node_hash = await data_store.insert( + insert_result = await data_store.insert( key=key, value=value, tree_id=tree_id, @@ -607,9 +649,9 @@ async def test_autoinsert_balances_gaps(data_store: DataStore, tree_id: bytes32) hint_keys_values=hint_keys_values, status=Status.COMMITTED, ) - ancestors = await data_store.get_ancestors_optimized(node_hash, tree_id) + ancestors = await data_store.get_ancestors_optimized(insert_result.node_hash, tree_id) assert len(ancestors) == i - hashes.append(node_hash) + hashes.append(insert_result.node_hash) heights = {node_hash: len(await data_store.get_ancestors_optimized(node_hash, tree_id)) for node_hash in hashes} too_tall = {hash: height for hash, height in heights.items() if height > 14} @@ -1081,7 +1123,7 @@ async def test_kv_diff(data_store: DataStore, tree_id: bytes32) -> None: @pytest.mark.asyncio async def test_kv_diff_2(data_store: DataStore, tree_id: bytes32) -> None: - node_hash = await data_store.insert( + insert_result = await data_store.insert( key=b"000", value=b"000", tree_id=tree_id, @@ -1090,11 +1132,11 @@ async def test_kv_diff_2(data_store: DataStore, tree_id: bytes32) -> None: ) empty_hash = bytes32([0] * 32) invalid_hash = bytes32([0] * 31 + [1]) - diff_1 = await data_store.get_kv_diff(tree_id, empty_hash, node_hash) + diff_1 = await data_store.get_kv_diff(tree_id, empty_hash, insert_result.node_hash) assert diff_1 == set([DiffData(OperationType.INSERT, b"000", b"000")]) - diff_2 = await data_store.get_kv_diff(tree_id, node_hash, empty_hash) + diff_2 = await data_store.get_kv_diff(tree_id, insert_result.node_hash, empty_hash) assert diff_2 == set([DiffData(OperationType.DELETE, b"000", b"000")]) - diff_3 = await data_store.get_kv_diff(tree_id, invalid_hash, node_hash) + diff_3 = await data_store.get_kv_diff(tree_id, invalid_hash, insert_result.node_hash) assert diff_3 == set() diff --git a/tests/core/data_layer/util.py b/tests/core/data_layer/util.py index 1cbd686090..77c5b029ad 100644 --- a/tests/core/data_layer/util.py +++ b/tests/core/data_layer/util.py @@ -34,7 +34,7 @@ async def general_insert( reference_node_hash: bytes32, side: Optional[Side], ) -> bytes32: - return await data_store.insert( + insert_result = await data_store.insert( key=key, value=value, tree_id=tree_id, @@ -42,6 +42,7 @@ async def general_insert( side=side, status=Status.COMMITTED, ) + return insert_result.node_hash @dataclass(frozen=True)