mirror of
https://github.com/SeedSigner/seedsigner.git
synced 2026-09-14 04:45:08 +00:00
Scope the derivation cache to parse()
Nothing outside parse() reads the cache, so it need not be an instance attribute. As a local, "the cache does not outlive the parse" is a property of scope rather than a teardown to maintain -- which also retires the finally block, both initializations, and the test guarding them.
This commit is contained in:
@@ -33,13 +33,13 @@ class PSBTParser():
|
||||
for and needs no parse.
|
||||
"""
|
||||
|
||||
# Upper bound on how many levels of derivation a single parse will cache in
|
||||
# _child_key_derivation_cache. 1000 is just slightly under a 3-of-5 multisig
|
||||
# consolidating 200 inputs and holds the cache to a max of about 600 kilobytes. A
|
||||
# psbt that requires more levels will still parse correctly, but may have to derive
|
||||
# some levels more than once. Capping the cache at a realistic upper bound protects
|
||||
# against a maliciously crafted psbt that would otherwise consume unbounded memory
|
||||
# while still providing cache wins for even atypically large real-world psbts.
|
||||
# Upper bound on how many levels of derivation a single parse will cache. 1000 is
|
||||
# just slightly under a 3-of-5 multisig consolidating 200 inputs and holds the cache
|
||||
# to a max of about 600 kilobytes. A psbt that requires more levels will still parse
|
||||
# correctly, but may have to derive some levels more than once. Capping the cache at
|
||||
# a realistic upper bound protects against a maliciously crafted psbt that would
|
||||
# otherwise consume unbounded memory while still providing cache wins for even
|
||||
# atypically large real-world psbts.
|
||||
MAX_CACHED_DERIVATIONS = 1000
|
||||
|
||||
|
||||
@@ -60,7 +60,6 @@ class PSBTParser():
|
||||
self.op_return_data: bytes = None
|
||||
|
||||
self.root = None
|
||||
self._child_key_derivation_cache = {}
|
||||
|
||||
if self.seed is not None:
|
||||
self.parse()
|
||||
@@ -109,8 +108,10 @@ class PSBTParser():
|
||||
traversals overlap heavily: everything in one account shares the same opening
|
||||
levels, differing only in the address at the end.
|
||||
|
||||
So every level derived during this parse is kept in _child_key_derivation_cache
|
||||
and reused. See _derive_with_cache.
|
||||
So every level derived during this parse is kept in a cache and reused. See
|
||||
_derive_with_cache.
|
||||
|
||||
Note that the cache is only useful within a single parse so it is not preserved.
|
||||
"""
|
||||
if self.psbt is None:
|
||||
logger.info(f"self.psbt is None!!")
|
||||
@@ -122,29 +123,23 @@ class PSBTParser():
|
||||
|
||||
self._set_root()
|
||||
|
||||
self._child_key_derivation_cache = {}
|
||||
child_key_derivation_cache = {}
|
||||
|
||||
try:
|
||||
# Try to fix missing fingerprints before parsing
|
||||
self._fill_missing_fingerprints()
|
||||
# Try to fix missing fingerprints before parsing
|
||||
self._fill_missing_fingerprints(child_key_derivation_cache)
|
||||
|
||||
rt = self._parse_inputs()
|
||||
if rt == False:
|
||||
return False
|
||||
rt = self._parse_inputs(child_key_derivation_cache)
|
||||
if rt == False:
|
||||
return False
|
||||
|
||||
rt = self._parse_outputs()
|
||||
if rt == False:
|
||||
return False
|
||||
rt = self._parse_outputs(child_key_derivation_cache)
|
||||
if rt == False:
|
||||
return False
|
||||
|
||||
return True
|
||||
finally:
|
||||
# The cache is only useful within a single parse and it holds keys derived
|
||||
# from the signing seed, so drop it now rather than letting it live on for
|
||||
# as long as this parser does.
|
||||
self._child_key_derivation_cache = {}
|
||||
return True
|
||||
|
||||
|
||||
def _parse_inputs(self):
|
||||
def _parse_inputs(self, child_key_derivation_cache: dict):
|
||||
self.input_amount = 0
|
||||
self.num_inputs = len(self.psbt.inputs)
|
||||
for inp in self.psbt.inputs:
|
||||
@@ -155,14 +150,14 @@ class PSBTParser():
|
||||
self.input_amount += inp.utxo.value
|
||||
script_pubkey = inp.script_pubkey
|
||||
|
||||
inp_policy = PSBTParser._get_policy(inp, script_pubkey, self.psbt.xpubs, self._child_key_derivation_cache)
|
||||
inp_policy = PSBTParser._get_policy(inp, script_pubkey, self.psbt.xpubs, child_key_derivation_cache)
|
||||
if self.policy == None:
|
||||
self.policy = inp_policy
|
||||
else:
|
||||
if self.policy != inp_policy:
|
||||
raise RuntimeError("Mixed inputs in the transaction")
|
||||
|
||||
def _parse_outputs(self):
|
||||
def _parse_outputs(self, child_key_derivation_cache: dict):
|
||||
self.spend_amount = 0
|
||||
self.change_amount = 0
|
||||
self.change_data = []
|
||||
@@ -176,7 +171,7 @@ class PSBTParser():
|
||||
vout = self.psbt.tx.vout
|
||||
|
||||
for i, out in enumerate(self.psbt.outputs):
|
||||
out_policy = PSBTParser._get_policy(out, vout[i].script_pubkey, self.psbt.xpubs, self._child_key_derivation_cache)
|
||||
out_policy = PSBTParser._get_policy(out, vout[i].script_pubkey, self.psbt.xpubs, child_key_derivation_cache)
|
||||
is_change = False
|
||||
|
||||
# if policy is the same - probably change
|
||||
@@ -209,7 +204,7 @@ class PSBTParser():
|
||||
# should be one or zero for single-key addresses
|
||||
if len(out.bip32_derivations.values()) > 0:
|
||||
der = list(out.bip32_derivations.values())[0].derivation
|
||||
my_pubkey = PSBTParser._derive_with_cache(self.root, der, self._child_key_derivation_cache)
|
||||
my_pubkey = PSBTParser._derive_with_cache(self.root, der, child_key_derivation_cache)
|
||||
|
||||
if self.policy["type"] == "p2pkh" and my_pubkey is not None:
|
||||
sc = script.p2pkh(my_pubkey)
|
||||
@@ -230,7 +225,7 @@ class PSBTParser():
|
||||
# TODO: Support keys in taptree leaves
|
||||
leaf_hashes, derivation = list(out.taproot_bip32_derivations.values())[0]
|
||||
der = derivation.derivation
|
||||
my_pubkey = PSBTParser._derive_with_cache(self.root, der, self._child_key_derivation_cache)
|
||||
my_pubkey = PSBTParser._derive_with_cache(self.root, der, child_key_derivation_cache)
|
||||
sc = script.p2tr(my_pubkey)
|
||||
|
||||
if sc.data == vout[i].script_pubkey.data:
|
||||
@@ -523,7 +518,7 @@ class PSBTParser():
|
||||
return is_owner
|
||||
|
||||
|
||||
def _fill_missing_fingerprints(self):
|
||||
def _fill_missing_fingerprints(self, child_key_derivation_cache: dict):
|
||||
"""
|
||||
Fix for when fingerprint is missing (defaults to all zeros). Happens when the user
|
||||
creates a new wallet in an external coordinator but only provides the xpub
|
||||
@@ -552,7 +547,7 @@ class PSBTParser():
|
||||
# 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, self._child_key_derivation_cache)
|
||||
self.root, derivation_path_obj.derivation, child_key_derivation_cache)
|
||||
if derived_key.key.sec() == public_key.sec():
|
||||
return DerivationPath(self.root.my_fingerprint, derivation_path_obj.derivation)
|
||||
return None
|
||||
|
||||
@@ -535,8 +535,8 @@ class TestPSBTParserOptimizations:
|
||||
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.
|
||||
|
||||
Reading the cache back once the parse is over depends on the parse disposing of
|
||||
it by rebinding the attribute; recording sizes as the parse runs does not.
|
||||
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
|
||||
|
||||
@@ -763,21 +763,3 @@ class TestPSBTParserOptimizations:
|
||||
assert max(capped_sizes) == cap
|
||||
|
||||
|
||||
def test_cache_is_dropped_when_the_parse_ends(self):
|
||||
"""
|
||||
The cache holds keys derived from the signing seed, so the parser must not still
|
||||
be holding it once the parse it belongs to is over.
|
||||
"""
|
||||
psbt = PSBT.parse(a2b_base64(PSBTTestData.MULTISIG_NATIVE_SEGWIT_1_INPUT))
|
||||
psbt.outputs.append(create_output(PSBTTestData.MULTISIG_NATIVE_SEGWIT_CHANGE, 10_000))
|
||||
|
||||
cache_sizes = []
|
||||
with patch.object(PSBTParser, "_derive_with_cache", staticmethod(self.cache_size_recorder(cache_sizes))):
|
||||
psbt_parser = PSBTParser(psbt, self.seed, network=SettingsConstants.REGTEST)
|
||||
|
||||
# Sanity check: there is something to drop, i.e. the parse really did fill the
|
||||
# cache it was handed.
|
||||
assert max(cache_sizes) > 0
|
||||
|
||||
# But since the parse is done, the PSBTParser should have an empty cache again
|
||||
assert psbt_parser._child_key_derivation_cache == {}
|
||||
|
||||
Reference in New Issue
Block a user