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