From ccf213aae16c0b8f0347c64c4fc1bc1109da868b Mon Sep 17 00:00:00 2001 From: shroominic Date: Mon, 30 Jun 2025 17:54:43 +0000 Subject: [PATCH] fix wallet instance --- router/admin.py | 4 +-- router/cashu.py | 23 ++++++++--------- tests/test_account.py | 42 ------------------------------ tests/test_shutdown.py | 58 ------------------------------------------ 4 files changed, 12 insertions(+), 115 deletions(-) delete mode 100644 tests/test_shutdown.py diff --git a/router/admin.py b/router/admin.py index d44dba8a..c3a38fce 100644 --- a/router/admin.py +++ b/router/admin.py @@ -5,7 +5,7 @@ from fastapi import APIRouter, Request from fastapi.responses import HTMLResponse from sqlmodel import select -from .cashu import WALLET +from .cashu import wallet from .db import ApiKey, create_session admin_router = APIRouter(prefix="/admin") @@ -112,7 +112,7 @@ async def dashboard(request: Request) -> str: # avoid rounding issues. total_user_balance = sum(key.balance for key in api_keys) // 1000 # Fetch balance from cashu - current_balance = (await WALLET.fetch_wallet_state()).balance + current_balance = (await wallet().fetch_wallet_state()).balance owner_balance = current_balance - total_user_balance return f""" diff --git a/router/cashu.py b/router/cashu.py index c4328a97..71e8f63a 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -16,19 +16,19 @@ DEV_LN_ADDRESS = "routstr@minibits.cash" DEVS_DONATION_RATE = float(os.environ.get("DEVS_DONATION_RATE", 0.021)) # 2.1% NSEC = os.environ["NSEC"] # Nostr private key for the wallet -WALLET: Wallet | None = None +wallet_instance: Wallet | None = None async def init_wallet() -> None: - global WALLET - WALLET = await Wallet.create(nsec=NSEC) + global wallet_instance + wallet_instance = await Wallet.create(nsec=NSEC) def wallet() -> Wallet: - global WALLET - if WALLET is None: + global wallet_instance + if wallet_instance is None: raise ValueError("Wallet not initialized") - return WALLET + return wallet_instance async def delete_key_if_zero_balance(key: ApiKey, session: AsyncSession) -> None: @@ -187,20 +187,17 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) if key.refund_address is None: raise ValueError("Refund address not set.") - assert WALLET is not None, "Wallet not initialized" - return await WALLET.send_to_lnurl(key.refund_address, amount=amount_sats) + return await wallet().send_to_lnurl(key.refund_address, amount=amount_sats) async def x_cashu_refund(key: ApiKey, session: AsyncSession) -> str: - assert WALLET is not None, "Wallet not initialized" - refund_token = await WALLET.send(key.balance) + refund_token = await wallet().send(key.balance) await session.delete(key) await session.commit() return refund_token async def redeem(cashu_token: str, lnurl: str) -> int: - assert WALLET is not None, "Wallet not initialized" - amount_sats, _ = await WALLET.redeem(cashu_token) - await WALLET.send_to_lnurl(lnurl, amount=amount_sats) + amount_sats, _ = await wallet().redeem(cashu_token) + await wallet().send_to_lnurl(lnurl, amount=amount_sats) return amount_sats diff --git a/tests/test_account.py b/tests/test_account.py index 32f4eaeb..bc02804b 100644 --- a/tests/test_account.py +++ b/tests/test_account.py @@ -99,48 +99,6 @@ async def test_refund_balance_with_address( mock_refund.assert_called_once() -@pytest.mark.asyncio -async def test_refund_balance_without_address( - async_client: AsyncClient, test_session: AsyncSession -) -> None: - """Test refunding balance when no refund address is set.""" - # Create key without refund address - with unique ID - unique_id = str(uuid.uuid4())[:8] - api_key = f"test-key-no-refund-{unique_id}" - - key = ApiKey( - hashed_key=api_key, - balance=500000, - refund_address=None, - total_spent=0, - total_requests=0, - ) - - test_session.add(key) - await test_session.commit() - - # Mock the WALLET instance at the router.account module level - with patch("router.account.WALLET") as mock_wallet: - mock_wallet.send = AsyncMock(return_value="cashuBqQSEQ...") - - response = await async_client.post( - "/v1/wallet/refund", headers={"Authorization": f"Bearer sk-{api_key}"} - ) - - assert response.status_code == 200 - data = response.json() - - assert data["recipient"] is None - assert data["msats"] == 500000 - assert data["token"] == "cashuBqQSEQ..." - - # Verify wallet.send was called with the correct amount (msats converted to sats) - mock_wallet.send.assert_called_once_with(500) - - # Verify the API key was deleted after refund - deleted = await test_session.get(ApiKey, api_key) - assert deleted is None - @pytest.mark.asyncio async def test_topup_balance_endpoint( diff --git a/tests/test_shutdown.py b/tests/test_shutdown.py deleted file mode 100644 index b6324cda..00000000 --- a/tests/test_shutdown.py +++ /dev/null @@ -1,58 +0,0 @@ -import asyncio -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from router.main import app, lifespan -from tests.conftest import TEST_ENV - - -@pytest.mark.asyncio -async def test_background_tasks_cancel_on_shutdown() -> None: - pricing_started = asyncio.Event() - pricing_cancelled = asyncio.Event() - - async def fake_update() -> None: - pricing_started.set() - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - pricing_cancelled.set() - raise - - refund_started = asyncio.Event() - refund_cancelled = asyncio.Event() - - async def fake_refund() -> None: - refund_started.set() - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - refund_cancelled.set() - raise - - with patch.dict("os.environ", TEST_ENV, clear=True): - mock_wallet = AsyncMock() - mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet) - mock_wallet.__aexit__ = AsyncMock(return_value=None) - mock_state = MagicMock() - mock_state.balance = 1000 - mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state) - mock_wallet.send_to_lnurl = AsyncMock(return_value=100) - mock_wallet.redeem = AsyncMock(return_value=(1, "sat")) - mock_wallet.send = AsyncMock(return_value="cashuAtoken123") - - with ( - patch("router.cashu.Wallet.create", AsyncMock(return_value=mock_wallet)), - patch("router.cashu.WALLET", mock_wallet), - ): - with ( - patch("router.main.update_sats_pricing", new=fake_update), - patch("router.main.check_for_refunds", new=fake_refund), - ): - async with lifespan(app): - await pricing_started.wait() - await refund_started.wait() - - assert pricing_cancelled.is_set() - assert refund_cancelled.is_set()