diff --git a/migrations/versions/b5e7c9d1f3a2_add_reserved_at_to_api_keys.py b/migrations/versions/b5e7c9d1f3a2_add_reserved_at_to_api_keys.py new file mode 100644 index 00000000..c658044a --- /dev/null +++ b/migrations/versions/b5e7c9d1f3a2_add_reserved_at_to_api_keys.py @@ -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") diff --git a/routstr/auth.py b/routstr/auth.py index b3f353d2..62a16da6 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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}") diff --git a/routstr/balance.py b/routstr/balance.py index 41de5fc8..7453aae7 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -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() diff --git a/routstr/core/db.py b/routstr/core/db.py index 0028234e..0ce92f6b 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -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) diff --git a/routstr/core/main.py b/routstr/core/main.py index d9fb00eb..2fca2338 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -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: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index ab8db07d..8faf190b 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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: diff --git a/routstr/proxy.py b/routstr/proxy.py index 373b80b2..5340c078 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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", diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py new file mode 100644 index 00000000..8c7d7c8f --- /dev/null +++ b/tests/unit/test_stale_reservations.py @@ -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)