mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-30 15:26:14 +00:00
fix: make reservation cleanup atomic
This commit is contained in:
+24
-6
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user