fix: make reservation cleanup atomic

This commit is contained in:
9qeklajc
2026-07-18 14:48:02 +02:00
parent fa0b366f9a
commit 2c218cce49
3 changed files with 84 additions and 9 deletions
+24 -6
View File
@@ -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(
+1 -2
View File
@@ -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={
@@ -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()