From 2c218cce492612731ef1bf7748a291a64e02c0bc Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 18 Jul 2026 14:48:02 +0200 Subject: [PATCH] fix: make reservation cleanup atomic --- routstr/auth.py | 30 ++++++++-- routstr/upstream/base.py | 3 +- .../test_streaming_billing_finalization.py | 60 ++++++++++++++++++- 3 files changed, 84 insertions(+), 9 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 1655eea2..26526e94 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -773,25 +773,43 @@ async def release_reservation( ) -> bool: """Release a request reservation without charging the key.""" billing_key = await get_billing_key(key, session) + cleared_reserved_at = case( + (col(ApiKey.reserved_balance) - reserved_msats > 0, col(ApiKey.reserved_at)), + else_=None, + ) release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .where(col(ApiKey.reserved_balance) >= reserved_msats) - .values(reserved_balance=col(ApiKey.reserved_balance) - reserved_msats) + .where(col(ApiKey.reserved_balance) == reserved_msats) + .values( + reserved_balance=col(ApiKey.reserved_balance) - reserved_msats, + reserved_at=cleared_reserved_at, + ) ) result = await session.exec(release_stmt) # type: ignore[call-overload] + if result.rowcount != 1: + await session.rollback() + return False if billing_key.hashed_key != key.hashed_key: child_release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.reserved_balance) >= reserved_msats) - .values(reserved_balance=col(ApiKey.reserved_balance) - reserved_msats) + .where(col(ApiKey.reserved_balance) == reserved_msats) + .values( + reserved_balance=col(ApiKey.reserved_balance) - reserved_msats, + reserved_at=cleared_reserved_at, + ) ) - await session.exec(child_release_stmt) # type: ignore[call-overload] + child_result = await session.exec( # type: ignore[call-overload] + child_release_stmt + ) + if child_result.rowcount != 1: + await session.rollback() + return False await session.commit() - return result.rowcount == 1 + return True async def adjust_payment_for_tokens( diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 9b2ecdfa..7ed09f04 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -13,7 +13,6 @@ import httpx from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel -from sqlalchemy.exc import SQLAlchemyError from ..auth import adjust_payment_for_tokens, release_reservation from ..core import get_logger @@ -1014,7 +1013,7 @@ class BaseUpstreamProvider: max_cost_for_model, ) usage_finalized = True - except (HTTPException, SQLAlchemyError) as e: + except BaseException as e: logger.critical( "Error during usage finalization — CRITICAL", extra={ diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index d599b268..ff3a4232 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -18,7 +18,9 @@ async def test_release_reservation_clears_reserved_balance() -> None: async with engine.begin() as connection: await connection.run_sync(SQLModel.metadata.create_all) - key = ApiKey(hashed_key="key", balance=1_000, reserved_balance=500) + key = ApiKey( + hashed_key="key", balance=1_000, reserved_balance=500, reserved_at=123 + ) async with AsyncSession(engine, expire_on_commit=False) as session: session.add(key) await session.commit() @@ -26,6 +28,62 @@ async def test_release_reservation_clears_reserved_balance() -> None: assert await release_reservation(key, session, 500) is True await session.refresh(key) assert key.reserved_balance == 0 + assert key.reserved_at is None + assert await release_reservation(key, session, 500) is False + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_release_reservation_updates_parent_and_child_atomically() -> None: + engine = create_async_engine("sqlite+aiosqlite://") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + + parent = ApiKey( + hashed_key="parent", balance=1_000, reserved_balance=500, reserved_at=123 + ) + child = ApiKey( + hashed_key="child", + parent_key_hash="parent", + balance=0, + reserved_balance=500, + reserved_at=123, + ) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add_all([parent, child]) + await session.commit() + + assert await release_reservation(child, session, 500) is True + await session.refresh(parent) + await session.refresh(child) + assert (parent.reserved_balance, child.reserved_balance) == (0, 0) + assert (parent.reserved_at, child.reserved_at) == (None, None) + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_release_reservation_rolls_back_partial_parent_child_update() -> None: + engine = create_async_engine("sqlite+aiosqlite://") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + + parent = ApiKey(hashed_key="parent", balance=1_000, reserved_balance=500) + child = ApiKey( + hashed_key="child", + parent_key_hash="parent", + balance=0, + reserved_balance=100, + ) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add_all([parent, child]) + await session.commit() + + assert await release_reservation(child, session, 500) is False + await session.refresh(parent) + await session.refresh(child) + assert (parent.reserved_balance, child.reserved_balance) == (500, 100) await engine.dispose()