add api-key history and fix race condition while topup

This commit is contained in:
9qeklajc
2026-04-22 21:50:22 +02:00
parent c0c8cafd00
commit 05115c3387
10 changed files with 339 additions and 16 deletions
@@ -0,0 +1,46 @@
"""add api key link to cashu_transactions
Revision ID: d4e5f6a7b8c9
Revises: c3d4e5f6a7b8
Create Date: 2026-04-20 00:00:00.000000
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "d4e5f6a7b8c9"
down_revision = "c3d4e5f6a7b8"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col["name"] for col in inspector.get_columns("cashu_transactions")]
indexes = {index["name"] for index in inspector.get_indexes("cashu_transactions")}
if "api_key_hashed_key" not in columns:
op.add_column(
"cashu_transactions",
sa.Column(
"api_key_hashed_key",
sqlmodel.sql.sqltypes.AutoString(),
nullable=True,
),
)
if "ix_cashu_transactions_api_key_hashed_key" not in indexes:
op.create_index(
"ix_cashu_transactions_api_key_hashed_key",
"cashu_transactions",
["api_key_hashed_key"],
unique=False,
)
def downgrade() -> None:
op.drop_index("ix_cashu_transactions_api_key_hashed_key", table_name="cashu_transactions")
op.drop_column("cashu_transactions", "api_key_hashed_key")
+83 -11
View File
@@ -7,7 +7,7 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from pydantic import BaseModel from pydantic import BaseModel
from sqlmodel import select from sqlmodel import col, select, update
from .auth import get_billing_key, validate_bearer_key from .auth import get_billing_key, validate_bearer_key
from .core.db import ( from .core.db import (
@@ -211,6 +211,26 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
_refund_cache[key] = (expiry, value) _refund_cache[key] = (expiry, value)
async def _restore_balance(
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int
) -> None:
"""Restore balance after a failed refund mint attempt."""
restore_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == hashed_key)
.values(
balance=col(ApiKey.balance) + balance,
reserved_balance=col(ApiKey.reserved_balance) + reserved_balance,
)
)
await session.exec(restore_stmt) # type: ignore[call-overload]
await session.commit()
logger.info(
"refund_wallet_endpoint: balance restored after mint failure",
extra={"hashed_key": hashed_key, "restored_balance": balance},
)
@router.post("/refund", response_model=None) @router.post("/refund", response_model=None)
async def refund_wallet_endpoint( async def refund_wallet_endpoint(
authorization: Annotated[str | None, Header()] = None, authorization: Annotated[str | None, Header()] = None,
@@ -292,7 +312,27 @@ async def refund_wallet_endpoint(
elif remaining_balance <= 0: elif remaining_balance <= 0:
raise HTTPException(status_code=400, detail="No balance to refund") raise HTTPException(status_code=400, detail="No balance to refund")
# Perform refund operation first, before modifying balance # --- DEBIT FIRST: atomically zero the balance before minting tokens ---
# This prevents the race where a concurrent topup/spend happens between
# reading the balance and minting the refund token (double-spend).
debit_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) == key.balance)
.where(col(ApiKey.reserved_balance) == key.reserved_balance)
.values(balance=0, reserved_balance=0)
)
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload]
await session.commit()
if debit_result.rowcount == 0:
# Balance changed between read and debit — another request is active
raise HTTPException(
status_code=409,
detail="Balance changed concurrently. Please retry the refund.",
)
# --- MINT: balance is locked at zero, safe to create the refund token ---
try: try:
if key.refund_address: if key.refund_address:
from .core.settings import settings as global_settings from .core.settings import settings as global_settings
@@ -328,10 +368,12 @@ async def refund_wallet_endpoint(
) )
except HTTPException: except HTTPException:
# Re-raise HTTP exceptions (like 400 for balance too small) # Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, key.balance, key.reserved_balance)
raise raise
except Exception as e: except Exception as e:
# If refund fails, don't modify the database # Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, key.balance, key.reserved_balance)
error_msg = str(e) error_msg = str(e)
if ( if (
"mint" in error_msg.lower() "mint" in error_msg.lower()
@@ -345,12 +387,6 @@ async def refund_wallet_endpoint(
await _refund_cache_set(bearer_value, result) await _refund_cache_set(bearer_value, result)
previous_reserved_balance = key.reserved_balance
key.balance = 0
key.reserved_balance = 0
session.add(key)
await session.commit()
if "token" in result: if "token" in result:
try: try:
await store_cashu_transaction( await store_cashu_transaction(
@@ -361,6 +397,7 @@ async def refund_wallet_endpoint(
typ="out", typ="out",
collected=False, collected=False,
source="apikey", source="apikey",
api_key_hashed_key=key.hashed_key,
) )
except Exception: except Exception:
pass # store_cashu_transaction already logs pass # store_cashu_transaction already logs
@@ -369,13 +406,48 @@ async def refund_wallet_endpoint(
"refund_wallet_endpoint: refund successful", "refund_wallet_endpoint: refund successful",
extra={ extra={
"refunded_msats": remaining_balance_msats, "refunded_msats": remaining_balance_msats,
"previous_reserved_balance": previous_reserved_balance, "previous_reserved_balance": key.reserved_balance,
}, },
) )
return result return result
@router.get("/history")
async def wallet_history(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, list[dict[str, str | int | bool | None]]]:
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot view child key history. Please use the parent key instead.",
)
result = await session.exec(
select(CashuTransaction)
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
.order_by(col(CashuTransaction.created_at).desc())
)
transactions = result.all()
return {
"transactions": [
{
"id": tx.id,
"type": tx.type,
"source": tx.source,
"amount": tx.amount,
"unit": tx.unit,
"mint_url": tx.mint_url,
"created_at": tx.created_at,
"collected": tx.collected,
"swept": tx.swept,
}
for tx in transactions
]
}
@router.post("/donate") @router.post("/donate")
async def donate(token: str, ref: str | None = None) -> str: async def donate(token: str, ref: str | None = None) -> str:
try: try:
+1
View File
@@ -1360,6 +1360,7 @@ async def get_transactions_api(
(col(CashuTransaction.id).like(search_pattern)) (col(CashuTransaction.id).like(search_pattern))
| (col(CashuTransaction.token).like(search_pattern)) | (col(CashuTransaction.token).like(search_pattern))
| (col(CashuTransaction.request_id).like(search_pattern)) | (col(CashuTransaction.request_id).like(search_pattern))
| (col(CashuTransaction.api_key_hashed_key).like(search_pattern))
) )
stmt = stmt.order_by(col(CashuTransaction.created_at).desc()).limit(limit) stmt = stmt.order_by(col(CashuTransaction.created_at).desc()).limit(limit)
+8
View File
@@ -159,6 +159,12 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
default="x-cashu", default="x-cashu",
description="Payment source: x-cashu or apikey", description="Payment source: x-cashu or apikey",
) )
api_key_hashed_key: str | None = Field(
default=None,
foreign_key="api_keys.hashed_key",
index=True,
description="Associated API key hash for wallet history",
)
async def store_cashu_transaction( async def store_cashu_transaction(
@@ -171,6 +177,7 @@ async def store_cashu_transaction(
collected: bool = False, collected: bool = False,
created_at: int | None = None, created_at: int | None = None,
source: str = "x-cashu", source: str = "x-cashu",
api_key_hashed_key: str | None = None,
) -> None: ) -> None:
try: try:
async with create_session() as session: async with create_session() as session:
@@ -184,6 +191,7 @@ async def store_cashu_transaction(
collected=collected, collected=collected,
created_at=created_at or int(time.time()), created_at=created_at or int(time.time()),
source=source, source=source,
api_key_hashed_key=api_key_hashed_key,
) )
session.add(tx) session.add(tx)
await session.commit() await session.commit()
+16
View File
@@ -8,6 +8,7 @@ from cashu.wallet.wallet import Wallet
from sqlmodel import col, select, update from sqlmodel import col, select, update
from .core import db, get_logger from .core import db, get_logger
from .core.db import store_cashu_transaction
from .core.settings import settings from .core.settings import settings
from .payment.lnurl import raw_send_to_lnurl from .payment.lnurl import raw_send_to_lnurl
@@ -279,6 +280,8 @@ async def credit_balance(
try: try:
amount, unit, mint_url = await recieve_token(cashu_token) amount, unit, mint_url = await recieve_token(cashu_token)
original_amount = amount
original_unit = unit
logger.info( logger.info(
"credit_balance: Token redeemed successfully", "credit_balance: Token redeemed successfully",
extra={"amount": amount, "unit": unit, "mint_url": mint_url}, extra={"amount": amount, "unit": unit, "mint_url": mint_url},
@@ -310,6 +313,19 @@ async def credit_balance(
extra={"new_balance": key.balance}, extra={"new_balance": key.balance},
) )
try:
await store_cashu_transaction(
token=cashu_token,
amount=original_amount,
unit=original_unit,
mint_url=mint_url,
typ="in",
source="apikey",
api_key_hashed_key=key.hashed_key,
)
except Exception:
pass
logger.info( logger.info(
"Cashu token successfully redeemed and stored", "Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url}, extra={"amount": amount, "unit": unit, "mint_url": mint_url},
+37 -1
View File
@@ -13,7 +13,7 @@ import pytest
from httpx import AsyncClient from httpx import AsyncClient
from sqlmodel import select from sqlmodel import select
from routstr.core.db import ApiKey from routstr.core.db import ApiKey, CashuTransaction
@pytest.mark.integration @pytest.mark.integration
@@ -394,6 +394,42 @@ async def test_refund_during_active_usage(
assert response.json()["balance"] == 0 assert response.json()["balance"] == 0
@pytest.mark.integration
@pytest.mark.asyncio
async def test_wallet_history_returns_apikey_transactions(
authenticated_client: AsyncClient,
testmint_wallet: Any,
integration_session: Any,
) -> None:
wallet_response = await authenticated_client.get("/v1/wallet/")
api_key = wallet_response.json()["api_key"]
hashed_key = api_key[3:] if api_key.startswith("sk-") else api_key
topup_token = await testmint_wallet.mint_tokens(250)
topup_response = await authenticated_client.post(
"/v1/wallet/topup", params={"cashu_token": topup_token}
)
assert topup_response.status_code == 200
refund_response = await authenticated_client.post("/v1/wallet/refund")
assert refund_response.status_code == 200
history_response = await authenticated_client.get("/v1/wallet/history")
assert history_response.status_code == 200
transactions = history_response.json()["transactions"]
assert len(transactions) >= 2
assert all("api_key_hashed_key" not in tx for tx in transactions)
assert {tx["type"] for tx in transactions} >= {"in", "out"}
db_result = await integration_session.execute(
select(CashuTransaction).where(
CashuTransaction.api_key_hashed_key == hashed_key
)
)
db_transactions = db_result.scalars().all()
assert len(db_transactions) >= 2
@pytest.mark.integration @pytest.mark.integration
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mint_unavailability_handling( async def test_mint_unavailability_handling(
+11 -1
View File
@@ -11,7 +11,7 @@ import pytest
from httpx import AsyncClient from httpx import AsyncClient
from sqlmodel import select from sqlmodel import select
from routstr.core.db import ApiKey from routstr.core.db import ApiKey, CashuTransaction
from .utils import ( from .utils import (
CashuTokenGenerator, CashuTokenGenerator,
@@ -71,6 +71,16 @@ async def test_topup_with_valid_token( # type: ignore[no-untyped-def]
assert db_key.balance == new_balance assert db_key.balance == new_balance
assert db_key.balance == initial_balance + (topup_amount * 1000) assert db_key.balance == initial_balance + (topup_amount * 1000)
tx_result = await integration_session.execute(
select(CashuTransaction).where(
CashuTransaction.token == token,
CashuTransaction.type == "in",
)
)
tx = tx_result.scalar_one()
assert tx.api_key_hashed_key == hashed_key
assert tx.source == "apikey"
@pytest.mark.integration @pytest.mark.integration
@pytest.mark.asyncio @pytest.mark.asyncio
+107
View File
@@ -6,6 +6,7 @@ from fastapi.responses import JSONResponse
from routstr.balance import refund_wallet_endpoint from routstr.balance import refund_wallet_endpoint
from routstr.core.db import ApiKey, CashuTransaction from routstr.core.db import ApiKey, CashuTransaction
from routstr.wallet import credit_balance
def _make_cashu_tx( def _make_cashu_tx(
@@ -29,6 +30,12 @@ def _exec_result(tx: CashuTransaction | None) -> MagicMock:
return result return result
def _update_result(rowcount: int) -> MagicMock:
result = MagicMock()
result.rowcount = rowcount
return result
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_refund_x_cashu_returns_token() -> None: async def test_refund_x_cashu_returns_token() -> None:
x_cashu_token = "cashuAtest_token_value" x_cashu_token = "cashuAtest_token_value"
@@ -161,6 +168,7 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No
refund_token = "cashuArefund_apikey_token" refund_token = "cashuArefund_apikey_token"
session = MagicMock() session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
@@ -186,6 +194,7 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No
assert call_kwargs["source"] == "apikey" assert call_kwargs["source"] == "apikey"
assert call_kwargs["token"] == refund_token assert call_kwargs["token"] == refund_token
assert call_kwargs["typ"] == "out" assert call_kwargs["typ"] == "out"
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -194,6 +203,7 @@ async def test_apikey_refund_logs_token() -> None:
refund_token = "cashuAlogged_token" refund_token = "cashuAlogged_token"
session = MagicMock() session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
@@ -222,6 +232,7 @@ async def test_apikey_refund_log_includes_path() -> None:
refund_token = "cashuApath_token" refund_token = "cashuApath_token"
session = MagicMock() session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
@@ -248,3 +259,99 @@ async def test_apikey_refund_log_includes_path() -> None:
assert len(token_issued_calls) == 1 assert len(token_issued_calls) == 1
extra = token_issued_calls[0].kwargs.get("extra", {}) extra = token_issued_calls[0].kwargs.get("extra", {})
assert extra.get("path") == "/v1/wallet/refund" assert extra.get("path") == "/v1/wallet/refund"
@pytest.mark.asyncio
async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
"""When the debit CAS fails (rowcount=0), no token is minted and 409 is returned."""
from fastapi import HTTPException
key = _make_api_key(balance=5000, refund_currency="sat")
session = MagicMock()
# Debit returns rowcount=0 → balance changed concurrently
session.exec = AsyncMock(return_value=_update_result(0))
session.commit = AsyncMock()
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", mock_send_token),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 409
# Crucially: send_token must NOT have been called
mock_send_token.assert_not_awaited()
@pytest.mark.asyncio
async def test_credit_balance_stores_apikey_transaction_history() -> None:
key = _make_api_key(balance=1000)
session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.commit = AsyncMock()
session.refresh = AsyncMock()
with (
patch(
"routstr.wallet.recieve_token",
AsyncMock(return_value=(100, "sat", "https://mint.example")),
),
patch("routstr.wallet.store_cashu_transaction", AsyncMock()) as mock_store,
):
amount = await credit_balance("cashuAtopup_token", key, session)
assert amount == 100_000
mock_store.assert_awaited_once()
call_kwargs = mock_store.call_args.kwargs
assert call_kwargs["typ"] == "in"
assert call_kwargs["source"] == "apikey"
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
assert call_kwargs["amount"] == 100
assert call_kwargs["unit"] == "sat"
assert call_kwargs["token"] == "cashuAtopup_token"
assert call_kwargs["mint_url"] == "https://mint.example"
@pytest.mark.asyncio
async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
"""When debit succeeds but minting fails, balance must be restored."""
from fastapi import HTTPException
key = _make_api_key(balance=5000, refund_currency="sat")
# First exec call = debit (succeeds), second = restore
session = MagicMock()
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 503
# Verify two exec calls: debit + restore
assert session.exec.await_count == 2
+29 -3
View File
@@ -92,6 +92,7 @@ function TransactionTable({
<TableHead>Type</TableHead> <TableHead>Type</TableHead>
<TableHead>Amount</TableHead> <TableHead>Amount</TableHead>
<TableHead>Status</TableHead> <TableHead>Status</TableHead>
<TableHead>API Key</TableHead>
<TableHead>Request ID</TableHead> <TableHead>Request ID</TableHead>
<TableHead>Mint</TableHead> <TableHead>Mint</TableHead>
<TableHead>Date</TableHead> <TableHead>Date</TableHead>
@@ -115,6 +116,31 @@ function TransactionTable({
{tx.amount} {tx.unit} {tx.amount} {tx.unit}
</TableCell> </TableCell>
<TableCell>{getStatusBadge(tx)}</TableCell> <TableCell>{getStatusBadge(tx)}</TableCell>
<TableCell>
{tx.api_key_hashed_key ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[120px] truncate font-mono'>
{tx.api_key_hashed_key.slice(0, 12)}...
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
onCopy(tx.api_key_hashed_key!, tx.id + '-apikey')
}
>
{copiedId === tx.id + '-apikey' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'>—</span>
)}
</TableCell>
<TableCell> <TableCell>
{tx.request_id ? ( {tx.request_id ? (
<div className='flex items-center gap-1 text-xs'> <div className='flex items-center gap-1 text-xs'>
@@ -325,7 +351,7 @@ export default function TransactionsPage() {
<Search className='text-muted-foreground absolute top-2.5 left-2.5 h-4 w-4' /> <Search className='text-muted-foreground absolute top-2.5 left-2.5 h-4 w-4' />
<Input <Input
id='search' id='search'
placeholder='Search by ID, token or request ID...' placeholder='Search by ID, token, request ID or key hash...'
className='pl-8' className='pl-8'
value={search} value={search}
onChange={(e) => setSearch(e.target.value)} onChange={(e) => setSearch(e.target.value)}
@@ -385,7 +411,7 @@ export default function TransactionsPage() {
</TabsTrigger> </TabsTrigger>
<TabsTrigger value='apikey' className='flex items-center gap-2'> <TabsTrigger value='apikey' className='flex items-center gap-2'>
<Key className='h-4 w-4' /> <Key className='h-4 w-4' />
API Key Refunds API Key
{data && ( {data && (
<Badge variant='secondary' className='ml-1'> <Badge variant='secondary' className='ml-1'>
{apikeyTxs.length} {apikeyTxs.length}
@@ -416,7 +442,7 @@ export default function TransactionsPage() {
<Card> <Card>
<CardHeader> <CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'> <div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>API Key Refund History</CardTitle> <CardTitle>API Key Transaction History</CardTitle>
{hasActiveFilters && ( {hasActiveFilters && (
<CardDescription> <CardDescription>
Filtered by {activeFilterDescription} Filtered by {activeFilterDescription}
+1
View File
@@ -1136,6 +1136,7 @@ export interface Transaction {
collected: boolean; collected: boolean;
swept: boolean; swept: boolean;
source: 'x-cashu' | 'apikey'; source: 'x-cashu' | 'apikey';
api_key_hashed_key?: string;
} }
export interface TransactionsResponse { export interface TransactionsResponse {