diff --git a/src/seedsigner/models/psbt_parser.py b/src/seedsigner/models/psbt_parser.py index 1b087e2a..e27f9fad 100644 --- a/src/seedsigner/models/psbt_parser.py +++ b/src/seedsigner/models/psbt_parser.py @@ -598,33 +598,32 @@ class PSBTParser(): """ seed_fingerprint = seed.get_fingerprint(network) - def check_fingerprint_match(public_key: PublicKey, derivation_path_obj: DerivationPath): + def check_fingerprint_match(public_key: PublicKey, derivation_path_obj: DerivationPath, is_taproot: bool): """Check fingerprint match with missing fingerprint fallback""" # If exact fingerprint match if hexlify(derivation_path_obj.fingerprint).decode() == seed_fingerprint: return True - + # Missing fingerprint fallback if derivation_path_obj.fingerprint == b"\x00\x00\x00\x00": root = bip32.HDKey.from_seed(seed.seed_bytes, version=NETWORKS[SettingsConstants.map_network_to_embit(network)]["xprv"]) try: - derived_key = root.derive(derivation_path_obj.derivation) - return derived_key.key.sec() == public_key.sec() # Public keys match + return PSBTParser.seed_owns_pubkey(root, derivation_path_obj.derivation, public_key, child_key_derivation_cache=None, is_taproot=is_taproot) except Exception as e: logger.debug("Fingerprint fallback derive failed: %s", e, exc_info=True) return False - + # Check all derivations in all inputs for input in psbt.inputs: # Check regular BIP32 derivations for public_key, derivation_path_obj in input.bip32_derivations.items(): - if check_fingerprint_match(public_key, derivation_path_obj): + if check_fingerprint_match(public_key, derivation_path_obj, is_taproot=False): return True - + # Check Taproot derivations for public_key, (leaf_hashes, derivation_path_obj) in input.taproot_bip32_derivations.items(): - if check_fingerprint_match(public_key, derivation_path_obj): + if check_fingerprint_match(public_key, derivation_path_obj, is_taproot=True): return True return False @@ -644,7 +643,12 @@ class PSBTParser(): derived_public_key = PSBTParser._derive_with_cache(root, claimed_derivation_path, child_key_derivation_cache).get_public_key() if is_taproot: - # Taproot keys are x-only + # A psbt carries a taproot key as its bare 32-byte x coordinate, but embit + # rebuilds a full key from it by just assuming even parity. The key derived + # from the seed carries its real parity, so a naive full-key comparison + # succeeds only when that real parity happens to be even, wrongly rejecting + # roughly half of the keys this seed genuinely owns. Only the x coordinate is + # real information: compare x-only. return derived_public_key.xonly() == public_key.xonly() # For ecdsa the parity byte IS part of the identity, so compare the full key. @@ -785,32 +789,29 @@ class PSBTParser(): """Helper function to fill missing fingerprints in a scope (input/output)""" # Helper function to check and fix fingerprint - def _get_updated_fingerprint(public_key: PublicKey, derivation_path_obj: DerivationPath) -> DerivationPath | None: + def _get_updated_fingerprint(public_key: PublicKey, derivation_path_obj: DerivationPath, is_taproot: bool) -> DerivationPath | None: if derivation_path_obj.fingerprint != b"\x00\x00\x00\x00": return None - - # Derive the public key from the currently loaded seed using the derivation - # contained in the PSBT. If the derived public key exactly matches - # the PSBT-provided public key, we can be confident that this input/output - # is owned by the signing seed. In that case we populate the missing (zero) - # fingerprint with the signing seed's master fingerprint so downstream - # parsing/signing can treat it as owned by this seed. - derived_key = PSBTParser._derive_with_cache( - self.root, derivation_path_obj.derivation, child_key_derivation_cache) - if derived_key.key.sec() == public_key.sec(): + + # If the signing seed really derives the psbt-provided public key at the + # claimed derivation path, this input/output is owned by the signing seed. + # In that case we populate the missing (zero) fingerprint with the signing + # seed's master fingerprint so downstream parsing/signing can treat it as + # owned by this seed. + if PSBTParser.seed_owns_pubkey(self.root, derivation_path_obj.derivation, public_key, child_key_derivation_cache, is_taproot=is_taproot): return DerivationPath(self.root.my_fingerprint, derivation_path_obj.derivation) return None - + # Handle regular BIP32 derivations for public_key, derivation_path_obj in list(scope.bip32_derivations.items()): - new_derivation = _get_updated_fingerprint(public_key, derivation_path_obj) + new_derivation = _get_updated_fingerprint(public_key, derivation_path_obj, is_taproot=False) if new_derivation: scope.bip32_derivations[public_key] = new_derivation logger.debug(f"Filled missing fingerprint for pubkey {public_key.sec().hex()} derivation {bip32.path_to_str(derivation_path_obj.derivation)}") - - # Handle Taproot derivations + + # Handle Taproot derivations for public_key, (leaf_hashes, derivation_path_obj) in list(scope.taproot_bip32_derivations.items()): - new_derivation = _get_updated_fingerprint(public_key, derivation_path_obj) + new_derivation = _get_updated_fingerprint(public_key, derivation_path_obj, is_taproot=True) if new_derivation: scope.taproot_bip32_derivations[public_key] = (leaf_hashes, new_derivation) logger.debug(f"Filled missing fingerprint for pubkey {public_key.sec().hex()} derivation {bip32.path_to_str(derivation_path_obj.derivation)}") diff --git a/tests/test_psbt_parser.py b/tests/test_psbt_parser.py index 8e464be2..0b770120 100644 --- a/tests/test_psbt_parser.py +++ b/tests/test_psbt_parser.py @@ -212,15 +212,64 @@ class TestPSBTParser: from binascii import hexlify fingerprint_hex = hexlify(derivation.fingerprint).decode() - # Check if this public key derives from the current seed + # Check if this public key derives from the current seed. A psbt + # carries a taproot key as its bare 32-byte x coordinate, and embit + # rebuilds a full key from it by just assuming even parity. The real + # derived key can be odd-parity, so a full-key compare would wrongly + # report a mismatch. Only the x coordinate is real data: compare + # x-only. derived_key = parser.root.derive(derivation.derivation) - if derived_key.key.sec() == pub.sec(): + if derived_key.xonly() == pub.xonly(): # This pubkey derives from current seed, should have current seed's fingerprint assert fingerprint_hex == seed_fingerprint, f"Expected {seed_fingerprint}, got {fingerprint_hex} for taproot pubkey that derives from current seed" else: # This pubkey doesn't derive from current seed, should remain 00000000 assert fingerprint_hex == "00000000" + # All of the above only proves the even-parity case. A psbt carries taproot keys + # as bare 32-byte x coordinates and embit rebuilds full keys from them by assuming + # even parity; that assumption happens to hold for the fixture's key at + # m/86h/1h/0h/0/0. Re-key the taproot input to a path whose key really derives + # with odd parity to prove the ownership fallback compares x-only rather than + # trusting embit's artificial parity. + root = root_for_seed(PSBTTestData.seed) + odd_parity_derivation_path = "m/86h/1h/0h/0/1" + odd_parity_public_key = root.derive(odd_parity_derivation_path).get_public_key() + assert odd_parity_public_key.sec()[0] == 0x03 # odd parity + + psbt = PSBT.parse(a2b_base64(PSBTTestData.SINGLE_SIG_TAPROOT_1_INPUT)) + taproot_input = psbt.inputs[0] + + # Present the key the way embit's psbt parsing yields it: rebuilt from just the + # x coordinate, carrying the assumed even parity (wrong for this key) + x_only_public_key = PublicKey.from_xonly(odd_parity_public_key.xonly()) + taproot_input.taproot_bip32_derivations.clear() + taproot_input.taproot_bip32_derivations[x_only_public_key] = ([], DerivationPath( + fingerprint=b"\x00\x00\x00\x00", + derivation=bip32.parse_path(odd_parity_derivation_path) + )) + taproot_input.taproot_internal_key = x_only_public_key + taproot_input.witness_utxo.script_pubkey = script.p2tr(x_only_public_key) + + # The zeroed-fingerprint fallback check must recognize this input as the seed's, + # even though embit's internal parity byte for the pubkey is wrong. Taproot + # pubkeys must be compared by their x-only representation. + assert PSBTParser.has_matching_input_fingerprint(psbt, PSBTTestData.seed, SettingsConstants.REGTEST) + + # Comparing x-only looks less strict than the full-key comparison used for + # non-taproot keys, but nothing is actually given up: a psbt never carries a + # parity byte for a taproot key, so the x coordinate is all the key material + # there is to compare. A completely wrong seed will still fail to match. + wrong_seed = Seed(["bacon"] * 24) + assert not PSBTParser.has_matching_input_fingerprint(psbt, wrong_seed, SettingsConstants.REGTEST) + + # Parsing should successfully fill the fingerprint and verify that the input + # belongs to the seed. + parser = PSBTParser(p=psbt, seed=PSBTTestData.seed, network=SettingsConstants.REGTEST) + (_, filled_derivation) = parser.psbt.inputs[0].taproot_bip32_derivations[x_only_public_key] + assert filled_derivation.fingerprint == parser.root.my_fingerprint + assert parser.verified_input_derivation_paths == [bip32.parse_path(odd_parity_derivation_path)] + def test_trim_and_sig_count(self): """