wallet_taproot.py raw
1 #!/usr/bin/env python3
2 # Copyright (c) 2021-present The Bitcoin Core developers
3 # Distributed under the MIT software license, see the accompanying
4 # file COPYING or http://www.opensource.org/licenses/mit-license.php.
5 """Test generation and spending of P2TR addresses."""
6
7 import random
8 import uuid
9
10 from decimal import Decimal
11 from test_framework.address import output_key_to_p2tr
12 from test_framework.key import H_POINT, compute_xonly_pubkey
13 from test_framework.test_framework import BitcoinTestFramework
14 from test_framework.util import assert_equal
15 from test_framework.descriptors import descsum_create
16 from test_framework.extendedkey import ExtendedPrivateKey
17 from test_framework.script import (
18 CScript,
19 MAX_PUBKEYS_PER_MULTI_A,
20 OP_CHECKSIG,
21 OP_CHECKSIGADD,
22 OP_NUMEQUAL,
23 taproot_construct,
24 )
25 from test_framework.segwit_addr import encode_segwit_address
26
27 def key(hex_key):
28 """Construct an x-only pubkey from its hex representation."""
29 return bytes.fromhex(hex_key)
30
31 def pk(hex_key):
32 """Construct a script expression for taproot_construct for pk(hex_key)."""
33 return (None, CScript([bytes.fromhex(hex_key), OP_CHECKSIG]))
34
35 def multi_a(k, hex_keys, sort=False):
36 """Construct a script expression for taproot_construct for a multi_a script."""
37 xkeys = [bytes.fromhex(hex_key) for hex_key in hex_keys]
38 if sort:
39 xkeys.sort()
40 ops = [xkeys[0], OP_CHECKSIG]
41 for i in range(1, len(hex_keys)):
42 ops += [xkeys[i], OP_CHECKSIGADD]
43 ops += [k, OP_NUMEQUAL]
44 return (None, CScript(ops))
45
46 def compute_taproot_address(pubkey, scripts):
47 """Compute the address for a taproot output with given inner key and scripts."""
48 return output_key_to_p2tr(taproot_construct(pubkey, scripts).output_pubkey)
49
50 def compute_raw_taproot_address(pubkey):
51 return encode_segwit_address("bcrt", 1, pubkey)
52
53 class WalletTaprootTest(BitcoinTestFramework):
54 """Test generation and spending of P2TR address outputs."""
55
56 def set_test_params(self):
57 self.num_nodes = 2
58 self.setup_clean_chain = True
59 self.extra_args = [['-keypool=100'], ['-keypool=100']]
60
61 def skip_test_if_missing_module(self):
62 self.skip_if_no_wallet()
63
64 def setup_network(self):
65 self.setup_nodes()
66
67 def init_wallet(self, *, node):
68 pass
69
70 @staticmethod
71 def make_desc(pattern, privmap, keys, pub_only = False):
72 pat = pattern.replace("$H", H_POINT)
73 for i in range(len(privmap)):
74 if privmap[i] and not pub_only:
75 pat = pat.replace("$%i" % (i + 1), keys[i]['xprv'])
76 else:
77 pat = pat.replace("$%i" % (i + 1), keys[i]['xpub'])
78 return descsum_create(pat)
79
80 @staticmethod
81 def make_addr(treefn, keys, i):
82 args = []
83 for j in range(len(keys)):
84 args.append(keys[j]['pubs'][i])
85 tree = treefn(*args)
86 if isinstance(tree, tuple):
87 return compute_taproot_address(*tree)
88 if isinstance(tree, bytes):
89 return compute_raw_taproot_address(tree)
90 assert False
91
92 def do_test_addr(self, comment, pattern, privmap, treefn, keys):
93 self.log.info("Testing %s address derivation" % comment)
94
95 # Create wallets
96 wallet_uuid = uuid.uuid4().hex
97 self.nodes[0].createwallet(wallet_name=f"privs_tr_enabled_{wallet_uuid}", blank=True)
98 self.nodes[0].createwallet(wallet_name=f"pubs_tr_enabled_{wallet_uuid}", blank=True, disable_private_keys=True)
99 self.nodes[0].createwallet(wallet_name=f"addr_gen_{wallet_uuid}", disable_private_keys=True, blank=True)
100 privs_tr_enabled = self.nodes[0].get_wallet_rpc(f"privs_tr_enabled_{wallet_uuid}")
101 pubs_tr_enabled = self.nodes[0].get_wallet_rpc(f"pubs_tr_enabled_{wallet_uuid}")
102 addr_gen = self.nodes[0].get_wallet_rpc(f"addr_gen_{wallet_uuid}")
103
104 desc = self.make_desc(pattern, privmap, keys, False)
105 desc_pub = self.make_desc(pattern, privmap, keys, True)
106 assert_equal(self.nodes[0].getdescriptorinfo(desc)['descriptor'], desc_pub)
107 result = addr_gen.importdescriptors([{"desc": desc_pub, "active": True, "timestamp": "now"}])
108 assert result[0]['success']
109 address_type = "bech32m" if "tr" in pattern else "bech32"
110 for i in range(4):
111 addr_g = addr_gen.getnewaddress(address_type=address_type)
112 if treefn is not None:
113 addr_r = self.make_addr(treefn, keys, i)
114 assert_equal(addr_g, addr_r)
115 desc_a = addr_gen.getaddressinfo(addr_g)['desc']
116 if desc.startswith("tr("):
117 assert desc_a.startswith("tr(")
118 rederive = self.nodes[1].deriveaddresses(desc_a)
119 assert_equal(len(rederive), 1)
120 assert_equal(rederive[0], addr_g)
121
122 # tr descriptors can be imported
123 result = privs_tr_enabled.importdescriptors([{"desc": desc, "timestamp": "now"}])
124 assert result[0]['success']
125 result = pubs_tr_enabled.importdescriptors([{"desc": desc_pub, "timestamp": "now"}])
126 assert result[0]["success"]
127
128 # Cleanup
129 privs_tr_enabled.unloadwallet()
130 pubs_tr_enabled.unloadwallet()
131 addr_gen.unloadwallet()
132
133 def do_test_sendtoaddress(self, comment, pattern, privmap, treefn, keys_pay, keys_change):
134 self.log.info("Testing %s through sendtoaddress" % comment)
135
136 # Create wallets
137 wallet_uuid = uuid.uuid4().hex
138 self.nodes[0].createwallet(wallet_name=f"rpc_online_{wallet_uuid}", blank=True)
139 rpc_online = self.nodes[0].get_wallet_rpc(f"rpc_online_{wallet_uuid}")
140
141 desc_pay = self.make_desc(pattern, privmap, keys_pay)
142 desc_change = self.make_desc(pattern, privmap, keys_change)
143 desc_pay_pub = self.make_desc(pattern, privmap, keys_pay, True)
144 desc_change_pub = self.make_desc(pattern, privmap, keys_change, True)
145 assert_equal(self.nodes[0].getdescriptorinfo(desc_pay)['descriptor'], desc_pay_pub)
146 assert_equal(self.nodes[0].getdescriptorinfo(desc_change)['descriptor'], desc_change_pub)
147 result = rpc_online.importdescriptors([{"desc": desc_pay, "active": True, "timestamp": "now"}])
148 assert result[0]['success']
149 result = rpc_online.importdescriptors([{"desc": desc_change, "active": True, "timestamp": "now", "internal": True}])
150 assert result[0]['success']
151 address_type = "bech32m" if "tr" in pattern else "bech32"
152 for i in range(4):
153 addr_g = rpc_online.getnewaddress(address_type=address_type)
154 if treefn is not None:
155 addr_r = self.make_addr(treefn, keys_pay, i)
156 assert_equal(addr_g, addr_r)
157 boring_balance = int(self.boring.getbalance() * 100000000)
158 to_amnt = random.randrange(1000000, boring_balance)
159 self.boring.sendtoaddress(address=addr_g, amount=Decimal(to_amnt) / 100000000, subtractfeefromamount=True)
160 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
161 test_balance = int(rpc_online.getbalance() * 100000000)
162 ret_amnt = random.randrange(100000, test_balance)
163 # Increase fee_rate to compensate for the wallet's inability to estimate fees for script path spends.
164 res = rpc_online.sendtoaddress(address=self.boring.getnewaddress(), amount=Decimal(ret_amnt) / 100000000, subtractfeefromamount=True, fee_rate=200)
165 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
166 assert rpc_online.gettransaction(res)["confirmations"] > 0
167
168 # Cleanup
169 txid = rpc_online.sendall(recipients=[self.boring.getnewaddress()])["txid"]
170 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
171 assert rpc_online.gettransaction(txid)["confirmations"] > 0
172 rpc_online.unloadwallet()
173
174 def do_test_psbt(self, comment, pattern, privmap, treefn, keys_pay, keys_change):
175 self.log.info("Testing %s through PSBT" % comment)
176
177 # Create wallets
178 wallet_uuid = uuid.uuid4().hex
179 self.nodes[0].createwallet(wallet_name=f"psbt_online_{wallet_uuid}", disable_private_keys=True, blank=True)
180 self.nodes[1].createwallet(wallet_name=f"psbt_offline_{wallet_uuid}", blank=True)
181 self.nodes[1].createwallet(f"key_only_wallet_{wallet_uuid}", blank=True)
182 psbt_online = self.nodes[0].get_wallet_rpc(f"psbt_online_{wallet_uuid}")
183 psbt_offline = self.nodes[1].get_wallet_rpc(f"psbt_offline_{wallet_uuid}")
184 key_only_wallet = self.nodes[1].get_wallet_rpc(f"key_only_wallet_{wallet_uuid}")
185
186 desc_pay = self.make_desc(pattern, privmap, keys_pay, False)
187 desc_change = self.make_desc(pattern, privmap, keys_change, False)
188 desc_pay_pub = self.make_desc(pattern, privmap, keys_pay, True)
189 desc_change_pub = self.make_desc(pattern, privmap, keys_change, True)
190 assert_equal(self.nodes[0].getdescriptorinfo(desc_pay)['descriptor'], desc_pay_pub)
191 assert_equal(self.nodes[0].getdescriptorinfo(desc_change)['descriptor'], desc_change_pub)
192 result = psbt_online.importdescriptors([{"desc": desc_pay_pub, "active": True, "timestamp": "now"}])
193 assert result[0]['success']
194 result = psbt_online.importdescriptors([{"desc": desc_change_pub, "active": True, "timestamp": "now", "internal": True}])
195 assert result[0]['success']
196 result = psbt_offline.importdescriptors([{"desc": desc_pay, "active": True, "timestamp": "now"}])
197 assert result[0]['success']
198 result = psbt_offline.importdescriptors([{"desc": desc_change, "active": True, "timestamp": "now", "internal": True}])
199 assert result[0]['success']
200 for key in keys_pay + keys_change:
201 result = key_only_wallet.importdescriptors([{"desc": descsum_create(f"wpkh({key['xprv']}/*)"), "timestamp":"now"}])
202 assert result[0]["success"]
203 address_type = "bech32m" if "tr" in pattern else "bech32"
204 for i in range(4):
205 addr_g = psbt_online.getnewaddress(address_type=address_type)
206 if treefn is not None:
207 addr_r = self.make_addr(treefn, keys_pay, i)
208 assert_equal(addr_g, addr_r)
209 boring_balance = int(self.boring.getbalance() * 100000000)
210 to_amnt = random.randrange(1000000, boring_balance)
211 self.boring.sendtoaddress(address=addr_g, amount=Decimal(to_amnt) / 100000000, subtractfeefromamount=True)
212 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
213 test_balance = int(psbt_online.getbalance() * 100000000)
214 ret_amnt = random.randrange(100000, test_balance)
215 # Increase fee_rate to compensate for the wallet's inability to estimate fees for script path spends.
216 psbt = psbt_online.walletcreatefundedpsbt([], [{self.boring.getnewaddress(): Decimal(ret_amnt) / 100000000}], None, {"subtractFeeFromOutputs":[0], "fee_rate": 200, "change_type": address_type})['psbt']
217 res = psbt_offline.walletprocesspsbt(psbt=psbt, finalize=False)
218 for wallet in [psbt_offline, key_only_wallet]:
219 res = wallet.walletprocesspsbt(psbt=psbt, finalize=False)
220
221 decoded = wallet.decodepsbt(res["psbt"])
222 if pattern.startswith("tr("):
223 for psbtin in decoded["inputs"]:
224 assert "non_witness_utxo" not in psbtin
225 assert "witness_utxo" in psbtin
226 assert "taproot_internal_key" in psbtin
227 assert "taproot_bip32_derivs" in psbtin
228 assert "taproot_key_path_sig" in psbtin or "taproot_script_path_sigs" in psbtin
229 if "taproot_script_path_sigs" in psbtin:
230 assert "taproot_merkle_root" in psbtin
231 assert "taproot_scripts" in psbtin
232
233 rawtx = self.nodes[0].finalizepsbt(res['psbt'])['hex']
234 res = self.nodes[0].testmempoolaccept([rawtx])
235 assert res[0]["allowed"]
236
237 txid = self.nodes[0].sendrawtransaction(rawtx)
238 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
239 assert psbt_online.gettransaction(txid)['confirmations'] > 0
240
241 # Cleanup
242 psbt = psbt_online.sendall(recipients=[self.boring.getnewaddress()], psbt=True)["psbt"]
243 res = psbt_offline.walletprocesspsbt(psbt=psbt, finalize=False)
244 rawtx = self.nodes[0].finalizepsbt(res['psbt'])['hex']
245 txid = self.nodes[0].sendrawtransaction(rawtx)
246 self.generatetoaddress(self.nodes[0], 1, self.boring.getnewaddress(), sync_fun=self.no_op)
247 assert psbt_online.gettransaction(txid)['confirmations'] > 0
248 psbt_online.unloadwallet()
249 psbt_offline.unloadwallet()
250
251 def do_test(self, comment, pattern, privmap, treefn):
252 nkeys = len(privmap)
253 keys = random.sample(self.keys, nkeys * 4)
254 self.do_test_addr(comment, pattern, privmap, treefn, keys[0:nkeys])
255 self.do_test_sendtoaddress(comment, pattern, privmap, treefn, keys[0:nkeys], keys[nkeys:2*nkeys])
256 self.do_test_psbt(comment, pattern, privmap, treefn, keys[2*nkeys:3*nkeys], keys[3*nkeys:4*nkeys])
257
258 def generate_test_keys(self):
259 xprvs = [ExtendedPrivateKey.generate() for _ in range(0, 13)]
260 return [{
261 "xprv": xprv.to_string(),
262 "xpub": xprv.pubkey().to_string(),
263 "pubs": [compute_xonly_pubkey(xprv.derive_path(f"m/{i}").key.get_bytes())[0].hex() for i in range(0, 4)]
264 } for xprv in xprvs]
265
266 def run_test(self):
267 self.keys = self.generate_test_keys()
268 self.nodes[0].createwallet(wallet_name="boring")
269 self.boring = self.nodes[0].get_wallet_rpc("boring")
270
271 self.log.info("Mining blocks...")
272 gen_addr = self.boring.getnewaddress()
273 self.generatetoaddress(self.nodes[0], 101, gen_addr, sync_fun=self.no_op)
274
275 self.do_test(
276 "tr(XPRV)",
277 "tr($1/*)",
278 [True],
279 lambda k1: (key(k1), [])
280 )
281 self.do_test(
282 "tr(H,XPRV)",
283 "tr($H,pk($1/*))",
284 [True],
285 lambda k1: (key(H_POINT), [pk(k1)])
286 )
287 self.do_test(
288 "wpkh(XPRV)",
289 "wpkh($1/*)",
290 [True],
291 None
292 )
293 self.do_test(
294 "tr(XPRV,{H,{H,XPUB}})",
295 "tr($1/*,{pk($H),{pk($H),pk($2/*)}})",
296 [True, False],
297 lambda k1, k2: (key(k1), [pk(H_POINT), [pk(H_POINT), pk(k2)]])
298 )
299 self.do_test(
300 "wsh(multi(1,XPRV,XPUB))",
301 "wsh(multi(1,$1/*,$2/*))",
302 [True, False],
303 None
304 )
305 self.do_test(
306 "tr(XPRV,{XPUB,XPUB})",
307 "tr($1/*,{pk($2/*),pk($2/*)})",
308 [True, False],
309 lambda k1, k2: (key(k1), [pk(k2), pk(k2)])
310 )
311 self.do_test(
312 "tr(XPRV,{{XPUB,H},{H,XPUB}})",
313 "tr($1/*,{{pk($2/*),pk($H)},{pk($H),pk($2/*)}})",
314 [True, False],
315 lambda k1, k2: (key(k1), [[pk(k2), pk(H_POINT)], [pk(H_POINT), pk(k2)]])
316 )
317 self.do_test(
318 "tr(XPUB,{{H,{H,XPUB}},{H,{H,{H,XPRV}}}})",
319 "tr($1/*,{{pk($H),{pk($H),pk($2/*)}},{pk($H),{pk($H),{pk($H),pk($3/*)}}}})",
320 [False, False, True],
321 lambda k1, k2, k3: (key(k1), [[pk(H_POINT), [pk(H_POINT), pk(k2)]], [pk(H_POINT), [pk(H_POINT), [pk(H_POINT), pk(k3)]]]])
322 )
323 self.do_test(
324 "tr(XPRV,{XPUB,{{XPUB,{H,H}},{{H,H},XPUB}}})",
325 "tr($1/*,{pk($2/*),{{pk($2/*),{pk($H),pk($H)}},{{pk($H),pk($H)},pk($2/*)}}})",
326 [True, False],
327 lambda k1, k2: (key(k1), [pk(k2), [[pk(k2), [pk(H_POINT), pk(H_POINT)]], [[pk(H_POINT), pk(H_POINT)], pk(k2)]]])
328 )
329 self.do_test(
330 "tr(H,multi_a(1,XPRV))",
331 "tr($H,multi_a(1,$1/*))",
332 [True],
333 lambda k1: (key(H_POINT), [multi_a(1, [k1])])
334 )
335 self.do_test(
336 "tr(H,sortedmulti_a(1,XPRV,XPUB))",
337 "tr($H,sortedmulti_a(1,$1/*,$2/*))",
338 [True, False],
339 lambda k1, k2: (key(H_POINT), [multi_a(1, [k1, k2], True)])
340 )
341 self.do_test(
342 "tr(H,{H,multi_a(1,XPUB,XPRV)})",
343 "tr($H,{pk($H),multi_a(1,$1/*,$2/*)})",
344 [False, True],
345 lambda k1, k2: (key(H_POINT), [pk(H_POINT), [multi_a(1, [k1, k2])]])
346 )
347 self.do_test(
348 "tr(H,sortedmulti_a(1,XPUB,XPRV,XPRV))",
349 "tr($H,sortedmulti_a(1,$1/*,$2/*,$3/*))",
350 [False, True, True],
351 lambda k1, k2, k3: (key(H_POINT), [multi_a(1, [k1, k2, k3], True)])
352 )
353 self.do_test(
354 "tr(H,multi_a(2,XPRV,XPUB,XPRV))",
355 "tr($H,multi_a(2,$1/*,$2/*,$3/*))",
356 [True, False, True],
357 lambda k1, k2, k3: (key(H_POINT), [multi_a(2, [k1, k2, k3])])
358 )
359 self.do_test(
360 "tr(XPUB,{{XPUB,{XPUB,sortedmulti_a(2,XPRV,XPUB,XPRV)}})",
361 "tr($2/*,{pk($2/*),{pk($2/*),sortedmulti_a(2,$1/*,$2/*,$3/*)}})",
362 [True, False, True],
363 lambda k1, k2, k3: (key(k2), [pk(k2), [pk(k2), multi_a(2, [k1, k2, k3], True)]])
364 )
365 rnd_pos = random.randrange(MAX_PUBKEYS_PER_MULTI_A)
366 self.do_test(
367 "tr(XPUB,multi_a(1,H...,XPRV,H...))",
368 "tr($2/*,multi_a(1" + (",$H" * rnd_pos) + ",$1/*" + (",$H" * (MAX_PUBKEYS_PER_MULTI_A - 1 - rnd_pos)) + "))",
369 [True, False],
370 lambda k1, k2: (key(k2), [multi_a(1, ([H_POINT] * rnd_pos) + [k1] + ([H_POINT] * (MAX_PUBKEYS_PER_MULTI_A - 1 - rnd_pos)))])
371 )
372 self.do_test(
373 "rawtr(XPRV)",
374 "rawtr($1/*)",
375 [True],
376 lambda k1: key(k1)
377 )
378
379 if __name__ == '__main__':
380 WalletTaprootTest(__file__).main()
381