Improve balance update logic

This commit is contained in:
shroominic
2025-06-10 14:35:19 +02:00
parent 1547ce7afd
commit 10ceb1d74e
2 changed files with 29 additions and 7 deletions
+3 -2
View File
@@ -108,8 +108,9 @@ async def dashboard(request: Request) -> str:
f"<tr><td>{key.hashed_key}</td><td>{key.balance}</td><td>{key.total_spent}</td><td>{key.total_requests}</td><td>{key.refund_address}</td><td>{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}</td></tr>"
)
# 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
+26 -5
View File
@@ -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.")