From a5cee796c986f5d152eaa7d967d76ea24e9717bf Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 12 Jul 2026 14:45:17 +0200 Subject: [PATCH] fix: make auto-topup tokens recoverable --- routstr/upstream/auto_topup.py | 44 +++++++++- tests/unit/test_auto_topup.py | 143 +++++++++++++++++++++++++++++++++ 2 files changed, 185 insertions(+), 2 deletions(-) create mode 100644 tests/unit/test_auto_topup.py diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index 3582b3fc..3517be7d 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -4,7 +4,12 @@ import json from sqlmodel import select from ..core import get_logger -from ..core.db import UpstreamProviderRow, create_session +from ..core.db import ( + CashuTransaction, + UpstreamProviderRow, + create_session, + store_cashu_transaction, +) from ..wallet import send_token from .routstr import RoutstrUpstreamProvider @@ -123,7 +128,6 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: }, ) - print(amount, mint_url) try: token = await send_token(amount, "sat", mint_url) except Exception as e: @@ -138,6 +142,22 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) return + stored = await store_cashu_transaction( + token=token, + amount=amount, + unit="sat", + mint_url=mint_url, + typ="out", + collected=False, + source="auto_topup", + ) + if not stored: + logger.critical( + "Aborting auto top-up because its cashu token could not be persisted", + extra={"provider_id": row.id, "mint_url": mint_url}, + ) + return + result = await provider.topup(token) if "error" in result: @@ -149,6 +169,26 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: }, ) else: + async with create_session() as session: + transaction = ( + await session.exec( + select(CashuTransaction).where( + CashuTransaction.token == token, + CashuTransaction.type == "out", + CashuTransaction.source == "auto_topup", + ) + ) + ).first() + if transaction is None: + logger.critical( + "Completed auto top-up transaction is missing from the database", + extra={"provider_id": row.id, "mint_url": mint_url}, + ) + else: + transaction.collected = True + session.add(transaction) + await session.commit() + logger.info( "Auto top-up completed successfully", extra={ diff --git a/tests/unit/test_auto_topup.py b/tests/unit/test_auto_topup.py new file mode 100644 index 00000000..c05c85e5 --- /dev/null +++ b/tests/unit/test_auto_topup.py @@ -0,0 +1,143 @@ +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from routstr.core.db import CashuTransaction +from routstr.upstream.auto_topup import _check_and_topup + + +def _row() -> MagicMock: + row = MagicMock() + row.id = "provider-1" + row.base_url = "https://provider.test" + row.api_key = "secret" + row.provider_settings = json.dumps( + { + "auto_topup": True, + "topup_threshold": 100, + "topup_amount_limit": 50, + "topup_mint_url": "https://mint.test", + } + ) + return row + + +class _Session: + def __init__(self, transaction: CashuTransaction) -> None: + self.transaction = transaction + self.commit = AsyncMock() + + async def __aenter__(self) -> "_Session": + return self + + async def __aexit__(self, *args: object) -> None: + return None + + async def exec(self, query: object) -> MagicMock: + result = MagicMock() + result.first.return_value = self.transaction + return result + + def add(self, transaction: CashuTransaction) -> None: + self.transaction = transaction + + +@pytest.mark.asyncio +async def test_auto_topup_persists_before_sending_and_marks_success_collected() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=0) + provider.topup = AsyncMock(return_value={"balance": 50}) + transaction = CashuTransaction( + token="cashu-token", amount=50, unit="sat", source="auto_topup" + ) + session = _Session(transaction) + + with ( + patch( + "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup.send_token", + AsyncMock(return_value="cashu-token"), + ), + patch( + "routstr.upstream.auto_topup.store_cashu_transaction", + AsyncMock(return_value=True), + ) as store, + patch("routstr.upstream.auto_topup.create_session", return_value=session), + ): + await _check_and_topup(_row()) + + store.assert_awaited_once_with( + token="cashu-token", + amount=50, + unit="sat", + mint_url="https://mint.test", + typ="out", + collected=False, + source="auto_topup", + ) + provider.topup.assert_awaited_once_with("cashu-token") + assert transaction.collected is True + session.commit.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", [{"error": "rejected"}, RuntimeError("network")]) +async def test_auto_topup_failure_leaves_persisted_token_uncollected( + outcome: object, +) -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=0) + provider.topup = AsyncMock( + side_effect=outcome if isinstance(outcome, Exception) else None, + return_value=outcome, + ) + + with ( + patch( + "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup.send_token", + AsyncMock(return_value="cashu-token"), + ), + patch( + "routstr.upstream.auto_topup.store_cashu_transaction", + AsyncMock(return_value=True), + ), + patch("routstr.upstream.auto_topup.create_session") as create_session, + ): + if isinstance(outcome, Exception): + with pytest.raises(RuntimeError): + await _check_and_topup(_row()) + else: + await _check_and_topup(_row()) + + create_session.assert_not_called() + + +@pytest.mark.asyncio +async def test_auto_topup_does_not_send_untracked_token() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=0) + provider.topup = AsyncMock() + with ( + patch( + "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup.send_token", + AsyncMock(return_value="cashu-token"), + ), + patch( + "routstr.upstream.auto_topup.store_cashu_transaction", + AsyncMock(return_value=False), + ), + ): + await _check_and_topup(_row()) + provider.topup.assert_not_awaited()