fix: identify reservation releases

This commit is contained in:
9qeklajc
2026-07-18 14:59:57 +02:00
parent f8125a8a2d
commit 4defe4f227
4 changed files with 106 additions and 43 deletions
@@ -0,0 +1,50 @@
"""add reservation release idempotency records
Revision ID: f9a0b1c2d3e4
Revises: d7e8f9a0b1c2
Create Date: 2026-07-18 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "f9a0b1c2d3e4"
down_revision = "d7e8f9a0b1c2"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"reservation_releases",
sa.Column("id", sa.String(), nullable=False),
sa.Column("key_hash", sa.String(), nullable=False),
sa.Column("billing_key_hash", sa.String(), nullable=False),
sa.Column("reserved_msats", sa.Integer(), nullable=False),
sa.Column("created_at", sa.Integer(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(
"ix_reservation_releases_key_hash",
"reservation_releases",
["key_hash"],
)
op.create_index(
"ix_reservation_releases_billing_key_hash",
"reservation_releases",
["billing_key_hash"],
)
def downgrade() -> None:
op.drop_index(
"ix_reservation_releases_billing_key_hash",
table_name="reservation_releases",
)
op.drop_index(
"ix_reservation_releases_key_hash",
table_name="reservation_releases",
)
op.drop_table("reservation_releases")
+38 -37
View File
@@ -3,6 +3,7 @@ import hashlib
import math
import random
import time
import uuid
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
@@ -14,7 +15,12 @@ from sqlalchemy.sql.dml import Update
from sqlmodel import col, select, update
from .core import get_logger
from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee
from .core.db import (
ApiKey,
AsyncSession,
ReservationRelease,
accumulate_routstr_fee,
)
from .core.settings import settings
from .payment.cost_calculation import (
CostData,
@@ -770,12 +776,9 @@ async def revert_pay_for_request(
@dataclass(frozen=True)
class ReservationSnapshot:
release_id: str
key_hash: str
key_reserved_balance: int
key_reserved_at: int | None
billing_key_hash: str
billing_reserved_balance: int
billing_reserved_at: int | None
async def get_reservation_snapshot(
@@ -784,40 +787,30 @@ async def get_reservation_snapshot(
"""Capture the reservation state used for idempotent cleanup."""
billing_key = await get_billing_key(key, session)
return ReservationSnapshot(
release_id=uuid.uuid4().hex,
key_hash=key.hashed_key,
key_reserved_balance=key.reserved_balance,
key_reserved_at=key.reserved_at,
billing_key_hash=billing_key.hashed_key,
billing_reserved_balance=billing_key.reserved_balance,
billing_reserved_at=billing_key.reserved_at,
)
def _reservation_release_statement(
key_hash: str,
expected_balance: int,
expected_reserved_at: int | None,
reserved_msats: int,
) -> Update:
stmt = (
return (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.where(col(ApiKey.reserved_balance) == expected_balance)
)
if expected_reserved_at is None:
stmt = stmt.where(col(ApiKey.reserved_at).is_(None))
else:
stmt = stmt.where(col(ApiKey.reserved_at) == expected_reserved_at)
remaining_balance = expected_balance - reserved_msats
next_reserved_at = (
None
if remaining_balance == 0
else (expected_reserved_at or int(time.time())) + 1
)
return stmt.values(
reserved_balance=remaining_balance,
reserved_at=next_reserved_at,
.where(col(ApiKey.reserved_balance) >= reserved_msats)
.values(
reserved_balance=col(ApiKey.reserved_balance) - reserved_msats,
reserved_at=case(
(
col(ApiKey.reserved_balance) - reserved_msats > 0,
col(ApiKey.reserved_at),
),
else_=None,
),
)
)
@@ -826,14 +819,27 @@ async def release_reservation(
session: AsyncSession,
reserved_msats: int,
) -> bool:
"""Release one snapshotted reservation exactly once without charging."""
if reserved_msats <= 0 or snapshot.billing_reserved_balance < reserved_msats:
"""Release one reservation exactly once without charging."""
if reserved_msats <= 0:
return False
session.add(
ReservationRelease(
id=snapshot.release_id,
key_hash=snapshot.key_hash,
billing_key_hash=snapshot.billing_key_hash,
reserved_msats=reserved_msats,
)
)
try:
await session.flush()
except IntegrityError:
await session.rollback()
existing = await session.get(ReservationRelease, snapshot.release_id)
return existing is not None and existing.reserved_msats == reserved_msats
release_stmt = _reservation_release_statement(
snapshot.billing_key_hash,
snapshot.billing_reserved_balance,
snapshot.billing_reserved_at,
reserved_msats,
)
result = await session.exec(release_stmt) # type: ignore[call-overload]
@@ -842,13 +848,8 @@ async def release_reservation(
return False
if snapshot.billing_key_hash != snapshot.key_hash:
if snapshot.key_reserved_balance < reserved_msats:
await session.rollback()
return False
child_release_stmt = _reservation_release_statement(
snapshot.key_hash,
snapshot.key_reserved_balance,
snapshot.key_reserved_at,
reserved_msats,
)
child_result = await session.exec( # type: ignore[call-overload]
+10
View File
@@ -348,6 +348,16 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
)
class ReservationRelease(SQLModel, table=True): # type: ignore
__tablename__ = "reservation_releases"
id: str = Field(primary_key=True)
key_hash: str = Field(index=True)
billing_key_hash: str = Field(index=True)
reserved_msats: int
created_at: int = Field(default_factory=lambda: int(time.time()))
class RoutstrFee(SQLModel, table=True): # type: ignore
__tablename__ = "routstr_fees"
id: int = Field(default=1, primary_key=True)
@@ -30,7 +30,7 @@ async def test_release_reservation_clears_reserved_balance() -> None:
await session.refresh(key)
assert key.reserved_balance == 0
assert key.reserved_at is None
assert await release_reservation(snapshot, session, 500) is False
assert await release_reservation(snapshot, session, 500) is True
await engine.dispose()
@@ -48,12 +48,14 @@ async def test_release_reservation_preserves_other_concurrent_reservations() ->
session.add(key)
await session.commit()
snapshot = await get_reservation_snapshot(key, session)
assert await release_reservation(snapshot, session, 500) is True
first_snapshot = await get_reservation_snapshot(key, session)
second_snapshot = await get_reservation_snapshot(key, session)
assert await release_reservation(first_snapshot, session, 400) is True
assert await release_reservation(second_snapshot, session, 400) is True
await session.refresh(key)
assert key.reserved_balance == 300
assert key.reserved_at == 124
assert await release_reservation(snapshot, session, 500) is False
assert key.reserved_balance == 0
assert key.reserved_at is None
assert await release_reservation(first_snapshot, session, 400) is True
await engine.dispose()