make key reservation reset predictable

This commit is contained in:
9qeklajc
2026-06-12 19:45:39 +02:00
parent a16bc1220c
commit 919dcf5535
8 changed files with 552 additions and 22 deletions
@@ -0,0 +1,25 @@
"""add reserved_at to api_keys
Revision ID: b5e7c9d1f3a2
Revises: a2b3c4d5e6f7
Create Date: 2026-06-12 00:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "b5e7c9d1f3a2"
down_revision = "a2b3c4d5e6f7"
branch_labels = None
depends_on = None
def upgrade() -> None:
# Nullable on purpose: existing keys keep NULL (no reservation recorded
# yet). New reservations populate it via pay_for_request.
op.add_column("api_keys", sa.Column("reserved_at", sa.Integer(), nullable=True))
def downgrade() -> None:
op.drop_column("api_keys", "reserved_at")
+38
View File
@@ -534,12 +534,14 @@ async def pay_for_request(
)
# Charge the base cost for the request atomically to avoid race conditions
reserved_at_now = int(time.time())
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request)
.values(
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
reserved_at=reserved_at_now,
total_requests=col(ApiKey.total_requests) + 1,
)
)
@@ -553,6 +555,7 @@ async def pay_for_request(
.values(
total_requests=col(ApiKey.total_requests) + 1,
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
reserved_at=reserved_at_now,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -620,12 +623,20 @@ async def revert_pay_for_request(
False if the reservation was already released (prevents negative reserved_balance)."""
billing_key = await get_billing_key(key, session)
# Keep reserved_at while other reservations remain; clear it once the
# reservation drains to zero so no stale-looking metadata lingers.
cleared_reserved_at = case(
(col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)),
else_=None,
)
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= cost_per_request)
.values(
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
reserved_at=cleared_reserved_at,
total_requests=col(ApiKey.total_requests) - 1,
)
)
@@ -641,6 +652,7 @@ async def revert_pay_for_request(
.values(
total_requests=col(ApiKey.total_requests) - 1,
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
reserved_at=cleared_reserved_at,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -1263,3 +1275,29 @@ async def periodic_key_reset() -> None:
break
except Exception as e:
logger.error(f"Error in periodic_key_reset: {e}")
STALE_RESERVATION_SWEEP_INTERVAL_SECONDS: int = 60
async def periodic_stale_reservation_sweep() -> None:
"""Background task that releases reservations leaked by client disconnects,
crashes or abandoned streams. Without it, a single interrupted request can
lock a key's balance (and block refunds) until the next process restart."""
from .core.db import create_session, release_stale_reservations
while True:
try:
await asyncio.sleep(STALE_RESERVATION_SWEEP_INTERVAL_SECONDS)
except asyncio.CancelledError:
break
try:
async with create_session() as session:
await release_stale_reservations(
session, settings.stale_reservation_timeout_seconds
)
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"Error in periodic_stale_reservation_sweep: {e}")
+35 -5
View File
@@ -7,7 +7,7 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from sqlmodel import col, select, update
from sqlmodel import col, or_, select, update
from .auth import get_billing_key, validate_bearer_key
from .core.db import (
@@ -313,9 +313,39 @@ async def refund_wallet_endpoint(
)
if key.reserved_balance > 0:
raise HTTPException(
status_code=400,
detail="Cannot refund key. There are ongoing requests for this api key.",
# Self-heal: reservations leaked by client disconnects or crashes would
# otherwise lock the user out of refunding forever. Release the
# reservation if it is stale (or predates reserved_at tracking) and
# proceed; only reject when a reservation is genuinely recent.
cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds
stale_release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) > 0)
.where(
or_(
col(ApiKey.reserved_at).is_(None),
col(ApiKey.reserved_at) < cutoff,
)
)
.values(reserved_balance=0, reserved_at=None)
)
stale_result = await session.exec(stale_release_stmt) # type: ignore[call-overload]
await session.commit()
if stale_result.rowcount == 0:
raise HTTPException(
status_code=400,
detail="Cannot refund key. There are ongoing requests for this api key.",
)
await session.refresh(key)
logger.warning(
"refund_wallet_endpoint: released stale reservation before refund",
extra={
"hashed_key": key.hashed_key,
"stale_timeout_seconds": settings.stale_reservation_timeout_seconds,
},
)
remaining_balance_msats: int = key.total_balance
@@ -342,7 +372,7 @@ async def refund_wallet_endpoint(
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) == pre_debit_balance)
.where(col(ApiKey.reserved_balance) == pre_debit_reserved)
.values(balance=0, reserved_balance=0)
.values(balance=0, reserved_balance=0, reserved_at=None)
)
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload]
await session.commit()
+38 -1
View File
@@ -33,6 +33,14 @@ class ApiKey(SQLModel, table=True): # type: ignore
reserved_balance: int = Field(
default=0, description="Reserved balance in millisatoshis (msats)"
)
reserved_at: int | None = Field(
default=None,
description=(
"Unix timestamp of the most recent balance reservation. Used to "
"detect and release stale reservations (e.g. after client "
"disconnects). NULL when no reservation has been made yet."
),
)
refund_address: str | None = Field(
default=None,
description="Lightning address to refund remaining balance after key expires",
@@ -87,12 +95,41 @@ class ApiKey(SQLModel, table=True): # type: ignore
async def reset_all_reserved_balances(session: AsyncSession) -> None:
stmt = update(ApiKey).values(reserved_balance=0)
stmt = update(ApiKey).values(reserved_balance=0, reserved_at=None)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
logger.info("Reset reserved balances on startup")
async def release_stale_reservations(
session: AsyncSession, max_age_seconds: int
) -> int:
"""Release reservations whose last reserve is older than max_age_seconds.
Only rows with a known reservation time are touched — NULL `reserved_at`
rows are left alone here so reservations made by instances running older
code (rolling deploys) are never killed mid-flight. Those rows are healed
on demand by the refund endpoint instead.
"""
cutoff = int(time.time()) - max_age_seconds
stmt = (
update(ApiKey)
.where(col(ApiKey.reserved_balance) > 0)
.where(col(ApiKey.reserved_at).is_not(None))
.where(col(ApiKey.reserved_at) < cutoff)
.values(reserved_balance=0, reserved_at=None)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
released = int(result.rowcount or 0)
if released:
logger.warning(
"Released stale balance reservations",
extra={"released_keys": released, "max_age_seconds": max_age_seconds},
)
return released
class ModelRow(SQLModel, table=True): # type: ignore
__tablename__ = "models"
id: str = Field(primary_key=True)
+9 -1
View File
@@ -11,7 +11,7 @@ from starlette.exceptions import HTTPException
from starlette.responses import Response as StarletteResponse
from starlette.types import Scope
from ..auth import periodic_key_reset
from ..auth import periodic_key_reset, periodic_stale_reservation_sweep
from ..balance import balance_router, deprecated_wallet_router
from ..lightning import lightning_router, periodic_invoice_watcher
from ..nostr import (
@@ -54,6 +54,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
models_refresh_task = None
model_maps_refresh_task = None
key_reset_task = None
stale_reservation_task = None
auto_topup_task = None
refund_sweep_task = None
routstr_fee_task = None
@@ -122,6 +123,9 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
if global_settings.providers_refresh_interval_seconds > 0:
providers_task = asyncio.create_task(providers_cache_refresher())
key_reset_task = asyncio.create_task(periodic_key_reset())
stale_reservation_task = asyncio.create_task(
periodic_stale_reservation_sweep()
)
auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
@@ -159,6 +163,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
model_maps_refresh_task.cancel()
if key_reset_task is not None:
key_reset_task.cancel()
if stale_reservation_task is not None:
stale_reservation_task.cancel()
if auto_topup_task is not None:
auto_topup_task.cancel()
if refund_sweep_task is not None:
@@ -188,6 +194,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(model_maps_refresh_task)
if key_reset_task is not None:
tasks_to_wait.append(key_reset_task)
if stale_reservation_task is not None:
tasks_to_wait.append(stale_reservation_task)
if auto_topup_task is not None:
tasks_to_wait.append(auto_topup_task)
if refund_sweep_task is not None:
+8 -1
View File
@@ -67,6 +67,12 @@ class Settings(BaseSettings):
reset_reserved_balance_on_startup: bool = Field(
default=True, env="RESET_RESERVED_BALANCE_ON_STARTUP"
) # deactivate in horizontal scaling setups
# Reservations older than this are considered leaked (client disconnect,
# crash, abandoned stream) and released by the background sweeper and the
# refund endpoint.
stale_reservation_timeout_seconds: int = Field(
default=300, env="STALE_RESERVATION_TIMEOUT_SECONDS"
)
# Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
@@ -167,7 +173,8 @@ def resolve_bootstrap() -> Settings:
pass
if not base.onion_url:
try:
from ..nostr.listing import discover_onion_url_from_tor # type: ignore
from ..nostr.listing import \
discover_onion_url_from_tor # type: ignore
discovered = discover_onion_url_from_tor()
if discovered:
+21 -14
View File
@@ -1,3 +1,4 @@
import asyncio
import json
from typing import Any
@@ -8,23 +9,14 @@ from sqlmodel import select
from .algorithm import create_model_mappings
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger
from .core.db import (
ApiKey,
AsyncSession,
ModelRow,
UpstreamProviderRow,
create_session,
get_session,
)
from .core.db import (ApiKey, AsyncSession, ModelRow, UpstreamProviderRow,
create_session, get_session)
from .core.exceptions import UpstreamError
from .core.not_found import build_not_found_response
from .core.settings import settings
from .payment.helpers import (
calculate_discounted_max_cost,
check_token_balance,
create_error_response,
get_max_cost_for_model,
)
from .payment.helpers import (calculate_discounted_max_cost,
check_token_balance, create_error_response,
get_max_cost_for_model)
from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.helpers import init_upstreams
@@ -492,6 +484,21 @@ async def proxy(
return response
except asyncio.CancelledError:
logger.warning(
"Client disconnected mid-request, reverting reservation",
extra={
"path": path,
"model": model_id,
"key_hash": key.hashed_key[:8] + "...",
"max_cost_for_model": max_cost_for_model,
},
)
await asyncio.shield(
revert_pay_for_request(key, session, max_cost_for_model)
)
raise
except UpstreamError as e:
logger.warning(
"Upstream %s failed for model=%s: %s",
+378
View File
@@ -0,0 +1,378 @@
"""Tests for stale reserved_balance handling (issue #551).
Covers:
- pay_for_request stamping reserved_at on billing and child keys
- release_stale_reservations sweeper semantics
- reset_all_reserved_balances clearing reserved_at
- refund endpoint self-healing stale/legacy reservations
- proxy reverting the reservation when the client disconnects (CancelledError)
"""
import asyncio
import time
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import pay_for_request
from routstr.balance import refund_wallet_endpoint
from routstr.core.db import (
ApiKey,
release_stale_reservations,
reset_all_reserved_balances,
)
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def session() -> "AsyncGenerator[AsyncSession, None]":
engine = _make_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
db_session = AsyncSession(engine, expire_on_commit=False)
try:
yield db_session
finally:
await db_session.close()
await engine.dispose()
# ---------------------------------------------------------------------------
# pay_for_request stamps reserved_at
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None:
key = ApiKey(hashed_key="paykey", balance=10_000)
session.add(key)
await session.commit()
before = int(time.time())
await pay_for_request(key, 1_000, session)
await session.refresh(key)
assert key.reserved_balance == 1_000
assert key.reserved_at is not None
assert key.reserved_at >= before
@pytest.mark.asyncio
async def test_pay_for_request_sets_reserved_at_on_child_key(session: AsyncSession) -> None:
parent = ApiKey(hashed_key="parentkey", balance=10_000)
child = ApiKey(hashed_key="childkey", balance=0, parent_key_hash="parentkey")
session.add(parent)
session.add(child)
await session.commit()
await pay_for_request(child, 1_000, session)
await session.refresh(parent)
await session.refresh(child)
assert parent.reserved_balance == 1_000
assert parent.reserved_at is not None
assert child.reserved_balance == 1_000
assert child.reserved_at is not None
@pytest.mark.asyncio
async def test_revert_clears_reserved_at_when_fully_released(
session: AsyncSession,
) -> None:
from routstr.auth import revert_pay_for_request
key = ApiKey(hashed_key="revertkey", balance=10_000)
session.add(key)
await session.commit()
await pay_for_request(key, 1_000, session)
reverted = await revert_pay_for_request(key, session, 1_000)
assert reverted is True
await session.refresh(key)
assert key.reserved_balance == 0
assert key.reserved_at is None
@pytest.mark.asyncio
async def test_revert_keeps_reserved_at_while_other_reservations_remain(
session: AsyncSession,
) -> None:
from routstr.auth import revert_pay_for_request
key = ApiKey(hashed_key="partialrevert", balance=10_000)
session.add(key)
await session.commit()
await pay_for_request(key, 1_000, session)
await pay_for_request(key, 1_000, session)
reverted = await revert_pay_for_request(key, session, 1_000)
assert reverted is True
await session.refresh(key)
assert key.reserved_balance == 1_000
assert key.reserved_at is not None
# ---------------------------------------------------------------------------
# release_stale_reservations sweeper
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_release_stale_reservations_releases_old(session: AsyncSession) -> None:
key = ApiKey(
hashed_key="stalekey",
balance=5_000,
reserved_balance=1_000,
reserved_at=int(time.time()) - 1_000,
)
session.add(key)
await session.commit()
released = await release_stale_reservations(session, max_age_seconds=300)
assert released == 1
await session.refresh(key)
assert key.reserved_balance == 0
assert key.reserved_at is None
@pytest.mark.asyncio
async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> None:
key = ApiKey(
hashed_key="freshkey",
balance=5_000,
reserved_balance=1_000,
reserved_at=int(time.time()),
)
session.add(key)
await session.commit()
released = await release_stale_reservations(session, max_age_seconds=300)
assert released == 0
await session.refresh(key)
assert key.reserved_balance == 1_000
assert key.reserved_at is not None
@pytest.mark.asyncio
async def test_release_stale_reservations_skips_null_reserved_at(session: AsyncSession) -> None:
# Reservations without a timestamp may belong to instances running older
# code (rolling deploy) — the background sweeper must not touch them.
key = ApiKey(
hashed_key="legacykey",
balance=5_000,
reserved_balance=1_000,
reserved_at=None,
)
session.add(key)
await session.commit()
released = await release_stale_reservations(session, max_age_seconds=300)
assert released == 0
await session.refresh(key)
assert key.reserved_balance == 1_000
@pytest.mark.asyncio
async def test_reset_all_reserved_balances_clears_reserved_at(session: AsyncSession) -> None:
key = ApiKey(
hashed_key="resetkey",
balance=5_000,
reserved_balance=1_000,
reserved_at=int(time.time()),
)
session.add(key)
await session.commit()
await reset_all_reserved_balances(session)
await session.refresh(key)
assert key.reserved_balance == 0
assert key.reserved_at is None
# ---------------------------------------------------------------------------
# Refund endpoint self-healing
# ---------------------------------------------------------------------------
def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def]
return (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_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()),
)
async def _add_key(session: AsyncSession, **kwargs) -> ApiKey: # type: ignore[no-untyped-def]
key = ApiKey(refund_currency="sat", **kwargs)
session.add(key)
await session.commit()
return key
@pytest.mark.asyncio
async def test_refund_self_heals_stale_reservation(session: AsyncSession) -> None:
key = await _add_key(
session,
hashed_key="stalerefund",
balance=5_000,
reserved_balance=2_000,
reserved_at=int(time.time()) - 10_000,
)
p1, p2, p3, p4 = _refund_patches()
with p1, p2, p3, p4:
result = await refund_wallet_endpoint(
authorization="Bearer sk-stalerefund",
x_cashu=None,
session=session,
)
assert isinstance(result, dict)
assert result["token"] == "cashuArefund"
# Full balance refunded (5000 msats -> 5 sats), reservation healed
assert result["sats"] == "5"
await session.refresh(key)
assert key.balance == 0
assert key.reserved_balance == 0
assert key.reserved_at is None
@pytest.mark.asyncio
async def test_refund_self_heals_legacy_null_reserved_at(session: AsyncSession) -> None:
# Keys stuck from before reserved_at existed must be refundable.
key = await _add_key(
session,
hashed_key="legacyrefund",
balance=5_000,
reserved_balance=2_000,
reserved_at=None,
)
p1, p2, p3, p4 = _refund_patches()
with p1, p2, p3, p4:
result = await refund_wallet_endpoint(
authorization="Bearer sk-legacyrefund",
x_cashu=None,
session=session,
)
assert isinstance(result, dict)
assert result["token"] == "cashuArefund"
await session.refresh(key)
assert key.balance == 0
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_refund_rejects_recent_reservation(session: AsyncSession) -> None:
from fastapi import HTTPException
await _add_key(
session,
hashed_key="activerefund",
balance=5_000,
reserved_balance=2_000,
reserved_at=int(time.time()),
)
p1, p2, p3, p4 = _refund_patches()
with p1, p2, p3, p4:
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-activerefund",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 400
assert "ongoing requests" in exc_info.value.detail
@pytest.mark.asyncio
async def test_refund_without_reservation_still_works(session: AsyncSession) -> None:
key = await _add_key(
session,
hashed_key="plainrefund",
balance=5_000,
reserved_balance=0,
)
p1, p2, p3, p4 = _refund_patches()
with p1, p2, p3, p4:
result = await refund_wallet_endpoint(
authorization="Bearer sk-plainrefund",
x_cashu=None,
session=session,
)
assert isinstance(result, dict)
assert result["token"] == "cashuArefund"
await session.refresh(key)
assert key.balance == 0
# ---------------------------------------------------------------------------
# Proxy reverts reservation on client disconnect
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
from routstr import proxy as proxy_module
key = ApiKey(hashed_key="cancelkey", balance=10_000)
request = MagicMock()
request.method = "POST"
request.headers = {"authorization": "Bearer sk-cancelkey"}
request.body = AsyncMock(return_value=b'{"model": "test-model"}')
upstream = MagicMock()
upstream.provider_type = "test"
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
upstream.forward_request = AsyncMock(side_effect=asyncio.CancelledError())
session = MagicMock()
revert_mock = AsyncMock(return_value=True)
with (
patch.object(proxy_module, "get_model_instance", return_value=MagicMock()),
patch.object(proxy_module, "get_provider_for_model", return_value=[upstream]),
patch.object(
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
),
patch.object(
proxy_module,
"calculate_discounted_max_cost",
AsyncMock(return_value=1_000),
),
patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(
proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)
),
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
):
with pytest.raises(asyncio.CancelledError):
await proxy_module.proxy(request, "v1/chat/completions", session=session)
revert_mock.assert_awaited_once_with(key, session, 1_000)