diff --git a/routstr/wallet.py b/routstr/wallet.py index fe34d7bb..1fef36c6 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -5,6 +5,7 @@ from typing import TypedDict from cashu.core.base import Proof, Token from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet +from sqlmodel import col, update from .core import db, get_logger from .core.settings import settings @@ -124,9 +125,17 @@ async def credit_balance( "credit_balance: Updating balance", extra={"old_balance": key.balance, "credit_amount": amount}, ) - key.balance += amount - session.add(key) + + # Use atomic SQL UPDATE to prevent race conditions during concurrent topups + stmt = ( + update(db.ApiKey) + .where(col(db.ApiKey.hashed_key) == key.hashed_key) + .values(balance=(db.ApiKey.balance) + amount) + ) + await session.exec(stmt) # type: ignore[call-overload] await session.commit() + await session.refresh(key) + logger.info( "credit_balance: Balance updated successfully", extra={"new_balance": key.balance}, diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 60b93b5a..f9a49434 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -4,6 +4,7 @@ from unittest.mock import AsyncMock, Mock, patch import pytest +from routstr.core.db import ApiKey from routstr.wallet import credit_balance, get_balance, recieve_token, send_token @@ -82,8 +83,15 @@ async def test_credit_balance() -> None: mock_key = Mock() mock_key.balance = 5000000 + mock_key.hashed_key = "test_hash" mock_session = AsyncMock() + # Mock session.refresh to update the balance (simulates DB reload) + async def mock_refresh(key: ApiKey) -> None: + key.balance = 6000000 + + mock_session.refresh.side_effect = mock_refresh + from routstr.core.settings import settings with patch.object(settings, "cashu_mints", ["http://mint:3338"]): @@ -93,9 +101,11 @@ async def test_credit_balance() -> None: ): amount = await credit_balance(token_str, mock_key, mock_session) assert amount == 1000000 # converted to msat - assert mock_key.balance == 6000000 - mock_session.add.assert_called_once_with(mock_key) - mock_session.commit.assert_called_once() + assert mock_key.balance == 6000000 # Should be updated after refresh + # Verify atomic operations were used + assert mock_session.exec.called # Atomic UPDATE statement + assert mock_session.commit.called + assert mock_session.refresh.called @pytest.mark.asyncio