diff --git a/router/admin.py b/router/admin.py index 8881551a..a4734de6 100644 --- a/router/admin.py +++ b/router/admin.py @@ -108,8 +108,9 @@ async def dashboard(request: Request) -> str: f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}" ) - # Calculate the total balance of all API keys - total_user_balance = int(sum(key.balance / 1000 for key in api_keys)) + # Calculate the total balance of all API keys using integer arithmetic to + # 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 owner_balance = current_balance - total_user_balance diff --git a/router/cashu.py b/router/cashu.py index baa68d47..d60fdc95 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -4,6 +4,7 @@ import time from sixty_nuts import Wallet from sqlmodel import select, func, col +from sqlalchemy import update from .db import ApiKey, AsyncSession, get_session @@ -76,9 +77,21 @@ async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) - await WALLET.redeem(cashu_token) state_after = await WALLET.fetch_wallet_state() amount = (state_after.balance - state_before.balance) * 1000 - key.balance += amount + + # Ensure the key is persisted so the update statement can succeed session.add(key) + await session.flush() + + # Apply the balance change atomically to avoid race conditions when topping + # up the same key concurrently. + stmt = ( + update(ApiKey) + .where(ApiKey.hashed_key == key.hashed_key) + .values(balance=ApiKey.balance + amount) + ) + await session.exec(stmt) await session.commit() + await session.refresh(key) return amount @@ -122,8 +135,6 @@ async def check_for_refunds() -> None: async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) -> int: - if key.balance < amount_msats: - raise ValueError("Insufficient balance.") if amount_msats <= 0: amount_msats = key.balance @@ -132,9 +143,19 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) if amount_sats == 0: raise ValueError("Amount too small to refund (less than 1 sat)") - key.balance -= amount_msats - session.add(key) + # Atomically deduct the balance to avoid race conditions when multiple + # refunds are triggered concurrently. + stmt = ( + update(ApiKey) + .where(ApiKey.hashed_key == key.hashed_key) + .where(ApiKey.balance >= amount_msats) + .values(balance=ApiKey.balance - amount_msats) + ) + result = await session.exec(stmt) await session.commit() + if result.rowcount == 0: + raise ValueError("Insufficient balance.") + await session.refresh(key) if key.refund_address is None: raise ValueError("Refund address not set.")