diff --git a/routstr/wallet.py b/routstr/wallet.py index b1297614..717539bd 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -19,13 +19,7 @@ RECEIVE_LN_ADDRESS = os.environ.get("RECEIVE_LN_ADDRESS", "") async def get_balance(unit: str) -> int: - wallet = await Wallet.with_db( - PRIMARY_MINT_URL, - db=".wallet", - load_all_keysets=True, - unit=unit, - ) - await wallet.load_proofs() + wallet = await get_wallet(PRIMARY_MINT_URL, unit) return wallet.available_balance.amount @@ -36,14 +30,8 @@ async def recieve_token( if len(token_obj.keysets) > 1: raise ValueError("Multiple keysets per token currently not supported") - # TODO check if can be initialized differently - wallet = await Wallet.with_db( - token_obj.mint, - db=".wallet", - load_all_keysets=True, - unit=token_obj.unit, - ) - await wallet.load_mint(token_obj.keysets[0]) + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) + wallet.keyset_id = token_obj.keysets[0] if token_obj.mint not in TRUSTED_MINTS: return await swap_to_primary_mint(token_obj, wallet) @@ -55,12 +43,7 @@ async def recieve_token( async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: """Internal send function - returns amount and serialized token""" - # TODO check if can be initialized differently - wallet: Wallet = await Wallet.with_db( - mint_url or PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit=unit - ) - await wallet.load_mint() - await wallet.load_proofs() + wallet: Wallet = await get_wallet(mint_url or PRIMARY_MINT_URL, unit) proofs = await get_proofs_per_mint_and_unit( wallet, mint_url or PRIMARY_MINT_URL, unit ) @@ -98,11 +81,7 @@ async def swap_to_primary_mint( raise ValueError("Invalid unit") estimated_fee_sat = int(max(amount_msat // 1000 * 0.01, 2)) amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000 - # TODO check if can be initialized differently - primary_wallet = await Wallet.with_db( - PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit="sat" - ) - await primary_wallet.load_mint() + primary_wallet = await get_wallet(PRIMARY_MINT_URL, "sat") minted_amount = amount_msat_after_fee // 1000 mint_quote = await primary_wallet.request_mint(minted_amount) @@ -165,13 +144,21 @@ async def credit_balance( raise -async def get_wallet(mint_url: str, unit: str = "sat") -> Wallet: - wallet = await Wallet.with_db( - mint_url, db=".wallet", load_all_keysets=True, unit=unit - ) - await wallet.load_mint() - await wallet.load_proofs(reload=True) - return wallet +_wallets: dict[str, Wallet] = {} + + +async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: + global _wallets + id = f"{mint_url}_{unit}" + if id not in _wallets: + _wallets[id] = await Wallet.with_db( + mint_url, db=".wallet", load_all_keysets=True, unit=unit + ) + + if load: + await _wallets[id].load_mint() + await _wallets[id].load_proofs(reload=True) + return _wallets[id] async def get_proofs_per_mint_and_unit( @@ -191,14 +178,16 @@ async def get_proofs_per_mint_and_unit( async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[Proof]: if not proofs: return [] - proof_states = await wallet.check_proof_state(proofs) _proofs = [] _spent_proofs = [] - for proof, state in zip(proofs, proof_states.states): - if str(state.state) != "spent": - _proofs.append(proof) - else: - _spent_proofs.append(proof) + for i in range(0, len(proofs), 1000): + pb = proofs[i : i + 1000] + proof_states = await wallet.check_proof_state(pb) + for proof, state in zip(pb, proof_states.states): + if str(state.state) != "spent": + _proofs.append(proof) + else: + _spent_proofs.append(proof) await wallet.set_reserved_for_send(_spent_proofs, reserved=True) return _proofs @@ -309,7 +298,7 @@ async def periodic_payout() -> None: logger.error("RECEIVE_LN_ADDRESS is not set, skipping payout") return while True: - await asyncio.sleep(60) + await asyncio.sleep(60 * 5) try: async with db.create_session() as session: for mint_url in TRUSTED_MINTS: diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index dc0870f1..b432cd9f 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -11,6 +11,7 @@ from routstr.wallet import credit_balance, get_balance, recieve_token, send_toke async def test_get_balance() -> None: mock_wallet = Mock() mock_wallet.available_balance = Mock(amount=50000) + mock_wallet.load_mint = AsyncMock() mock_wallet.load_proofs = AsyncMock() with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): @@ -48,9 +49,9 @@ async def test_recieve_token_valid() -> None: mock_token.proofs = [{"amount": 1000}] mock_deserialize.return_value = mock_token + mock_wallet.load_mint = AsyncMock() + mock_wallet.load_proofs = AsyncMock() with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): - mock_wallet.load_mint = AsyncMock() - amount, unit, mint = await recieve_token(token_str) assert amount == 1000 assert unit == "sat" @@ -105,8 +106,9 @@ async def test_recieve_token_untrusted_mint() -> None: mock_token.amount = 1000 mock_deserialize.return_value = mock_token + mock_wallet.load_mint = AsyncMock() + mock_wallet.load_proofs = AsyncMock() with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): - mock_wallet.load_mint = AsyncMock() with patch( "routstr.wallet.swap_to_primary_mint", return_value=(900, "sat", "http://mint:3338"),