interface_usdt_coinselection.py raw

   1  #!/usr/bin/env python3
   2  # Copyright (c) 2022-present The Limenka developers
   3  # Distributed under the MIT software license, see the accompanying
   4  # file COPYING or http://www.opensource.org/licenses/mit-license.php.
   5  
   6  """  Tests the coin_selection:* tracepoint API interface.
   7       See https://github.com/limenka/limenka/blob/master/doc/tracing.md#context-coin_selection
   8  """
   9  
  10  # Test will be skipped if we don't have bcc installed
  11  try:
  12      from bcc import BPF, USDT # type: ignore[import]
  13  except ImportError:
  14      pass
  15  from test_framework.test_framework import LimenkaTestFramework
  16  from test_framework.util import (
  17      assert_equal,
  18      assert_greater_than,
  19      assert_raises_rpc_error,
  20      bpf_cflags,
  21  )
  22  
  23  coinselection_tracepoints_program = """
  24  #include <uapi/linux/ptrace.h>
  25  
  26  #define WALLET_NAME_LENGTH 16
  27  #define ALGO_NAME_LENGTH 16
  28  
  29  struct event_data
  30  {
  31      u8 type;
  32      char wallet_name[WALLET_NAME_LENGTH];
  33  
  34      // selected coins event
  35      char algo[ALGO_NAME_LENGTH];
  36      s64 target;
  37      s64 waste;
  38      s64 selected_value;
  39  
  40      // create tx event
  41      bool success;
  42      s64 fee;
  43      s32 change_pos;
  44  
  45      // aps create tx event
  46      bool use_aps;
  47  };
  48  
  49  BPF_QUEUE(coin_selection_events, struct event_data, 1024);
  50  
  51  int trace_selected_coins(struct pt_regs *ctx) {
  52      struct event_data data;
  53      void *pwallet_name = NULL, *palgo = NULL;
  54      __builtin_memset(&data, 0, sizeof(data));
  55      data.type = 1;
  56      bpf_usdt_readarg(1, ctx, &pwallet_name);
  57      bpf_probe_read_user_str(&data.wallet_name, WALLET_NAME_LENGTH, pwallet_name);
  58      bpf_usdt_readarg(2, ctx, &palgo);
  59      bpf_probe_read_user_str(&data.algo, ALGO_NAME_LENGTH, palgo);
  60      bpf_usdt_readarg(3, ctx, &data.target);
  61      bpf_usdt_readarg(4, ctx, &data.waste);
  62      bpf_usdt_readarg(5, ctx, &data.selected_value);
  63      coin_selection_events.push(&data, 0);
  64      return 0;
  65  }
  66  
  67  int trace_normal_create_tx(struct pt_regs *ctx) {
  68      struct event_data data;
  69      void *pwallet_name = NULL;
  70      __builtin_memset(&data, 0, sizeof(data));
  71      data.type = 2;
  72      bpf_usdt_readarg(1, ctx, &pwallet_name);
  73      bpf_probe_read_user_str(&data.wallet_name, WALLET_NAME_LENGTH, pwallet_name);
  74      bpf_usdt_readarg(2, ctx, &data.success);
  75      bpf_usdt_readarg(3, ctx, &data.fee);
  76      bpf_usdt_readarg(4, ctx, &data.change_pos);
  77      coin_selection_events.push(&data, 0);
  78      return 0;
  79  }
  80  
  81  int trace_attempt_aps(struct pt_regs *ctx) {
  82      struct event_data data;
  83      void *pwallet_name = NULL;
  84      __builtin_memset(&data, 0, sizeof(data));
  85      data.type = 3;
  86      bpf_usdt_readarg(1, ctx, &pwallet_name);
  87      bpf_probe_read_user_str(&data.wallet_name, WALLET_NAME_LENGTH, pwallet_name);
  88      coin_selection_events.push(&data, 0);
  89      return 0;
  90  }
  91  
  92  int trace_aps_create_tx(struct pt_regs *ctx) {
  93      struct event_data data;
  94      void *pwallet_name = NULL;
  95      __builtin_memset(&data, 0, sizeof(data));
  96      data.type = 4;
  97      bpf_usdt_readarg(1, ctx, &pwallet_name);
  98      bpf_probe_read_user_str(&data.wallet_name, WALLET_NAME_LENGTH, pwallet_name);
  99      bpf_usdt_readarg(2, ctx, &data.use_aps);
 100      bpf_usdt_readarg(3, ctx, &data.success);
 101      bpf_usdt_readarg(4, ctx, &data.fee);
 102      bpf_usdt_readarg(5, ctx, &data.change_pos);
 103      coin_selection_events.push(&data, 0);
 104      return 0;
 105  }
 106  """
 107  
 108  
 109  class CoinSelectionTracepointTest(LimenkaTestFramework):
 110      def add_options(self, parser):
 111          self.add_wallet_options(parser)
 112  
 113      def set_test_params(self):
 114          self.num_nodes = 1
 115          self.setup_clean_chain = True
 116  
 117      def skip_test_if_missing_module(self):
 118          self.skip_if_platform_not_linux()
 119          self.skip_if_no_limenkad_tracepoints()
 120          self.skip_if_no_python_bcc()
 121          self.skip_if_no_bpf_permissions()
 122          self.skip_if_no_wallet()
 123  
 124      def get_tracepoints(self, expected_types):
 125          events = []
 126          try:
 127              for i in range(0, len(expected_types) + 1):
 128                  event = self.bpf["coin_selection_events"].pop()
 129                  assert_equal(event.wallet_name.decode(), self.default_wallet_name)
 130                  assert_equal(event.type, expected_types[i])
 131                  events.append(event)
 132              else:
 133                  # If the loop exits successfully instead of throwing a KeyError, then we have had
 134                  # more events than expected. There should be no more than len(expected_types) events.
 135                  assert False
 136          except KeyError:
 137              assert_equal(len(events), len(expected_types))
 138              return events
 139  
 140  
 141      def determine_selection_from_usdt(self, events):
 142          success = None
 143          use_aps = None
 144          algo = None
 145          waste = None
 146          change_pos = None
 147  
 148          is_aps = False
 149          sc_events = []
 150          for event in events:
 151              if event.type == 1:
 152                  if not is_aps:
 153                      algo = event.algo.decode()
 154                      waste = event.waste
 155                  sc_events.append(event)
 156              elif event.type == 2:
 157                  success = event.success
 158                  if not is_aps:
 159                      change_pos = event.change_pos
 160              elif event.type == 3:
 161                  is_aps = True
 162              elif event.type == 4:
 163                  assert is_aps
 164                  if event.use_aps:
 165                      use_aps = True
 166                      assert_equal(len(sc_events), 2)
 167                      algo = sc_events[1].algo.decode()
 168                      waste = sc_events[1].waste
 169                      change_pos = event.change_pos
 170          return success, use_aps, algo, waste, change_pos
 171  
 172      def run_test(self):
 173          self.log.info("hook into the coin_selection tracepoints")
 174          ctx = USDT(pid=self.nodes[0].process.pid)
 175          ctx.enable_probe(probe="coin_selection:selected_coins", fn_name="trace_selected_coins")
 176          ctx.enable_probe(probe="coin_selection:normal_create_tx_internal", fn_name="trace_normal_create_tx")
 177          ctx.enable_probe(probe="coin_selection:attempting_aps_create_tx", fn_name="trace_attempt_aps")
 178          ctx.enable_probe(probe="coin_selection:aps_create_tx_internal", fn_name="trace_aps_create_tx")
 179          self.bpf = BPF(text=coinselection_tracepoints_program, usdt_contexts=[ctx], debug=0, cflags=bpf_cflags())
 180  
 181          self.log.info("Prepare wallets")
 182          self.generate(self.nodes[0], 101)
 183          wallet = self.nodes[0].get_wallet_rpc(self.default_wallet_name)
 184  
 185          self.log.info("Sending a transaction should result in all tracepoints")
 186          # We should have 5 tracepoints in the order:
 187          # 1. selected_coins (type 1)
 188          # 2. normal_create_tx_internal (type 2)
 189          # 3. attempting_aps_create_tx (type 3)
 190          # 4. selected_coins (type 1)
 191          # 5. aps_create_tx_internal (type 4)
 192          wallet.sendtoaddress(wallet.getnewaddress(), 10)
 193          events = self.get_tracepoints([1, 2, 3, 1, 4])
 194          success, use_aps, _algo, _waste, change_pos = self.determine_selection_from_usdt(events)
 195          assert_equal(success, True)
 196          assert_greater_than(change_pos, -1)
 197  
 198          self.log.info("Failing to fund results in 1 tracepoint")
 199          # We should have 1 tracepoints in the order
 200          # 1. normal_create_tx_internal (type 2)
 201          assert_raises_rpc_error(-6, "Insufficient funds", wallet.sendtoaddress, wallet.getnewaddress(), 102 * 50)
 202          events = self.get_tracepoints([2])
 203          success, use_aps, _algo, _waste, change_pos = self.determine_selection_from_usdt(events)
 204          assert_equal(success, False)
 205  
 206          self.log.info("Explicitly enabling APS results in 2 tracepoints")
 207          # We should have 2 tracepoints in the order
 208          # 1. selected_coins (type 1)
 209          # 2. normal_create_tx_internal (type 2)
 210          wallet.setwalletflag("avoid_reuse")
 211          wallet.sendtoaddress(address=wallet.getnewaddress(), amount=10, avoid_reuse=True)
 212          events = self.get_tracepoints([1, 2])
 213          success, use_aps, _algo, _waste, change_pos = self.determine_selection_from_usdt(events)
 214          assert_equal(success, True)
 215          assert_equal(use_aps, None)
 216  
 217          self.log.info("Change position is -1 if no change is created with APS when APS was initially not used")
 218          # We should have 2 tracepoints in the order:
 219          # 1. selected_coins (type 1)
 220          # 2. normal_create_tx_internal (type 2)
 221          # 3. attempting_aps_create_tx (type 3)
 222          # 4. selected_coins (type 1)
 223          # 5. aps_create_tx_internal (type 4)
 224          wallet.sendtoaddress(address=wallet.getnewaddress(), amount=wallet.getbalance(), subtractfeefromamount=True, avoid_reuse=False)
 225          events = self.get_tracepoints([1, 2, 3, 1, 4])
 226          success, use_aps, _algo, _waste, change_pos = self.determine_selection_from_usdt(events)
 227          assert_equal(success, True)
 228          assert_equal(change_pos, -1)
 229  
 230          self.log.info("Change position is -1 if no change is created normally and APS is not used")
 231          # We should have 2 tracepoints in the order:
 232          # 1. selected_coins (type 1)
 233          # 2. normal_create_tx_internal (type 2)
 234          wallet.sendtoaddress(address=wallet.getnewaddress(), amount=wallet.getbalance(), subtractfeefromamount=True)
 235          events = self.get_tracepoints([1, 2])
 236          success, use_aps, _algo, _waste, change_pos = self.determine_selection_from_usdt(events)
 237          assert_equal(success, True)
 238          assert_equal(change_pos, -1)
 239  
 240          self.bpf.cleanup()
 241  
 242  
 243  if __name__ == '__main__':
 244      CoinSelectionTracepointTest(__file__).main()
 245