Merge pull request #1044 from kdmukai/psbt_multisig_output_claims

[security] Store multiple verified derivation path entries per input/output; reject all decoy keys
This commit is contained in:
Nick Klockenga
2026-09-25 16:05:28 -04:00
committed by GitHub
2 changed files with 221 additions and 143 deletions
+74 -53
View File
@@ -183,10 +183,11 @@ class PSBTParser():
self.is_high_fee: bool = False self.is_high_fee: bool = False
# Contains one entry per input in psbt.inputs and per output in psbt.outputs. Each # Contains one entry per input in psbt.inputs and per output in psbt.outputs. Each
# entry is either the derivation path the seed genuinely owns there, or it is set # entry lists every derivation path the seed genuinely owns there, in the order
# to `None`. # the psbt lists them. An input or output that does not claim any of our keys
self.verified_input_derivation_paths: List[List[int] | None] = [] # gets an empty list.
self.verified_output_derivation_paths: List[List[int] | None] = [] self.verified_input_derivation_paths: List[List[DerivationPath]] = []
self.verified_output_derivation_paths: List[List[DerivationPath]] = []
self.root = None self.root = None
@@ -255,8 +256,8 @@ class PSBTParser():
via: via:
- single-sig: Rebuild the output script from the seed and match it against - single-sig: Rebuild the output script from the seed and match it against
the committed scriptPubKey. the committed scriptPubKey.
- multisig: Match the seed's verified key against the pubkeys in the script - multisig: Match each of the seed's verified keys against the pubkeys in
the output commits to. the script the output commits to.
Every change_data entry after this point will carry a derivation path that Every change_data entry after this point will carry a derivation path that
our seed provably owns. our seed provably owns.
@@ -276,7 +277,7 @@ class PSBTParser():
levels, differing only in the address at the end. levels, differing only in the address at the end.
So every level derived during this parse is kept in a cache and reused. See So every level derived during this parse is kept in a cache and reused. See
_derive_with_cache. _derive_with_cache_via_indices.
Note that the cache is only useful within a single parse so it is not preserved. Note that the cache is only useful within a single parse so it is not preserved.
""" """
@@ -448,8 +449,8 @@ class PSBTParser():
# Rebuild the scriptPubKey from the key at the claimed derivation path # Rebuild the scriptPubKey from the key at the claimed derivation path
if len(out.bip32_derivations.values()) == 1: if len(out.bip32_derivations.values()) == 1:
singlesig_derivation_path = list(out.bip32_derivations.values())[0].derivation singlesig_derivation_path = list(out.bip32_derivations.values())[0]
seed_public_key = PSBTParser._derive_with_cache(self.root, singlesig_derivation_path, child_key_derivation_cache).get_public_key() seed_public_key = PSBTParser._derive_with_cache_via_derivation_path(self.root, singlesig_derivation_path, child_key_derivation_cache).get_public_key()
rebuilt_script_pubkey = PSBTParser._build_singlesig_script(self.policy["type"], seed_public_key) rebuilt_script_pubkey = PSBTParser._build_singlesig_script(self.policy["type"], seed_public_key)
else: else:
# There's nothing for us to verify against so this output will be # There's nothing for us to verify against so this output will be
@@ -474,9 +475,8 @@ class PSBTParser():
raise PSBTSurplusDerivationPathsError("Taproot output claims more than one internal key") raise PSBTSurplusDerivationPathsError("Taproot output claims more than one internal key")
if len(taproot_entries) == 1 and internal_key_claims == 1: if len(taproot_entries) == 1 and internal_key_claims == 1:
leaf_hashes, derivation = taproot_entries[0] leaf_hashes, singlesig_derivation_path = taproot_entries[0]
singlesig_derivation_path = derivation.derivation seed_public_key = PSBTParser._derive_with_cache_via_derivation_path(self.root, singlesig_derivation_path, child_key_derivation_cache).get_public_key()
seed_public_key = PSBTParser._derive_with_cache(self.root, singlesig_derivation_path, child_key_derivation_cache).get_public_key()
rebuilt_script_pubkey = PSBTParser._build_singlesig_script(self.policy["type"], seed_public_key) rebuilt_script_pubkey = PSBTParser._build_singlesig_script(self.policy["type"], seed_public_key)
else: else:
# This output has at least one derivation path entry for a key # This output has at least one derivation path entry for a key
@@ -495,19 +495,19 @@ class PSBTParser():
# which is also caught here. # which is also caught here.
raise RuntimeError(f"Unsupported policy type: {self.policy['type']}") raise RuntimeError(f"Unsupported policy type: {self.policy['type']}")
verified_derivation_path = self.verified_output_derivation_paths[i] verified_derivation_paths = self.verified_output_derivation_paths[i]
if rebuilt_script_pubkey.data == vout[i].script_pubkey.data: if rebuilt_script_pubkey.data == vout[i].script_pubkey.data:
# The scriptPubKey we created using our own seed matched what this # The scriptPubKey we created using our own seed matched what this
# output is actually committing to. # output is actually committing to.
if singlesig_derivation_path is not None: if singlesig_derivation_path is not None:
if verified_derivation_path is None: if verified_derivation_paths == []:
# The output pays this seed but the psbt claimed a different # The output pays this seed but the psbt claimed a different
# fingerprint here. We treat this deception as an attack. # fingerprint here. We treat this deception as an attack.
raise PSBTOutputOwnershipContradictionError(f"Output pays this seed at {bip32.path_to_str(singlesig_derivation_path)} but does not claim it there") raise PSBTOutputOwnershipContradictionError(f"Output pays this seed at {bip32.path_to_str(singlesig_derivation_path.derivation)} but does not claim it there")
if verified_derivation_path != singlesig_derivation_path: if verified_derivation_paths != [singlesig_derivation_path]:
# Shouldn't be able to reach here: the surplus check above # Shouldn't be able to reach here: the surplus check above
# allows only one entry, and the ownership scan refuses a # allows only one entry, and the ownership scan refuses a
# scope populating both derivation path maps, so the scan can # scope populating both derivation path maps, so the scan can
@@ -520,7 +520,7 @@ class PSBTParser():
is_presumed_change = True is_presumed_change = True
elif multisig_script is not None: elif multisig_script is not None:
if verified_derivation_path is None: if verified_derivation_paths == []:
# No entry claimed this seed's fingerprint, but we already # No entry claimed this seed's fingerprint, but we already
# have everything we need to see if our seed is actually in # have everything we need to see if our seed is actually in
# the output script. # the output script.
@@ -529,7 +529,7 @@ class PSBTParser():
# the coordinator says sits there. Both are its own # the coordinator says sits there. Both are its own
# claims, so we read only the path and derive the key # claims, so we read only the path and derive the key
# ourselves. # ourselves.
seed_public_key = PSBTParser._derive_with_cache(self.root, derivation_path_obj.derivation, child_key_derivation_cache).get_public_key() seed_public_key = PSBTParser._derive_with_cache_via_derivation_path(self.root, derivation_path_obj, child_key_derivation_cache).get_public_key()
if PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key): if PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key):
# The output pays a multisig this seed is part # The output pays a multisig this seed is part
@@ -545,14 +545,17 @@ class PSBTParser():
else: else:
# This output claimed that our seed is part of the receiving # This output claimed that our seed is part of the receiving
# multisig, at a specific path. So now we verify that the key # multisig, at one or more specific paths. So now we verify
# at that path is in the committed script. # that the key at every claimed path is in the committed
seed_public_key = PSBTParser._derive_with_cache(self.root, verified_derivation_path, child_key_derivation_cache).get_public_key() # script.
if not PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key): for verified_derivation_path in verified_derivation_paths:
# The psbt said this output was coming back to our seed seed_public_key = PSBTParser._derive_with_cache_via_derivation_path(self.root, verified_derivation_path, child_key_derivation_cache).get_public_key()
# at that path, but the key there is not in the committed if not PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key):
# script. We treat this deception as an attack. # The psbt said this output was coming back to our
raise PSBTOutputOwnershipContradictionError(f"Output claims this seed at {bip32.path_to_str(verified_derivation_path)} but its committed script does not hold that key") # seed at that path, but the key there is not in the
# committed script. We treat this deception as an
# attack.
raise PSBTOutputOwnershipContradictionError(f"Output claims this seed at {bip32.path_to_str(verified_derivation_path.derivation)} but its committed script does not hold that key")
# The output should not describe more keys than are actually # The output should not describe more keys than are actually
# used in its script. We check for the more serious deceptions # used in its script. We check for the more serious deceptions
@@ -582,7 +585,7 @@ class PSBTParser():
if input_cosigners is not None and input_cosigners != output_cosigners: if input_cosigners is not None and input_cosigners != output_cosigners:
is_presumed_change = False is_presumed_change = False
elif verified_derivation_path is not None and self.policy["type"] != "p2tr": elif verified_derivation_paths != [] and self.policy["type"] != "p2tr":
# The psbt claims one of this seed's keys on this output, yet the # The psbt claims one of this seed's keys on this output, yet the
# output does NOT pay what that claim describes. We treat this # output does NOT pay what that claim describes. We treat this
# deception as an attack. # deception as an attack.
@@ -602,7 +605,7 @@ class PSBTParser():
# output verifiable change, one that does not is a contradiction to # output verifiable change, one that does not is a contradiction to
# refuse here, and an output supplying no tree stays exempt, since # refuse here, and an output supplying no tree stays exempt, since
# an omitted optional field is not a contradiction. # an omitted optional field is not a contradiction.
raise PSBTOutputOwnershipContradictionError(f"Output claims this seed at {bip32.path_to_str(verified_derivation_path)} but its committed script contradicts that") raise PSBTOutputOwnershipContradictionError(f"Output claims this seed at {bip32.path_to_str(verified_derivation_paths[0].derivation)} but its committed script contradicts that")
if vout[i].script_pubkey.data[0] == OPCODES.OP_RETURN: if vout[i].script_pubkey.data[0] == OPCODES.OP_RETURN:
# The data is written as: OP_RETURN + OP_PUSHDATA1 + len(payload) + payload # The data is written as: OP_RETURN + OP_PUSHDATA1 + len(payload) + payload
@@ -618,7 +621,7 @@ class PSBTParser():
"output_index": i, "output_index": i,
"address": addr, "address": addr,
"amount": vout[i].value, "amount": vout[i].value,
"verified_derivation_path": self.verified_output_derivation_paths[i], "verified_derivation_path": self.verified_output_derivation_paths[i][0].derivation,
}) })
self.change_amount += vout[i].value self.change_amount += vout[i].value
@@ -787,7 +790,7 @@ class PSBTParser():
@staticmethod @staticmethod
def _derive_with_cache(parent_key: bip32.HDKey, derivation_path: List[int], child_key_derivation_cache: dict | None = None) -> bip32.HDKey: def _derive_with_cache_via_indices(parent_key: bip32.HDKey, derivation_path: List[int], child_key_derivation_cache: dict | None = None) -> bip32.HDKey:
""" """
Derives the key that sits at the given derivation path below parent_key, reusing Derives the key that sits at the given derivation path below parent_key, reusing
any levels along the way that have already been derived during this parse. any levels along the way that have already been derived during this parse.
@@ -837,6 +840,19 @@ class PSBTParser():
return derived_key return derived_key
@staticmethod
def _derive_with_cache_via_derivation_path(parent_key: bip32.HDKey, derivation_path: DerivationPath, child_key_derivation_cache: dict | None = None) -> bip32.HDKey:
"""
_derive_with_cache_via_indices for a psbt entry: derives at the entry's full
derivation path below parent_key.
The DerivationPath.fingerprint is completely ignored; this function allows for
deriving a key even when it's known that the fingerprint doesn't match (e.g. to
catch a false claim).
"""
return PSBTParser._derive_with_cache_via_indices(parent_key, derivation_path.derivation, child_key_derivation_cache)
@staticmethod @staticmethod
def _get_cosigners(pubkeys, derivations, xpubs, child_key_derivation_cache: dict | None): def _get_cosigners(pubkeys, derivations, xpubs, child_key_derivation_cache: dict | None):
""" """
@@ -895,7 +911,7 @@ class PSBTParser():
if origin_der.derivation == der.derivation[:-2]: if origin_der.derivation == der.derivation[:-2]:
# Derive the child key that sits two indices below the xpub (i.e. at # Derive the child key that sits two indices below the xpub (i.e. at
# the full derivation path). # the full derivation path).
derived_key = PSBTParser._derive_with_cache(xpub, der.derivation[-2:], child_key_derivation_cache) derived_key = PSBTParser._derive_with_cache_via_indices(xpub, der.derivation[-2:], child_key_derivation_cache)
# Finally, compare that key with the target pubkey # Finally, compare that key with the target pubkey
if derived_key.key == pubkey: if derived_key.key == pubkey:
@@ -987,7 +1003,7 @@ class PSBTParser():
say anything. Ownership is established here and only here, by deriving the key say anything. Ownership is established here and only here, by deriving the key
again from the seed and comparing the actual key material. again from the seed and comparing the actual key material.
""" """
derived_public_key = PSBTParser._derive_with_cache(root, claimed_derivation_path, child_key_derivation_cache).get_public_key() derived_public_key = PSBTParser._derive_with_cache_via_indices(root, claimed_derivation_path, child_key_derivation_cache).get_public_key()
if is_taproot: if is_taproot:
# A psbt carries a taproot key as its bare 32-byte x coordinate, but embit # A psbt carries a taproot key as its bare 32-byte x coordinate, but embit
@@ -1008,12 +1024,13 @@ class PSBTParser():
@staticmethod @staticmethod
def _get_seed_derivation_path(scope: InputScope | OutputScope, root: bip32.HDKey, child_key_derivation_cache: dict) -> List[int] | None: def _get_seed_derivation_paths(scope: InputScope | OutputScope, root: bip32.HDKey, child_key_derivation_cache: dict) -> List[DerivationPath]:
""" """
Scans the derivation path(s) in the provided input or output scope to determine Scans the derivation path(s) in the provided input or output scope to determine
which, if any, are provably derived from the signing seed (for multisig a path is which, if any, are provably derived from the signing seed (for multisig a path is
provided per key; if the seed is part of the multisig, one of the n paths will provided per key; if the seed is part of the multisig, one of the n paths will
match). Returns the verified derivation path (as a list of ints) or None. match). Returns every verified DerivationPath entry, in the order the psbt lists
them. An input or output that does not claim any of our keys yields an empty list.
Every key in the scope that claims this seed's fingerprint is re-derived and Every key in the scope that claims this seed's fingerprint is re-derived and
checked. A false claim raises PSBT[Output|Input]OwnershipClaimError. checked. A false claim raises PSBT[Output|Input]OwnershipClaimError.
@@ -1027,22 +1044,26 @@ class PSBTParser():
Note that neither BIP-174 nor BIP-371 forbids the combination. And embit will Note that neither BIP-174 nor BIP-371 forbids the combination. And embit will
parse and even sign such a psbt. We disallow it by opinionated choice. parse and even sign such a psbt. We disallow it by opinionated choice.
One edge case: One edge case: A scope may carry more than one entry that verifies against this
* A multisig could use this seed in more than one cosigner slot, each seed.
at its own derivation path. The scope then carries several entries that all * Foolish as it may be, a multisig could honestly use this seed in two cosigner
verify against this seed; we return the first but still check the rest. slots, each at its own derivation path.
* More importantly: a malicious psbt could list a second claim of ours as a
decoy, at a path our seed really does derive but whose key the committed
script has no use for.
This function only checks and returns the DerivationPath entry for each key that
derives from our seed. What those entries mean for the psbt is determined
elsewhere.
The path itself is still whatever the psbt supplied: it can be any length or The path itself is still whatever the psbt supplied: it can be any length or
shape, since any path that derives from the seed will pass. Whether the path is shape, since any path that derives from the seed will pass. Whether the path is
one the user's wallet would ever look at is a separate question, answered one the user's wallet would ever look at is a separate question.
elsewhere.
""" """
seed_fingerprint = root.my_fingerprint seed_fingerprint = root.my_fingerprint
verified_derivation_path = None verified_derivation_paths = []
def _check_claim(public_key: PublicKey, derivation_path_obj: DerivationPath, is_taproot: bool): def _check_claim(public_key: PublicKey, derivation_path_obj: DerivationPath, is_taproot: bool):
nonlocal verified_derivation_path
if derivation_path_obj.fingerprint != seed_fingerprint: if derivation_path_obj.fingerprint != seed_fingerprint:
# Claims to belong to some other key. Nothing to prove or disprove here. # Claims to belong to some other key. Nothing to prove or disprove here.
return return
@@ -1051,9 +1072,7 @@ class PSBTParser():
error_class = (PSBTInputOwnershipClaimError if isinstance(scope, InputScope) else PSBTOutputOwnershipClaimError) error_class = (PSBTInputOwnershipClaimError if isinstance(scope, InputScope) else PSBTOutputOwnershipClaimError)
raise error_class(f"Key at {bip32.path_to_str(derivation_path_obj.derivation)} claims this seed's fingerprint but does not derive from it") raise error_class(f"Key at {bip32.path_to_str(derivation_path_obj.derivation)} claims this seed's fingerprint but does not derive from it")
# Store only the first verified path verified_derivation_paths.append(derivation_path_obj)
if verified_derivation_path is None:
verified_derivation_path = derivation_path_obj.derivation
# Note that both loops check EVERY claim # Note that both loops check EVERY claim
for public_key, derivation_path_obj in scope.bip32_derivations.items(): for public_key, derivation_path_obj in scope.bip32_derivations.items():
@@ -1069,14 +1088,15 @@ class PSBTParser():
if scope.bip32_derivations and scope.taproot_bip32_derivations: if scope.bip32_derivations and scope.taproot_bip32_derivations:
raise PSBTMixedDerivationPathTypesError("Scope declares both ecdsa and taproot derivation paths") raise PSBTMixedDerivationPathTypesError("Scope declares both ecdsa and taproot derivation paths")
return verified_derivation_path return verified_derivation_paths
def _verify_claimed_derivation_paths(self, child_key_derivation_cache: dict): def _verify_claimed_derivation_paths(self, child_key_derivation_cache: dict):
""" """
Verifies every derivation path entry that claims this seed's fingerprint. The Verifies every derivation path entry that claims this seed's fingerprint. The
result, stored in verified_[input|output]_derivation_paths, is either the verified result, stored in verified_[input|output]_derivation_paths, is the list of
derivation path or None (no entry claimed this seed) for each input/output scope. verified DerivationPath entries for each input/output scope (empty where no entry
claimed this seed).
The coordinator-supplied fingerprints cannot be trusted as-is. We must derive and The coordinator-supplied fingerprints cannot be trusted as-is. We must derive and
verify the ownership of each one that claims to belong to this seed. verify the ownership of each one that claims to belong to this seed.
@@ -1088,12 +1108,12 @@ class PSBTParser():
Raises PSBT[Output|Input]OwnershipClaimError on the first false claim detected. Raises PSBT[Output|Input]OwnershipClaimError on the first false claim detected.
""" """
self.verified_output_derivation_paths = [ self.verified_output_derivation_paths = [
PSBTParser._get_seed_derivation_path(out, self.root, child_key_derivation_cache) PSBTParser._get_seed_derivation_paths(out, self.root, child_key_derivation_cache)
for out in self.psbt.outputs for out in self.psbt.outputs
] ]
self.verified_input_derivation_paths = [ self.verified_input_derivation_paths = [
PSBTParser._get_seed_derivation_path(inp, self.root, child_key_derivation_cache) PSBTParser._get_seed_derivation_paths(inp, self.root, child_key_derivation_cache)
for inp in self.psbt.inputs for inp in self.psbt.inputs
] ]
@@ -1119,8 +1139,9 @@ class PSBTParser():
# proved the seed derives it (single-sig: one such key; multisig: one per # proved the seed derives it (single-sig: one such key; multisig: one per
# cosigner, ours among them). One verified input path is enough for the psbt to # cosigner, ours among them). One verified input path is enough for the psbt to
# be signable. # be signable.
if any(path is not None for path in self.verified_input_derivation_paths): for verified_derivation_paths in self.verified_input_derivation_paths:
return if len(verified_derivation_paths) > 0:
return
# There's nothing for this seed to sign # There's nothing for this seed to sign
raise PSBTSeedCannotSignError() raise PSBTSeedCannotSignError()
+147 -90
View File
@@ -288,7 +288,7 @@ class TestPSBTParser:
parser = PSBTParser(p=psbt, seed=PSBTTestData.seed, network=SettingsConstants.REGTEST) parser = PSBTParser(p=psbt, seed=PSBTTestData.seed, network=SettingsConstants.REGTEST)
(_, filled_derivation) = parser.psbt.inputs[0].taproot_bip32_derivations[x_only_public_key] (_, filled_derivation) = parser.psbt.inputs[0].taproot_bip32_derivations[x_only_public_key]
assert filled_derivation.fingerprint == parser.root.my_fingerprint assert filled_derivation.fingerprint == parser.root.my_fingerprint
assert parser.verified_input_derivation_paths == [bip32.parse_path(odd_parity_derivation_path)] assert parser.verified_input_derivation_paths == [[filled_derivation]]
def test_trim_and_sig_count(self): def test_trim_and_sig_count(self):
@@ -685,13 +685,14 @@ class TestPSBTParserOptimizations:
def cache_size_recorder(self, cache_sizes: list): def cache_size_recorder(self, cache_sizes: list):
""" """
Returns a stand-in for _derive_with_cache that derives exactly as the real one Returns a stand-in for _derive_with_cache_via_indices that derives exactly as the
does, but appends the cache's size to cache_sizes on the way out of every call. real one does, but appends the cache's size to cache_sizes on the way out of every
call.
The cache is a local inside parse(), so intercepting the calls it gets handed to The cache is a local inside parse(), so intercepting the calls it gets handed to
is the only way to see how large it grew. is the only way to see how large it grew.
""" """
real_derive_with_cache = PSBTParser._derive_with_cache real_derive_with_cache = PSBTParser._derive_with_cache_via_indices
def recorded(parent_key, derivation_path, cache=None): def recorded(parent_key, derivation_path, cache=None):
derived_key = real_derive_with_cache(parent_key, derivation_path, cache) derived_key = real_derive_with_cache(parent_key, derivation_path, cache)
@@ -787,8 +788,8 @@ class TestPSBTParserOptimizations:
# The cache here isn't providing any speedup (there are no derivations in the # The cache here isn't providing any speedup (there are no derivations in the
# cache to take advantage of), but we're just testing that the cache doesn't # cache to take advantage of), but we're just testing that the cache doesn't
# confuse/combine the two cosigners' derivation data. # confuse/combine the two cosigners' derivation data.
from_a = PSBTParser._derive_with_cache(cosigner_a_xpub, receive_index_5, cache) from_a = PSBTParser._derive_with_cache_via_indices(cosigner_a_xpub, receive_index_5, cache)
from_b = PSBTParser._derive_with_cache(cosigner_b_xpub, receive_index_5, cache) from_b = PSBTParser._derive_with_cache_via_indices(cosigner_b_xpub, receive_index_5, cache)
# Two levels should have been added for each cosigner # Two levels should have been added for each cosigner
assert len(cache) == 4 assert len(cache) == 4
@@ -802,6 +803,26 @@ class TestPSBTParserOptimizations:
assert from_b.key.sec() == cosigner_b_xpub.derive(receive_index_5).key.sec() assert from_b.key.sec() == cosigner_b_xpub.derive(receive_index_5).key.sec()
def test_derive_with_cache_via_derivation_path_ignores_the_entry_fingerprint(self):
"""
The helper derives the key at the entry's path, whatever fingerprint the entry
lists.
"""
root = self._root()
derivation_path = bip32.parse_path("m/84h/1h/0h/1/7")
expected_key = root.derive(derivation_path).key.sec()
# Our own fingerprint, the all-zero placeholder, and a stranger's
for fingerprint in [root.my_fingerprint, b"\x00\x00\x00\x00", bytes.fromhex("deadbeef")]:
entry = DerivationPath(fingerprint, derivation_path)
derived = PSBTParser._derive_with_cache_via_derivation_path(root, entry, {})
assert derived.key.sec() == expected_key
# Same answer with the cache disabled
derived = PSBTParser._derive_with_cache_via_derivation_path(root, entry, None)
assert derived.key.sec() == expected_key
def test_get_cosigners_identical_with_and_without_cache(self): def test_get_cosigners_identical_with_and_without_cache(self):
""" """
The cache is transparent to callers: _get_cosigners returns the same cosigner The cache is transparent to callers: _get_cosigners returns the same cosigner
@@ -833,9 +854,9 @@ class TestPSBTParserOptimizations:
def test_cache_does_not_change_parse_output(self): def test_cache_does_not_change_parse_output(self):
""" """
The whole point of the cache is that it changes nothing at all. Parse the same The whole point of the cache is that it changes nothing at all. Parse the same
psbt twice — once normally, once with the cache discarded so that every derivation psbt twice: once normally, once with the cache discarded so that every derivation
falls through to embit's own HDKey.derive() — and require identical parser state falls through to embit's own HDKey.derive(). The two runs must yield the identical
and identical resulting psbt bytes. parser state and resulting psbt bytes.
Single-sig and multisig each get a run because they reach the cache from different Single-sig and multisig each get a run because they reach the cache from different
starting points: single-sig traverses down from our own root, multisig down from starting points: single-sig traverses down from our own root, multisig down from
@@ -854,7 +875,7 @@ class TestPSBTParserOptimizations:
def assert_cache_makes_no_difference(input_base64: str, change_hex: str): def assert_cache_makes_no_difference(input_base64: str, change_hex: str):
# Store the real function before the patches below replace it. Each replacement # Store the real function before the patches below replace it. Each replacement
# still needs access to the real function to do the actual deriving. # still needs access to the real function to do the actual deriving.
real_derive_with_cache = PSBTParser._derive_with_cache real_derive_with_cache = PSBTParser._derive_with_cache_via_indices
# This version of the replacement will derive exactly as the real cache-backed # This version of the replacement will derive exactly as the real cache-backed
# function does, but will also record the cache it was handed on each call. # function does, but will also record the cache it was handed on each call.
@@ -863,7 +884,7 @@ class TestPSBTParserOptimizations:
caches_received.append(cache) caches_received.append(cache)
return real_derive_with_cache(parent_key, derivation_path, cache) return real_derive_with_cache(parent_key, derivation_path, cache)
with patch.object(PSBTParser, "_derive_with_cache", staticmethod(recording_derive_with_cache)): with patch.object(PSBTParser, "_derive_with_cache_via_indices", staticmethod(recording_derive_with_cache)):
with_cache = PSBTParser( with_cache = PSBTParser(
build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST) build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST)
@@ -875,7 +896,7 @@ class TestPSBTParserOptimizations:
def cache_free_derive(parent_key, derivation_path, cache=None): def cache_free_derive(parent_key, derivation_path, cache=None):
return real_derive_with_cache(parent_key, derivation_path) return real_derive_with_cache(parent_key, derivation_path)
with patch.object(PSBTParser, "_derive_with_cache", staticmethod(cache_free_derive)): with patch.object(PSBTParser, "_derive_with_cache_via_indices", staticmethod(cache_free_derive)):
without_cache = PSBTParser( without_cache = PSBTParser(
build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST) build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST)
@@ -907,7 +928,7 @@ class TestPSBTParserOptimizations:
# Record how large the cache grew over the course of each parse # Record how large the cache grew over the course of each parse
unconstrained_sizes = [] unconstrained_sizes = []
with patch.object(PSBTParser, "_derive_with_cache", staticmethod(self.cache_size_recorder(unconstrained_sizes))): with patch.object(PSBTParser, "_derive_with_cache_via_indices", staticmethod(self.cache_size_recorder(unconstrained_sizes))):
multisig_unconstrained = PSBTParser(build_psbt(multisig_case), self.seed, network=SettingsConstants.REGTEST) multisig_unconstrained = PSBTParser(build_psbt(multisig_case), self.seed, network=SettingsConstants.REGTEST)
singlesig_unconstrained = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST) singlesig_unconstrained = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST)
@@ -916,7 +937,7 @@ class TestPSBTParserOptimizations:
cap = 3 cap = 3
capped_sizes = [] capped_sizes = []
with patch.object(PSBTParser, "MAX_CACHED_DERIVATIONS", cap): with patch.object(PSBTParser, "MAX_CACHED_DERIVATIONS", cap):
with patch.object(PSBTParser, "_derive_with_cache", staticmethod(self.cache_size_recorder(capped_sizes))): with patch.object(PSBTParser, "_derive_with_cache_via_indices", staticmethod(self.cache_size_recorder(capped_sizes))):
multisig_capped = PSBTParser(build_psbt(multisig_case), self.seed, network=SettingsConstants.REGTEST) multisig_capped = PSBTParser(build_psbt(multisig_case), self.seed, network=SettingsConstants.REGTEST)
singlesig_capped = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST) singlesig_capped = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST)
@@ -959,6 +980,9 @@ class PSBTParserOwnershipTestBase:
def _parse(self, psbt: PSBT) -> PSBTParser: def _parse(self, psbt: PSBT) -> PSBTParser:
# TODO: Rename this helper. "parse" does not convey that a new PSBTParser instance
# is being created and it creates confusion in tests that also call PSBT.parse()
# (embit's deserializer).
return PSBTParser(psbt, self.seed, network=SettingsConstants.REGTEST) return PSBTParser(psbt, self.seed, network=SettingsConstants.REGTEST)
@@ -1043,20 +1067,20 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
assert len(psbt_parser.verified_output_derivation_paths) == len(psbt.outputs) assert len(psbt_parser.verified_output_derivation_paths) == len(psbt.outputs)
# Every recorded path is one the seed really does derive the scope's key at # Every recorded path is one the seed really does derive the scope's key at
for scopes, verified_derivation_paths in [ for scopes, verified_derivation_paths_per_scope in [
(psbt.inputs, psbt_parser.verified_input_derivation_paths), (psbt.inputs, psbt_parser.verified_input_derivation_paths),
(psbt.outputs, psbt_parser.verified_output_derivation_paths), (psbt.outputs, psbt_parser.verified_output_derivation_paths),
]: ]:
for scope, verified_derivation_path in zip(scopes, verified_derivation_paths): for scope, verified_derivation_paths in zip(scopes, verified_derivation_paths_per_scope):
assert verified_derivation_path is not None assert len(verified_derivation_paths) == 1
public_key = list(scope.bip32_derivations.keys())[0] public_key = list(scope.bip32_derivations.keys())[0]
assert PSBTParser.seed_owns_pubkey(psbt_parser.root, verified_derivation_path, public_key, child_key_derivation_cache=None) is True assert PSBTParser.seed_owns_pubkey(psbt_parser.root, verified_derivation_paths[0].derivation, public_key, child_key_derivation_cache=None) is True
def test__parse__verified_derivation_paths_none_for_not_owned_output(self): def test__parse__verified_derivation_paths_empty_for_not_owned_output(self):
""" """
An output paying someone else is not a failure; the seed simply owns nothing An output paying someone else is not a failure; the seed simply owns nothing
there so the matching verified_output_derivation_paths should be None. there so the matching verified_output_derivation_paths entry should be empty.
""" """
psbt = self._psbt_with_change() psbt = self._psbt_with_change()
@@ -1065,11 +1089,11 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
assert psbt_parser.verified_output_derivation_paths[0] is None assert psbt_parser.verified_output_derivation_paths[0] == []
assert psbt_parser.verified_output_derivation_paths[1] is not None assert psbt_parser.verified_output_derivation_paths[1] != []
def test__parse__verified_derivation_paths_none_for_not_owned_input(self): def test__parse__verified_derivation_paths_empty_for_not_owned_input(self):
""" """
A collaborative spend also includes an input belonging to another party, in two A collaborative spend also includes an input belonging to another party, in two
shapes: a payjoin counterparty's input arrives finalized with no derivation info shapes: a payjoin counterparty's input arrives finalized with no derivation info
@@ -1087,15 +1111,15 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
# The payjoin shape # The payjoin shape
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
assert psbt_parser.verified_input_derivation_paths[0] is not None assert psbt_parser.verified_input_derivation_paths[0] != []
assert psbt_parser.verified_input_derivation_paths[1] is None assert psbt_parser.verified_input_derivation_paths[1] == []
# The coordinated shape: the derivation entry is truthful, naming the other # The coordinated shape: the derivation entry is truthful, naming the other
# party's fingerprint and a key that party really controls. # party's fingerprint and a key that party really controls.
claim_seed_owns_key(foreign_input, "m/84h/1h/0h/0/0", foreign_public_key(), seed=PSBTTestData.recipient_seed) claim_seed_owns_key(foreign_input, "m/84h/1h/0h/0/0", foreign_public_key(), seed=PSBTTestData.recipient_seed)
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
assert psbt_parser.verified_input_derivation_paths[0] is not None assert psbt_parser.verified_input_derivation_paths[0] != []
assert psbt_parser.verified_input_derivation_paths[1] is None assert psbt_parser.verified_input_derivation_paths[1] == []
def test__parse__rejects_a_forged_claim_on_an_input(self): def test__parse__rejects_a_forged_claim_on_an_input(self):
@@ -1228,7 +1252,7 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
# The seed still owns its own key in that input, via the scope's genuine # The seed still owns its own key in that input, via the scope's genuine
# derivation. # derivation.
assert psbt_parser.verified_input_derivation_paths[0] is not None assert psbt_parser.verified_input_derivation_paths[0] != []
def test_genuine_fingerprint_collision_is_rejected_like_a_forgery(self): def test_genuine_fingerprint_collision_is_rejected_like_a_forgery(self):
@@ -1291,8 +1315,9 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
# Sanity check: the scan really did run over all ten inputs and the change output # Sanity check: the scan really did run over all ten inputs and the change output
assert len(psbt_parser.verified_input_derivation_paths) == 10 assert len(psbt_parser.verified_input_derivation_paths) == 10
assert all(path is not None for path in psbt_parser.verified_input_derivation_paths) for verified_derivation_paths in psbt_parser.verified_input_derivation_paths:
assert psbt_parser.verified_output_derivation_paths[0] is not None assert verified_derivation_paths != []
assert psbt_parser.verified_output_derivation_paths[0] != []
# The inputs were cloned so they all use the same path with num_levels depth. The # The inputs were cloned so they all use the same path with num_levels depth. The
# change output differs only in its last two levels. Verify that each of these # change output differs only in its last two levels. Verify that each of these
@@ -1326,14 +1351,10 @@ class TestPSBTParserSeedOwnership(PSBTParserOwnershipTestBase):
""" """
psbt = self._psbt_with_change(PSBTTestData.MULTISIG_NATIVE_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NATIVE_SEGWIT_CHANGE) psbt = self._psbt_with_change(PSBTTestData.MULTISIG_NATIVE_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NATIVE_SEGWIT_CHANGE)
psbt_parser = PSBTParser(psbt, PSBTTestData.seed, network=SettingsConstants.REGTEST) # The fixture has one input; each cosigner's seed must verify on it
assert any(path is not None for path in psbt_parser.verified_input_derivation_paths) for seed in [PSBTTestData.seed, PSBTTestData.multisig_key_2, PSBTTestData.multisig_key_3]:
psbt_parser = PSBTParser(psbt, seed, network=SettingsConstants.REGTEST)
psbt_parser = PSBTParser(psbt, PSBTTestData.multisig_key_2, network=SettingsConstants.REGTEST) assert psbt_parser.verified_input_derivation_paths[0] != []
assert any(path is not None for path in psbt_parser.verified_input_derivation_paths)
psbt_parser = PSBTParser(psbt, PSBTTestData.multisig_key_3, network=SettingsConstants.REGTEST)
assert any(path is not None for path in psbt_parser.verified_input_derivation_paths)
def test_a_psbt_with_no_utxos_is_rejected_rather_than_crashing(self): def test_a_psbt_with_no_utxos_is_rejected_rather_than_crashing(self):
@@ -1539,7 +1560,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
# Trivial confirmation: none of the output's three derivation path entries claimed # Trivial confirmation: none of the output's three derivation path entries claimed
# to belong to this seed. # to belong to this seed.
assert psbt_parser.verified_output_derivation_paths[0] is None assert psbt_parser.verified_output_derivation_paths[0] == []
# The parser correctly categorized the output as an external spend # The parser correctly categorized the output as an external spend
assert psbt_parser.change_data == [] assert psbt_parser.change_data == []
@@ -1578,7 +1599,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
# With the derivation paths present, we verified that the output did name a key # With the derivation paths present, we verified that the output did name a key
# that this seed owns (which also enabled the parser to verify that our key was # that this seed owns (which also enabled the parser to verify that our key was
# indeed part of the script). # indeed part of the script).
assert psbt_parser.verified_output_derivation_paths[0] is not None assert psbt_parser.verified_output_derivation_paths[0] != []
# And the output was correctly categorized as change # And the output was correctly categorized as change
assert psbt_parser.change_amount == 10_000 assert psbt_parser.change_amount == 10_000
@@ -1593,7 +1614,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
# The output provided no derivation paths to verify (leaving the parser unable to # The output provided no derivation paths to verify (leaving the parser unable to
# determine if our seed owns any of the keys in the output's script). # determine if our seed owns any of the keys in the output's script).
assert psbt_parser.verified_output_derivation_paths[0] is None assert psbt_parser.verified_output_derivation_paths[0] == []
# Because we couldn't do proper verification, the parser correctly categorized the # Because we couldn't do proper verification, the parser correctly categorized the
# output as an external spend. # output as an external spend.
@@ -1620,7 +1641,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
# The claim itself still verifies # The claim itself still verifies
assert psbt_parser.verified_output_derivation_paths[0] is not None assert psbt_parser.verified_output_derivation_paths[0] != []
# But with no script there is no m-of-n to compare, so the output never becomes a # But with no script there is no m-of-n to compare, so the output never becomes a
# change candidate at all. # change candidate at all.
@@ -1668,7 +1689,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
# Even though the parser verified that our seed owns the internal key... # Even though the parser verified that our seed owns the internal key...
assert psbt_parser.verified_output_derivation_paths[0] is not None assert psbt_parser.verified_output_derivation_paths[0] != []
# ...the parser can't fully verify the output as change, so has to report it as an # ...the parser can't fully verify the output as change, so has to report it as an
# external spend. # external spend.
@@ -1712,7 +1733,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
psbt_parser = self._parse(psbt) psbt_parser = self._parse(psbt)
# The parser verified that we own the tapleaf key... # The parser verified that we own the tapleaf key...
assert psbt_parser.verified_output_derivation_paths[0] is not None assert psbt_parser.verified_output_derivation_paths[0] != []
# ...but the output still has to be reported as an external spend # ...but the output still has to be reported as an external spend
assert psbt_parser.change_amount == 0 assert psbt_parser.change_amount == 0
@@ -1899,33 +1920,29 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
self._parse(psbt) self._parse(psbt)
def test__parse__refuses_a_multisig_decoy_entry_in_either_position(self): def test__parse__refuses_a_multisig_decoy_entry(self):
""" """
In this scenario the multisig change output is a legitimate change output that In this scenario the multisig change output is a legitimate change output that
genuinely belongs to our seed, but a decoy derivation path entry is added. The genuinely belongs to our seed, but a decoy derivation path entry is added as an
decoy is ALSO a key that our seed owns, but it is not used in the output's script. extra entry or as a replacement for another cosigner's. The decoy is ALSO a key
that our seed owns, but it is not used in the output's script.
We don't need to decide if such a psbt has malicious intent; the fact that it We don't need to decide if such a psbt has malicious intent; the decoy is a
contradicts itself is unacceptable regardless: contradiction provable from the psbt alone, so we reject the psbt.
* it names a key on an output whose script does not use it.
* it names more keys than that script has.
Both are provable from the psbt alone, so we reject the psbt.
This is similar to the single sig test earlier in this class, but is more This is similar to the single sig test earlier in this class, but is more
complicated for multisig since it's the norm for multiple derivation paths to be complicated for multisig since it's the norm for multiple derivation paths to be
provided for each multisig change output. provided for each multisig change output.
The derivation path entries are provided in a coordinator-controlled order, so The derivation path entries are provided in a coordinator-controlled order, so
this test covers decoy entries that are listed before or after the seed's actual this test covers:
cosigner entry, across all three multisig script types. * Decoy listed before the seed's actual cosigner entry.
* Decoy listed after it.
* Decoy substituted for another cosigner's entry (so the output still lists
exactly as many entries as its script has keys).
Both orderings are refused. The ordering only decides which problem we report. The presence of a decoy in any of the placements should raise
We record the first entry that verifies against our seed, so: PSBTOutputOwnershipContradictionError.
* When the decoy is listed first, the decoy is what we record and it is not in
the script.
* When the decoy is listed last, the key we record is our real one and nothing
is wrong with it; what gives the decoy away instead is that the output named
more keys than its script has.
""" """
root = self._root() root = self._root()
@@ -1935,39 +1952,79 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
(PSBTTestData.MULTISIG_NESTED_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NESTED_SEGWIT_CHANGE), (PSBTTestData.MULTISIG_NESTED_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NESTED_SEGWIT_CHANGE),
(PSBTTestData.MULTISIG_LEGACY_P2SH_1_INPUT, PSBTTestData.MULTISIG_LEGACY_P2SH_CHANGE), (PSBTTestData.MULTISIG_LEGACY_P2SH_1_INPUT, PSBTTestData.MULTISIG_LEGACY_P2SH_CHANGE),
]: ]:
# ...run both versions of the test: decoy listed first and decoy last psbt = self._psbt_with_change(input_base64, change_hex)
for decoy_first in [True, False]: cosigner_entries = dict(psbt.outputs[0].bip32_derivations)
psbt = self._psbt_with_change(input_base64, change_hex)
cosigner_entries = dict(psbt.outputs[0].bip32_derivations) # Build the decoy from the cosigners' baseline, then make one minor
# derivation path change.
genuine_derivation_path = list(cosigner_entries.values())[0].derivation
decoy_derivation_path = genuine_derivation_path[:-1] + [genuine_derivation_path[-1] + 1]
decoy_public_key = root.derive(decoy_derivation_path).get_public_key()
decoy_entry = DerivationPath(root.my_fingerprint, decoy_derivation_path)
# Build the decoy from the cosigners' baseline, then make one minor # Decoy listed first (note: dicts preserve insertion order)
# derivation path change. decoy_first = {decoy_public_key: decoy_entry}
genuine_derivation_path = list(cosigner_entries.values())[0].derivation decoy_first.update(cosigner_entries)
decoy_derivation_path = genuine_derivation_path[:-1] + [genuine_derivation_path[-1] + 1]
decoy_public_key = root.derive(decoy_derivation_path).get_public_key()
decoy_entry = DerivationPath(root.my_fingerprint, decoy_derivation_path)
# Add the decoy to the existing 3 derivations # Decoy listed last
entries = psbt.outputs[0].bip32_derivations decoy_last = dict(cosigner_entries)
if decoy_first: decoy_last[decoy_public_key] = decoy_entry
entries.clear()
entries[decoy_public_key] = decoy_entry
entries.update(cosigner_entries)
else:
entries[decoy_public_key] = decoy_entry
if decoy_first: # Decoy in place of another cosigner's entry
# The parser uses the decoy as the comparison against which keys are decoy_substituted = dict(cosigner_entries)
# actually in the script. for public_key, entry in cosigner_entries.items():
expected_error = PSBTOutputOwnershipContradictionError if entry.fingerprint != root.my_fingerprint:
else: del decoy_substituted[public_key]
# The original cosigner is verified but then the parser detects the break
# decoy as a surplus derivation path. decoy_substituted[decoy_public_key] = decoy_entry
expected_error = PSBTSurplusDerivationPathsError assert len(decoy_substituted) == len(cosigner_entries)
with pytest.raises(expected_error): # Run all three placements of the decoy
self._parse(psbt) for entries in [decoy_first, decoy_last, decoy_substituted]:
psbt.outputs[0].bip32_derivations = entries
# Prep the modified psbt in embit
tampered_psbt = PSBT.parse(psbt.serialize())
with pytest.raises(PSBTOutputOwnershipContradictionError):
PSBTParser(tampered_psbt, self.seed, network=SettingsConstants.REGTEST)
# def test__parse__accepts_a_multisig_output_holding_this_seed_in_two_slots(self):
# """
# An edge case 2-of-3 that uses the same seed for two of its keys, each at its own
# derivation path. A legitimate change output for such a multisig should be
# recognized as change.
# Test not built; the setup complexity for this test is more effort than it's
# worth for a wallet nobody would / should set up.
# """
# pass
def test__parse__rejects_a_multisig_output_padded_with_a_strangers_entry(self):
"""
An honest multisig change output, plus one extra derivation path entry claiming a
stranger's fingerprint. Our own entry verifies and our key is in the script, so
the output's account of itself holds up as far as this seed can check. But the
script has only as many keys as it has cosigners, so the extra entry describes a
key the script never uses. The parser rejects the psbt with
PSBTSurplusDerivationPathsError.
"""
for input_base64, change_hex in [
(PSBTTestData.MULTISIG_NATIVE_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NATIVE_SEGWIT_CHANGE),
(PSBTTestData.MULTISIG_NESTED_SEGWIT_1_INPUT, PSBTTestData.MULTISIG_NESTED_SEGWIT_CHANGE),
(PSBTTestData.MULTISIG_LEGACY_P2SH_1_INPUT, PSBTTestData.MULTISIG_LEGACY_P2SH_CHANGE),
]:
psbt = self._psbt_with_change(input_base64, change_hex)
out = psbt.outputs[0]
assert len(out.bip32_derivations) == 3
claim_seed_owns_key(out, "m/48h/1h/0h/2h/1/0", foreign_public_key(), seed=PSBTTestData.recipient_multisig_key_2)
assert len(out.bip32_derivations) == 4
with pytest.raises(PSBTSurplusDerivationPathsError):
PSBTParser(psbt, self.seed, network=SettingsConstants.REGTEST)
def test__parse__rejects_a_multisig_output_whose_supplied_script_is_not_its_own(self): def test__parse__rejects_a_multisig_output_whose_supplied_script_is_not_its_own(self):
@@ -2140,7 +2197,7 @@ class TestPSBTParserOutputOwnership(PSBTParserOwnershipTestBase):
# This seed's key really is in the committed script and the psbt's claim of # This seed's key really is in the committed script and the psbt's claim of
# this seed verified. # this seed verified.
assert psbt_parser.verified_output_derivation_paths[0] is not None assert psbt_parser.verified_output_derivation_paths[0] != []
# But the output pays a different quorum than the inputs spend from, so it # But the output pays a different quorum than the inputs spend from, so it
# is counted as a spend. # is counted as a spend.