Rename _derive_with_cache to _derive_with_cache_via_indices

Names the input shape: a list of child indices below parent_key, which
is what the cache is keyed on. A companion that takes a psbt entry's
DerivationPath object follows; the two names then say which one a caller
holds.
This commit is contained in:
kdmukai
2026-09-19 15:41:48 -05:00
parent 45a5eabb9d
commit 98b46f5a29
2 changed files with 19 additions and 18 deletions
+8 -8
View File
@@ -267,7 +267,7 @@ class PSBTParser():
levels, differing only in the address at the end.
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.
"""
@@ -437,7 +437,7 @@ class PSBTParser():
# Rebuild the scriptPubKey from the key at the claimed derivation path
if len(out.bip32_derivations.values()) == 1:
singlesig_derivation_path = list(out.bip32_derivations.values())[0].derivation
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_indices(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)
else:
# There's nothing for us to verify against so this output will be
@@ -464,7 +464,7 @@ class PSBTParser():
if len(taproot_entries) == 1 and internal_key_claims == 1:
leaf_hashes, derivation = taproot_entries[0]
singlesig_derivation_path = derivation.derivation
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_indices(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)
else:
# This output has at least one derivation path entry for a key
@@ -517,7 +517,7 @@ class PSBTParser():
# the coordinator says sits there. Both are its own
# claims, so we read only the path and derive the key
# 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_indices(self.root, derivation_path_obj.derivation, child_key_derivation_cache).get_public_key()
if PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key):
# The output pays a multisig this seed is part
@@ -535,7 +535,7 @@ class PSBTParser():
# This output claimed that our seed is part of the receiving
# multisig, at a specific path. So now we verify that the key
# at that path is in the committed script.
seed_public_key = PSBTParser._derive_with_cache(self.root, verified_derivation_path, child_key_derivation_cache).get_public_key()
seed_public_key = PSBTParser._derive_with_cache_via_indices(self.root, verified_derivation_path, child_key_derivation_cache).get_public_key()
if not PSBTParser._multisig_script_contains_key(multisig_script, seed_public_key):
# The psbt said this output was coming back to our seed
# at that path, but the key there is not in the committed
@@ -775,7 +775,7 @@ class PSBTParser():
@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
any levels along the way that have already been derived during this parse.
@@ -883,7 +883,7 @@ class PSBTParser():
if origin_der.derivation == der.derivation[:-2]:
# Derive the child key that sits two indices below the xpub (i.e. at
# 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
if derived_key.key == pubkey:
@@ -975,7 +975,7 @@ class PSBTParser():
say anything. Ownership is established here and only here, by deriving the key
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:
# A psbt carries a taproot key as its bare 32-byte x coordinate, but embit
+11 -10
View File
@@ -569,13 +569,14 @@ class TestPSBTParserOptimizations:
def cache_size_recorder(self, cache_sizes: list):
"""
Returns a stand-in for _derive_with_cache that derives exactly as the real one
does, but appends the cache's size to cache_sizes on the way out of every call.
Returns a stand-in for _derive_with_cache_via_indices that derives exactly as the
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
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):
derived_key = real_derive_with_cache(parent_key, derivation_path, cache)
@@ -671,8 +672,8 @@ class TestPSBTParserOptimizations:
# 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
# confuse/combine the two cosigners' derivation data.
from_a = PSBTParser._derive_with_cache(cosigner_a_xpub, receive_index_5, cache)
from_b = PSBTParser._derive_with_cache(cosigner_b_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_via_indices(cosigner_b_xpub, receive_index_5, cache)
# Two levels should have been added for each cosigner
assert len(cache) == 4
@@ -738,7 +739,7 @@ class TestPSBTParserOptimizations:
def assert_cache_makes_no_difference(input_base64: str, change_hex: str):
# Store the real function before the patches below replace it. Each replacement
# 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
# function does, but will also record the cache it was handed on each call.
@@ -747,7 +748,7 @@ class TestPSBTParserOptimizations:
caches_received.append(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(
build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST)
@@ -759,7 +760,7 @@ class TestPSBTParserOptimizations:
def cache_free_derive(parent_key, derivation_path, cache=None):
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(
build_psbt(input_base64, change_hex), self.seed, network=SettingsConstants.REGTEST)
@@ -791,7 +792,7 @@ class TestPSBTParserOptimizations:
# Record how large the cache grew over the course of each parse
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)
singlesig_unconstrained = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST)
@@ -800,7 +801,7 @@ class TestPSBTParserOptimizations:
cap = 3
capped_sizes = []
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)
singlesig_capped = PSBTParser(build_psbt(singlesig_case), self.seed, network=SettingsConstants.REGTEST)