* Hanlde NFT reorg

* Fix pre-commit

* Change approach & add more unit tests

* Add adjust test timeout

* Fix unit tests

* Fix unit tests

* Revert unit test changes

* Unit test

* Unit test

* Unit test

* Remove limit

* Resolve comments

* Fix pre-commit
This commit is contained in:
Kronus91
2022-08-05 13:27:08 -05:00
committed by GitHub
parent 48a8347c83
commit 2fa9f768ea
7 changed files with 187 additions and 35 deletions
+1
View File
@@ -1696,6 +1696,7 @@ class WalletRpcApi:
None,
full_puzzle,
launcher_coin[0].spent_height,
coin_state.created_height if coin_state.created_height else uint32(0),
)
)
except Exception as e:
+9
View File
@@ -79,11 +79,20 @@ class NFTInfo(Streamable):
@streamable
@dataclass(frozen=True)
class NFTCoinInfo(Streamable):
"""The launcher coin ID of the NFT"""
nft_id: bytes32
"""The latest coin of the NFT"""
coin: Coin
"""NFT lineage proof"""
lineage_proof: Optional[LineageProof]
"""NFT full puzzle"""
full_puzzle: Program
"""NFT minting block height"""
mint_height: uint32
"""The block height of the latest coin"""
latest_height: uint32 = uint32(0)
"""If the NFT is in the transaction"""
pending_transaction: bool = False
+17 -13
View File
@@ -224,16 +224,13 @@ class NFTWallet:
# all is well, lets add NFT to our local db
parent_coin = None
coin_record = await self.wallet_state_manager.coin_store.get_coin_record(coin_name)
if coin_record is None:
coin_states: Optional[List[CoinState]] = await self.wallet_state_manager.wallet_node.get_coin_state(
[coin_name]
)
if coin_states is not None:
parent_coin = coin_states[0].coin
if coin_record is not None:
parent_coin = coin_record.coin
if parent_coin is None:
confirmed_height = None
coin_states: Optional[List[CoinState]] = await self.wallet_state_manager.wallet_node.get_coin_state([coin_name])
if coin_states is not None:
parent_coin = coin_states[0].coin
confirmed_height = coin_states[0].spent_height
if parent_coin is None or confirmed_height is None:
raise ValueError("Error finding parent")
await self.add_coin(
@@ -242,6 +239,7 @@ class NFTWallet:
child_puzzle,
LineageProof(parent_coin.parent_coin_info, parent_inner_puzhash, uint64(parent_coin.amount)),
mint_height,
confirmed_height,
)
async def add_coin(
@@ -251,24 +249,25 @@ class NFTWallet:
puzzle: Program,
lineage_proof: LineageProof,
mint_height: uint32,
confirmed_height: uint32,
) -> None:
my_nft_coins = self.my_nft_coins
for coin_info in my_nft_coins:
if coin_info.coin == coin:
my_nft_coins.remove(coin_info)
new_nft = NFTCoinInfo(nft_id, coin, lineage_proof, puzzle, mint_height)
new_nft = NFTCoinInfo(nft_id, coin, lineage_proof, puzzle, mint_height, confirmed_height)
my_nft_coins.append(new_nft)
await self.wallet_state_manager.nft_store.save_nft(self.id(), self.get_did(), new_nft)
await self.wallet_state_manager.add_interested_coin_ids([coin.name()])
self.wallet_state_manager.state_changed("nft_coin_added", self.wallet_info.id)
return
async def remove_coin(self, coin: Coin) -> None:
async def remove_coin(self, coin: Coin, height: uint32) -> None:
my_nft_coins = self.my_nft_coins
for coin_info in my_nft_coins:
if coin_info.coin == coin:
my_nft_coins.remove(coin_info)
await self.wallet_state_manager.nft_store.delete_nft(coin_info.nft_id)
await self.wallet_state_manager.nft_store.delete_nft(coin_info.nft_id, height)
self.wallet_state_manager.state_changed("nft_coin_removed", self.wallet_info.id)
return
@@ -471,6 +470,10 @@ class NFTWallet:
def get_current_nfts(self) -> List[NFTCoinInfo]:
return self.my_nft_coins
async def load_current_nft(self) -> List[NFTCoinInfo]:
self.my_nft_coins = await self.wallet_state_manager.nft_store.get_nft_list(wallet_id=self.wallet_id)
return self.my_nft_coins
async def update_coin_status(self, coin_id: bytes32, pending_transaction: bool) -> None:
my_nft_coins = self.my_nft_coins
target_nft: Optional[NFTCoinInfo] = None
@@ -486,6 +489,7 @@ class NFTWallet:
target_nft.lineage_proof,
target_nft.full_puzzle,
target_nft.mint_height,
target_nft.latest_height,
pending_transaction,
)
my_nft_coins.append(new_nft)
+56 -9
View File
@@ -10,6 +10,7 @@ from chia.wallet.lineage_proof import LineageProof
from chia.wallet.nft_wallet.nft_info import DEFAULT_STATUS, IN_TRANSACTION_STATUS, NFTCoinInfo
_T_WalletNftStore = TypeVar("_T_WalletNftStore", bound="WalletNftStore")
REMOVE_BUFF_BLOCKS = 1000
class WalletNftStore:
@@ -41,16 +42,28 @@ class WalletNftStore:
await conn.execute("CREATE INDEX IF NOT EXISTS nft_coin_id on users_nfts(nft_coin_id)")
await conn.execute("CREATE INDEX IF NOT EXISTS nft_wallet_id on users_nfts(wallet_id)")
await conn.execute("CREATE INDEX IF NOT EXISTS nft_did_id on users_nfts(did_id)")
try:
# These are patched columns for resolving reorg issue
await conn.execute("ALTER TABLE users_nfts ADD COLUMN removed_height bigint")
await conn.execute("ALTER TABLE users_nfts ADD COLUMN latest_height bigint")
await conn.execute("CREATE INDEX IF NOT EXISTS removed_nft_height on users_nfts(removed_height)")
await conn.execute("CREATE INDEX IF NOT EXISTS latest_nft_height on users_nfts(latest_height)")
except Exception:
pass
return self
async def delete_nft(self, nft_id: bytes32) -> None:
async def delete_nft(self, nft_id: bytes32, height: uint32) -> None:
async with self.db_wrapper.writer_maybe_transaction() as conn:
await (await conn.execute("DELETE FROM users_nfts where nft_id=?", (nft_id.hex(),))).close()
# Remove NFT in the users_nfts table
await (
await conn.execute("UPDATE users_nfts SET removed_height=? WHERE nft_id=?", (int(height), nft_id.hex()))
).close()
async def save_nft(self, wallet_id: uint32, did_id: Optional[bytes32], nft_coin_info: NFTCoinInfo) -> None:
async with self.db_wrapper.writer_maybe_transaction() as conn:
cursor = await conn.execute(
"INSERT or REPLACE INTO users_nfts VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)",
"INSERT or REPLACE INTO users_nfts VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
(
nft_coin_info.nft_id.hex(),
nft_coin_info.coin.name().hex(),
@@ -63,21 +76,35 @@ class WalletNftStore:
int(nft_coin_info.mint_height),
IN_TRANSACTION_STATUS if nft_coin_info.pending_transaction else DEFAULT_STATUS,
bytes(nft_coin_info.full_puzzle),
None,
int(nft_coin_info.latest_height),
),
)
await cursor.close()
# Rotate the old removed NFTs, they are not possible to be reorged
await (
await conn.execute(
"DELETE FROM users_nfts WHERE removed_height is not NULL and removed_height<?",
(int(nft_coin_info.latest_height) - REMOVE_BUFF_BLOCKS,),
)
).close()
async def get_nft_list(
self, wallet_id: Optional[uint32] = None, did_id: Optional[bytes32] = None
) -> List[NFTCoinInfo]:
sql: str = "SELECT nft_id, coin, lineage_proof, mint_height, status, full_puzzle from users_nfts"
sql: str = (
"SELECT nft_id, coin, lineage_proof, mint_height, status, full_puzzle, latest_height"
" from users_nfts WHERE"
)
if wallet_id is not None and did_id is None:
sql += f" where wallet_id={wallet_id}"
sql += f" wallet_id={wallet_id}"
if wallet_id is None and did_id is not None:
sql += f" where did_id='{did_id.hex()}'"
sql += f" did_id='{did_id.hex()}'"
if wallet_id is not None and did_id is not None:
sql += f" where did_id='{did_id.hex()}' and wallet_id={wallet_id}"
sql += f" did_id='{did_id.hex()}' and wallet_id={wallet_id}"
if wallet_id is not None or did_id is not None:
sql += " and"
sql += " removed_height is NULL"
async with self.db_wrapper.reader_no_transaction() as conn:
rows = await conn.execute_fetchall(sql)
@@ -88,6 +115,7 @@ class WalletNftStore:
None if row[2] is None else LineageProof.from_json_dict(json.loads(row[2])),
Program.from_bytes(row[5]),
uint32(row[3]),
uint32(row[6]),
row[4] == IN_TRANSACTION_STATUS,
)
for row in rows
@@ -97,7 +125,8 @@ class WalletNftStore:
async with self.db_wrapper.reader_no_transaction() as conn:
row = await execute_fetchone(
conn,
"SELECT nft_id, coin, lineage_proof, mint_height, status, full_puzzle from users_nfts WHERE nft_id=?",
"SELECT nft_id, coin, lineage_proof, mint_height, status, full_puzzle, latest_height"
" from users_nfts WHERE removed_height is NULL and nft_id=?",
(nft_id.hex(),),
)
@@ -110,5 +139,23 @@ class WalletNftStore:
None if row[2] is None else LineageProof.from_json_dict(json.loads(row[2])),
Program.from_bytes(row[5]),
uint32(row[3]),
uint32(row[6]),
row[4] == IN_TRANSACTION_STATUS,
)
async def rollback_to_block(self, height: int) -> None:
"""
Rolls back the blockchain to block_index. All coins confirmed after this point are removed.
All coins spent after this point are set to unspent. Can be -1 (rollback all)
"""
async with self.db_wrapper.writer_maybe_transaction() as conn:
# Remove reorged NFTs
await (await conn.execute("DELETE FROM users_nfts WHERE latest_height>?", (height,))).close()
# Retrieve removed NFTs
await (
await conn.execute(
"UPDATE users_nfts SET removed_height = null WHERE removed_height>?",
(height,),
)
).close()
+12 -5
View File
@@ -651,7 +651,7 @@ class WalletStateManager:
# First spend where 1 mojo coin -> Singleton launcher -> NFT -> NFT
uncurried_nft = UncurriedNFT.uncurry(mod, curried_args)
if uncurried_nft is not None:
return await self.handle_nft(coin_spend, uncurried_nft)
return await self.handle_nft(coin_spend, uncurried_nft, parent_coin_state)
# Check if the coin is a DID
did_curried_args = match_did_puzzle(mod, curried_args)
@@ -790,12 +790,13 @@ class WalletStateManager:
return wallet_id, wallet_type
async def handle_nft(
self, coin_spend: CoinSpend, uncurried_nft: UncurriedNFT
self, coin_spend: CoinSpend, uncurried_nft: UncurriedNFT, parent_coin_state: CoinState
) -> Tuple[Optional[uint32], Optional[WalletType]]:
"""
Handle the new coin when it is a NFT
:param coin_spend: New coin spend
:param uncurried_nft: Uncurried NFT
:param parent_coin_state: Parent coin state
:return: Wallet ID & Wallet Type
"""
wallet_id = None
@@ -836,7 +837,6 @@ class WalletStateManager:
uncurried_nft.singleton_launcher_id.hex(),
)
return wallet_id, wallet_type
for wallet_info in await self.get_all_wallet_info_entries(wallet_type=WalletType.NFT):
nft_wallet_info: NFTWalletInfo = NFTWalletInfo.from_json_dict(json.loads(wallet_info.data))
if nft_wallet_info.did_id == old_did_id:
@@ -846,7 +846,8 @@ class WalletStateManager:
old_did_id,
)
nft_wallet: NFTWallet = self.wallets[wallet_info.id]
await nft_wallet.remove_coin(coin_spend.coin)
if parent_coin_state.spent_height is not None:
await nft_wallet.remove_coin(coin_spend.coin, parent_coin_state.spent_height)
if nft_wallet_info.did_id == new_did_id:
self.log.info(
"Adding new NFT, NFT_ID:%s, DID_ID:%s",
@@ -1112,7 +1113,7 @@ class WalletStateManager:
elif record.wallet_type == WalletType.NFT:
if coin_state.spent_height is not None:
nft_wallet = self.wallets[uint32(record.wallet_id)]
await nft_wallet.remove_coin(coin_state.coin)
await nft_wallet.remove_coin(coin_state.coin, coin_state.spent_height)
# Check if a child is a singleton launcher
if children is None:
@@ -1403,6 +1404,7 @@ class WalletStateManager:
Rolls back and updates the coin_store and transaction store. It's possible this height
is the tip, or even beyond the tip.
"""
await self.nft_store.rollback_to_block(height)
await self.coin_store.rollback_to_block(height)
reorged: List[TransactionRecord] = await self.tx_store.get_transaction_above(height)
await self.tx_store.rollback_to_block(height)
@@ -1421,6 +1423,11 @@ class WalletStateManager:
remove: bool = await wallet.rewind(height)
if remove:
remove_ids.append(wallet_id)
if wallet.type() == WalletType.NFT.value:
# Refresh the NFTs list
await wallet.load_current_nft()
if len(wallet.my_nft_coins) == 0:
remove_ids.append(wallet_id)
for wallet_id in remove_ids:
await self.user_store.delete_wallet(wallet_id)
+38 -6
View File
@@ -9,7 +9,7 @@ from chia.consensus.block_rewards import calculate_base_farmer_reward, calculate
from chia.full_node.mempool_manager import MempoolManager
from chia.rpc.wallet_rpc_api import WalletRpcApi
from chia.simulator.full_node_simulator import FullNodeSimulator
from chia.simulator.simulator_protocol import FarmNewBlockProtocol
from chia.simulator.simulator_protocol import FarmNewBlockProtocol, ReorgProtocol
from chia.simulator.time_out_assert import time_out_assert, time_out_assert_not_none
from chia.types.blockchain_format.program import Program
from chia.types.blockchain_format.sized_bytes import bytes32
@@ -23,6 +23,7 @@ from chia.wallet.nft_wallet.nft_wallet import NFTWallet
from chia.wallet.util.address_type import AddressType
from chia.wallet.util.compute_memos import compute_memos
from chia.wallet.util.wallet_types import WalletType
from chia.wallet.wallet_state_manager import WalletStateManager
async def tx_in_pool(mempool: MempoolManager, tx_id: bytes32) -> bool:
@@ -32,6 +33,14 @@ async def tx_in_pool(mempool: MempoolManager, tx_id: bytes32) -> bool:
return True
async def get_nft_number(wallet: NFTWallet) -> int:
return len(await wallet.load_current_nft())
async def get_wallet_number(manager: WalletStateManager) -> int:
return len(manager.wallets)
async def wait_rpc_state_condition(
timeout: int,
coroutine: Callable[[Dict[str, Any]], Awaitable[Dict]],
@@ -225,7 +234,21 @@ async def test_nft_wallet_creation_and_transfer(two_wallet_nodes: Any, trusted:
for i in range(1, num_blocks):
await full_node_api.farm_new_transaction_block(FarmNewBlockProtocol(ph))
await time_out_assert(20, len, 1, nft_wallet_0.my_nft_coins)
await time_out_assert(10, len, 1, nft_wallet_0.my_nft_coins)
await time_out_assert(10, wallet_0.get_unconfirmed_balance, 4000000000000 - 1)
await time_out_assert(10, wallet_0.get_confirmed_balance, 4000000000000 - 1)
# Test Reorg mint
height = full_node_api.full_node.blockchain.get_peak_height()
if height is None:
assert False
await full_node_api.reorg_from_index_to_new_index(ReorgProtocol(uint32(height - 1), uint32(height + 1), ph1))
await time_out_assert(15, get_nft_number, 0, nft_wallet_0)
await time_out_assert(15, get_wallet_number, 1, wallet_node_0.wallet_state_manager)
nft_wallet_0 = await NFTWallet.create_new_nft_wallet(
wallet_node_0.wallet_state_manager, wallet_0, name="NFT WALLET 1"
)
metadata = Program.to(
[
("u", ["https://www.test.net/logo.svg"]),
@@ -233,8 +256,9 @@ async def test_nft_wallet_creation_and_transfer(two_wallet_nodes: Any, trusted:
]
)
await time_out_assert(20, wallet_0.get_unconfirmed_balance, 4000000000000 - 1)
await time_out_assert(20, wallet_0.get_confirmed_balance, 4000000000000 - 1)
await time_out_assert(10, wallet_0.get_unconfirmed_balance, 4000000000000 - 1)
await time_out_assert(10, wallet_0.get_confirmed_balance, 4000000000000)
sb = await nft_wallet_0.generate_new_nft(metadata)
assert sb
# ensure hints are generated
@@ -265,10 +289,12 @@ async def test_nft_wallet_creation_and_transfer(two_wallet_nodes: Any, trusted:
await full_node_api.farm_new_transaction_block(FarmNewBlockProtocol(ph1))
await time_out_assert(15, len, 1, nft_wallet_0.my_nft_coins)
await time_out_assert(15, len, 1, nft_wallet_1.my_nft_coins)
coins = nft_wallet_1.my_nft_coins
assert len(coins) == 1
await time_out_assert(15, wallet_1.get_pending_change_balance, 0)
# Send it back to original owner
txs = await nft_wallet_1.generate_signed_transaction([uint64(coins[0].coin.amount)], [ph], coins={coins[0].coin})
assert len(txs) == 1
@@ -281,13 +307,19 @@ async def test_nft_wallet_creation_and_transfer(two_wallet_nodes: Any, trusted:
for i in range(1, num_blocks):
await full_node_api.farm_new_transaction_block(FarmNewBlockProtocol(ph1))
for i in range(1, num_blocks):
await full_node_api.farm_new_transaction_block(FarmNewBlockProtocol(ph))
await time_out_assert(30, wallet_node_0.wallet_state_manager.lock.locked, False)
await time_out_assert(15, len, 2, nft_wallet_0.my_nft_coins)
await time_out_assert(15, len, 0, nft_wallet_1.my_nft_coins)
# Test Reorg
height = full_node_api.full_node.blockchain.get_peak_height()
if height is None:
assert False
await full_node_api.reorg_from_index_to_new_index(ReorgProtocol(uint32(height - 1), uint32(height + 1), ph1))
await time_out_assert(15, get_nft_number, 1, nft_wallet_0)
await time_out_assert(15, get_nft_number, 1, nft_wallet_1)
@pytest.mark.parametrize(
"trusted",
+54 -2
View File
@@ -12,7 +12,7 @@ from tests.util.db_connection import DBConnection
class TestNftStore:
@pytest.mark.asyncio
async def test_nft_store(self) -> None:
async def test_nft_insert(self) -> None:
async with DBConnection(1) as wrapper:
db = await WalletNftStore.create(wrapper)
a_bytes32 = bytes32.fromhex("09287c75377c63fd6a3a4d6658abed03e9a521e0436b1f83cdf4af99341ce8f1")
@@ -22,6 +22,7 @@ class TestNftStore:
Coin(a_bytes32, a_bytes32, uint64(1)),
LineageProof(a_bytes32, a_bytes32, uint64(1)),
puzzle,
uint32(1),
uint32(10),
)
# Test save
@@ -32,6 +33,57 @@ class TestNftStore:
assert nft == (await db.get_nft_list(did_id=a_bytes32))[0]
assert nft == (await db.get_nft_list(wallet_id=uint32(1), did_id=a_bytes32))[0]
assert nft == await db.get_nft_by_id(a_bytes32)
@pytest.mark.asyncio
async def test_nft_remove(self) -> None:
async with DBConnection(1) as wrapper:
db = await WalletNftStore.create(wrapper)
a_bytes32 = bytes32.fromhex("09287c75377c63fd6a3a4d6658abed03e9a521e0436b1f83cdf4af99341ce8f1")
puzzle = Program.to(["A Test puzzle"])
nft = NFTCoinInfo(
a_bytes32,
Coin(a_bytes32, a_bytes32, uint64(1)),
LineageProof(a_bytes32, a_bytes32, uint64(1)),
puzzle,
uint32(1),
uint32(10),
)
# Test save
await db.save_nft(uint32(1), a_bytes32, nft)
# Test delete
await db.delete_nft(a_bytes32)
await db.delete_nft(a_bytes32, uint32(11))
assert await db.get_nft_by_id(a_bytes32) is None
@pytest.mark.asyncio
async def test_nft_reorg(self) -> None:
async with DBConnection(1) as wrapper:
db = await WalletNftStore.create(wrapper)
a_bytes32 = bytes32.fromhex("09287c75377c63fd6a3a4d6658abed03e9a521e0436b1f83cdf4af99341ce8f1")
a_bytes32_1 = bytes32.fromhex("09287c75377c63fd6a3a4d6658abed03e9a521e0436b1f83cdf4af99341ce8f2")
puzzle = Program.to(["A Test puzzle"])
nft = NFTCoinInfo(
a_bytes32,
Coin(a_bytes32, a_bytes32, uint64(1)),
LineageProof(a_bytes32, a_bytes32, uint64(1)),
puzzle,
uint32(1),
uint32(10),
)
# Test save
await db.save_nft(uint32(1), a_bytes32, nft)
# Test delete
await db.delete_nft(a_bytes32, uint32(11))
assert await db.get_nft_by_id(a_bytes32) is None
# Test reorg
nft1 = NFTCoinInfo(
a_bytes32_1,
Coin(a_bytes32, a_bytes32, uint64(1)),
LineageProof(a_bytes32, a_bytes32, uint64(1)),
puzzle,
uint32(1),
uint32(12),
)
await db.save_nft(uint32(1), a_bytes32_1, nft1)
assert nft1 == (await db.get_nft_list(wallet_id=uint32(1)))[0]
await db.rollback_to_block(10)
assert nft == (await db.get_nft_list(wallet_id=uint32(1)))[0]