feat: add a warning for high tx fee

This commit is contained in:
Chaitanya-Keyal
2026-09-20 16:03:45 +05:30
committed by okaybro
parent b225ae77e9
commit cd0b30db2f
4 changed files with 184 additions and 2 deletions
+48
View File
@@ -156,6 +156,10 @@ class PSBTParser():
# just stops getting cache hits once the cache is full.
MAX_CACHED_DERIVATIONS = 1000
# Warn when the fee exceeds this percentage of what is being sent (outputs other than
# change). TODO: Possibly make this configurable via settings.
HIGH_FEES_WARNING_THRESHOLD = 25
def __init__(self, p: PSBT, seed: Seed, network: str = SettingsConstants.MAINNET):
self.psbt: PSBT = p
@@ -1182,3 +1186,47 @@ class PSBTParser():
for out in self.psbt.outputs:
_fill_scope(out)
def get_total_output_value(self, include_change: bool = False):
"""
Returns the sum of all outputs (fee not included).
`change_data` holds every output that comes back to this seed, which is two
different things: change, and self-transfers to one of our own receive
addresses. With `include_change=False` only the change is subtracted;
self-transfers stay in the total, since the user chose to send funds there
just as they did for any external recipient. The two are told apart the way
the views do it, by `is_change_branch` on the derivation path the parse
proved this seed owns.
Used to decide whether the fee is high relative to what is actually being
sent, and whether to warn.
"""
total = sum(out.value for out in self.psbt.tx.vout)
if include_change:
return total
# Subtract the change; keep self-transfers, they count as recipients.
true_change = sum(
entry["amount"]
for entry in self.change_data
if PSBTParser.is_change_branch(entry["verified_derivation_path"])
)
return total - true_change
def has_high_fee(self):
"""
Returns True if the fee is high.
i.e. fee amount > <HIGH_FEES_WARNING_THRESHOLD>% of total outputs excluding change
"""
total_output_value_excluding_change = self.get_total_output_value()
# If there are no outputs other than change, then it can't be a high fee
if total_output_value_excluding_change <= 0:
return False
else:
return self.fee_amount > ((self.HIGH_FEES_WARNING_THRESHOLD / 100) * total_output_value_excluding_change)
+38
View File
@@ -163,6 +163,8 @@ class PSBTOverviewView(View):
else:
num_self_transfer_outputs += 1
is_high_fee_tx = psbt_parser.has_high_fee()
# Run the overview screen
selected_menu_num = self.run_screen(
PSBTOverviewScreen,
@@ -185,6 +187,9 @@ class PSBTOverviewView(View):
if psbt_parser.policy == None:
return Destination(PSBTUnsupportedScriptTypeWarningView)
elif is_high_fee_tx:
return Destination(PSBTHighFeeWarningView, view_args={"warning_threshold_percent": psbt_parser.HIGH_FEES_WARNING_THRESHOLD})
elif psbt_parser.change_amount == 0:
return Destination(PSBTNoChangeWarningView)
@@ -235,6 +240,39 @@ class PSBTNoChangeWarningView(View):
class PSBTHighFeeWarningView(View):
def __init__(self, warning_threshold_percent: int):
super().__init__()
self.warning_threshold_percent = warning_threshold_percent
def run(self):
selected_menu_num = self.run_screen(
DireWarningScreen,
status_headline=_("High Fee!"),
# TRANSLATOR_NOTE: Variable is the percentage of the total output value (excluding change) that the fee exceeds. (e.g. "This transaction has a fee higher than 25% of the total output value (excluding change).")
text=_("This transaction has a fee higher than {}% of the total output value (excluding change).").format(self.warning_threshold_percent),
button_data=[ButtonOption("Continue")],
)
if selected_menu_num == RET_CODE__BACK_BUTTON:
return Destination(BackStackView)
# PSBT may have high fee + no change
if self.controller.psbt_parser.change_amount == 0:
return Destination(
PSBTNoChangeWarningView,
skip_current_view=True, # Prevent going BACK to WarningViews
)
else:
return Destination(
PSBTMathView,
skip_current_view=True, # Prevent going BACK to WarningViews
)
class PSBTMathView(View):
"""
Follows the Overview pictogram. Shows:
+1
View File
@@ -439,6 +439,7 @@ def generate_screenshots(locale):
ScreenshotConfig(psbt_views.PSBTOverviewView, mock_context_manager=mock_multisig_psbt_loaded),
ScreenshotConfig(psbt_views.PSBTUnsupportedScriptTypeWarningView),
ScreenshotConfig(psbt_views.PSBTNoChangeWarningView),
ScreenshotConfig(psbt_views.PSBTHighFeeWarningView, dict(warning_threshold_percent=25)),
ScreenshotConfig(psbt_views.PSBTMathView, mock_context_manager=mock_multisig_psbt_loaded),
ScreenshotConfig(psbt_views.PSBTAddressDetailsView, dict(address_num=0), mock_context_manager=mock_multisig_psbt_loaded),
+95
View File
@@ -3,6 +3,7 @@ import random
from binascii import a2b_base64
from copy import deepcopy
from types import SimpleNamespace
from unittest.mock import patch
from embit import bip32, script
from embit.ec import PublicKey
@@ -60,6 +61,10 @@ class TestPSBTParser:
assert psbt_parser.fee_amount == fee_amount
assert psbt_parser.input_amount == psbt_parser.spend_amount + psbt_parser.change_amount + psbt_parser.fee_amount
# No self-transfer here, so all change is true change
assert psbt_parser.get_total_output_value() == psbt_parser.spend_amount
assert psbt_parser.get_total_output_value(include_change=True) == psbt_parser.spend_amount + psbt_parser.change_amount
# Internally cycle the input(s) back to sender via the `self_transfer_data`
psbt.outputs.clear()
psbt.outputs.append(create_output(self_transfer_data, input_amount - fee_amount))
@@ -75,6 +80,11 @@ class TestPSBTParser:
assert psbt_parser.fee_amount == fee_amount
assert psbt_parser.input_amount == psbt_parser.spend_amount + psbt_parser.change_amount + psbt_parser.fee_amount
# Only self-transfer, and no "true" change, so `change_amount` is included in total output
# Both calls should return the same value
assert psbt_parser.get_total_output_value() == psbt_parser.spend_amount + psbt_parser.change_amount
assert psbt_parser.get_total_output_value(include_change=True) == psbt_parser.spend_amount + psbt_parser.change_amount
# Now do full spends with no change
fee_amount = random.randint(5_000, 100_000)
recipient_amount = input_amount - fee_amount
@@ -94,6 +104,10 @@ class TestPSBTParser:
assert psbt_parser.fee_amount == fee_amount
assert psbt_parser.input_amount == psbt_parser.spend_amount + psbt_parser.change_amount + psbt_parser.fee_amount
# No self-transfer here, so all change is true change
assert psbt_parser.get_total_output_value() == psbt_parser.spend_amount
assert psbt_parser.get_total_output_value(include_change=True) == psbt_parser.spend_amount + psbt_parser.change_amount
# Now try a single mega psbt with ALL the outputs at once
psbt.outputs.clear()
change_amount = input_amount - fee_amount
@@ -116,6 +130,10 @@ class TestPSBTParser:
assert psbt_parser.fee_amount == fee_amount
assert psbt_parser.input_amount == psbt_parser.spend_amount + psbt_parser.change_amount + psbt_parser.fee_amount
# No self-transfer here, so all change is true change
assert psbt_parser.get_total_output_value() == psbt_parser.spend_amount
assert psbt_parser.get_total_output_value(include_change=True) == psbt_parser.spend_amount + psbt_parser.change_amount
def test_singlesig_native_segwit(self):
self.run_basic_test(PSBTTestData.SINGLE_SIG_NATIVE_SEGWIT_1_INPUT, PSBTTestData.SINGLE_SIG_NATIVE_SEGWIT_CHANGE, PSBTTestData.SINGLE_SIG_NATIVE_SEGWIT_SELF_TRANSFER)
@@ -348,6 +366,83 @@ class TestPSBTParser:
"""
assert PSBTParser.is_change_branch(bip32.parse_path("m/84h/1h/0h/1/0")) is True
assert PSBTParser.is_change_branch(bip32.parse_path("m/84h/1h/0h/0/0")) is False
@pytest.mark.parametrize("vout_values, change_data, expected", [
# single destination + single change
([100, 200], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 200}], 100),
# multiple destinations + single change
([50, 75, 25], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 25}], 50 + 75),
# no change outputs at all
([10, 20, 30], [], 10 + 20 + 30),
# only change outputs
([123], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 123}], 0),
# mix of true change and self-transfer
(
[100, 200, 300],
[
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/0/0"), "amount": 200}, # self-transfer
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 300}, # true change
],
100 + 200 # Only subtract true change (300)
),
])
def test_get_total_output_value(self, vout_values, change_data, expected):
"""
get_total_output_value() should return
sum(vout_values) - sum(true_change),
where true change is determined by derivation path having chain index 1.
"""
# Build a dummy parser without running .parse()
parser = PSBTParser.__new__(PSBTParser)
# Stub out parser.psbt.tx.vout as list of objects with a .value attribute
parser.psbt = SimpleNamespace(
tx=SimpleNamespace(
vout=[SimpleNamespace(value=v) for v in vout_values]
)
)
parser.change_data = change_data
assert parser.get_total_output_value() == expected
@pytest.mark.parametrize("vin, vout_values, change_data, expected", [
# fee=30 (24%) -> NOT high
(180, [50, 75, 25], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 25}], False),
# fee=40 (32%) -> HIGH
(190, [50, 75, 25], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 25}], True),
# fee=15 (exactly 25%) -> NOT high
(75, [10, 20, 30], [], False),
# only change outputs: excluding change=0 -> never high by definition
(130, [123], [{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/5"), "amount": 123}], False),
# fee=60 (20%) -> NOT high
(660, [100, 200, 300], [
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/0/0"), "amount": 200}, # self-transfer
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 300}, # true change
], False),
# same mix but high fee: fee=100 (33%) -> HIGH
(700, [100, 200, 300], [
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/0/0"), "amount": 200},
{"verified_derivation_path": bip32.parse_path("m/84h/0h/0h/1/0"), "amount": 300},
], True),
])
def test_has_high_fee(self, vin, vout_values, change_data, expected):
"""
Should correctly identify if a PSBT has a high fee.
"""
# Build a dummy parser without running .parse()
parser = PSBTParser.__new__(PSBTParser)
parser.HIGH_FEES_WARNING_THRESHOLD = 25
# Stub out parser.psbt.tx.vout as list of objects with a .value attribute
parser.psbt = SimpleNamespace(
tx=SimpleNamespace(
vout=[SimpleNamespace(value=v) for v in vout_values]
)
)
parser.change_data = change_data
parser.fee_amount = vin - sum(vout_values)
assert parser.has_high_fee() is expected