Compare commits

...

7 Commits

Author SHA1 Message Date
9qeklajc
4defe4f227 fix: identify reservation releases 2026-07-18 14:59:57 +02:00
9qeklajc
f8125a8a2d Merge branch 'fix/fee-payout-crash-guard' into fix/streaming-billing-finalization 2026-07-18 14:57:15 +02:00
9qeklajc
999a5634fa fix: snapshot reservation cleanup state 2026-07-18 14:54:06 +02:00
9qeklajc
2c218cce49 fix: make reservation cleanup atomic 2026-07-18 14:48:02 +02:00
9qeklajc
fa0b366f9a fix: fail safely on streaming billing errors 2026-07-18 14:42:03 +02:00
9qeklajc
be5e323c68 test: cover payout restart and migration safety 2026-07-18 14:13:20 +02:00
9qeklajc
f6d1a41728 fix: checkpoint fee payouts before sending 2026-07-18 14:08:41 +02:00
9 changed files with 677 additions and 22 deletions

View File

@@ -0,0 +1,37 @@
"""add fee payout checkpoint
Revision ID: d7e8f9a0b1c2
Revises: c6d7e8f9a0b1
Create Date: 2026-07-18 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "d7e8f9a0b1c2"
down_revision = "c6d7e8f9a0b1"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"routstr_fees",
sa.Column(
"payout_in_progress_msats",
sa.Integer(),
nullable=False,
server_default="0",
),
)
op.add_column(
"routstr_fees",
sa.Column("payout_started_at", sa.Integer(), nullable=True),
)
def downgrade() -> None:
op.drop_column("routstr_fees", "payout_started_at")
op.drop_column("routstr_fees", "payout_in_progress_msats")

View File

@@ -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")

View File

@@ -3,16 +3,24 @@ import hashlib
import math
import random
import time
import uuid
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
from fastapi import HTTPException
from sqlalchemy import case
from sqlalchemy.exc import IntegrityError
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,
@@ -766,6 +774,95 @@ async def revert_pay_for_request(
return True
@dataclass(frozen=True)
class ReservationSnapshot:
release_id: str
key_hash: str
billing_key_hash: str
async def get_reservation_snapshot(
key: ApiKey, session: AsyncSession
) -> ReservationSnapshot:
"""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,
billing_key_hash=billing_key.hashed_key,
)
def _reservation_release_statement(
key_hash: str,
reserved_msats: int,
) -> Update:
return (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.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,
),
)
)
async def release_reservation(
snapshot: ReservationSnapshot,
session: AsyncSession,
reserved_msats: int,
) -> bool:
"""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,
reserved_msats,
)
result = await session.exec(release_stmt) # type: ignore[call-overload]
if result.rowcount != 1:
await session.rollback()
return False
if snapshot.billing_key_hash != snapshot.key_hash:
child_release_stmt = _reservation_release_statement(
snapshot.key_hash,
reserved_msats,
)
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 True
async def adjust_payment_for_tokens(
key: ApiKey,
response_data: dict,

View File

@@ -348,12 +348,24 @@ 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)
accumulated_msats: int = Field(default=0)
total_paid_msats: int = Field(default=0)
last_paid_at: int | None = Field(default=None)
payout_in_progress_msats: int = Field(default=0)
payout_started_at: int | None = Field(default=None)
class CliToken(SQLModel, table=True): # type: ignore
@@ -394,18 +406,42 @@ async def get_routstr_fee(session: AsyncSession) -> RoutstrFee:
return fee
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> None:
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool:
"""Checkpoint a fee payout before making the external payment."""
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.where(col(RoutstrFee.payout_in_progress_msats) == 0)
.where(col(RoutstrFee.accumulated_msats) >= paid_msats)
.values(
accumulated_msats=RoutstrFee.accumulated_msats - paid_msats,
payout_in_progress_msats=paid_msats,
payout_started_at=int(time.time()),
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
return result.rowcount == 1
async def complete_routstr_fee_payout(
session: AsyncSession, paid_msats: int
) -> bool:
"""Mark a checkpointed payout complete after the external payment succeeds."""
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.where(col(RoutstrFee.payout_in_progress_msats) == paid_msats)
.values(
payout_in_progress_msats=0,
payout_started_at=None,
total_paid_msats=RoutstrFee.total_paid_msats + paid_msats,
last_paid_at=int(time.time()),
)
)
await session.exec(stmt) # type: ignore[call-overload]
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
return result.rowcount == 1
async def balances_for_mint_and_unit(

View File

@@ -14,7 +14,11 @@ from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from pydantic.v1 import BaseModel
from ..auth import adjust_payment_for_tokens
from ..auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
release_reservation,
)
from ..core import get_logger
from ..core.db import (
ApiKey,
@@ -996,6 +1000,9 @@ class BaseUpstreamProvider:
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
reservation_snapshot = await get_reservation_snapshot(
fresh_key, session
)
cost_data: dict
try:
adjustment_input = (
@@ -1013,25 +1020,40 @@ class BaseUpstreamProvider:
max_cost_for_model,
)
usage_finalized = True
except Exception as e:
logger.exception(
"Error during usage finalization",
except BaseException as e:
logger.critical(
"Error during usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
# Fall back so we still emit a non-zero sats cost downstream.
cost_data = {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
}
try:
await session.rollback()
released = await release_reservation(
reservation_snapshot,
session,
max_cost_for_model,
)
if not released:
logger.critical(
"Billing reservation could not be released",
extra={
"key_hash": key.hashed_key[:8] + "...",
"reserved_balance": fresh_key.reserved_balance,
},
)
except Exception as release_error:
logger.critical(
"Billing reservation release failed",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(release_error),
},
exc_info=True,
)
raise
if usage_chunk_data is None:
if not hasattr(self, "_current_stream_id"):

View File

@@ -1067,17 +1067,56 @@ async def periodic_routstr_fee_payout() -> None:
try:
async with db.create_session() as session:
fee = await db.get_routstr_fee(session)
if fee.payout_in_progress_msats:
logger.critical(
"Routstr fee payout requires manual reconciliation",
extra={
"payout_in_progress_msats": fee.payout_in_progress_msats,
"payout_started_at": fee.payout_started_at,
},
)
continue
accumulated_sats = fee.accumulated_msats // 1000
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True
)
amount_received = await raw_send_to_lnurl(
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
)
paid_msats = accumulated_sats * 1000
await db.reset_routstr_fee(session, paid_msats)
payout_checkpointed = await db.reset_routstr_fee(
session, paid_msats
)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
)
except Exception:
logger.critical(
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=True,
)
continue
payout_completed = await db.complete_routstr_fee_payout(
session, paid_msats
)
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
logger.info(
"Routstr fee payout sent",
extra={

View File

@@ -0,0 +1,160 @@
import asyncio
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr import wallet
from routstr.core import db
@asynccontextmanager
async def _session_context(session: Mock) -> AsyncIterator[Mock]:
yield session
@pytest.mark.asyncio
async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
async with AsyncSession(engine) as session:
session.add(db.RoutstrFee(id=1, accumulated_msats=5_000))
await session.commit()
assert await db.reset_routstr_fee(session, 5_000) is True
assert await db.reset_routstr_fee(session, 5_000) is False
fee = await db.get_routstr_fee(session)
await session.refresh(fee)
assert fee.accumulated_msats == 0
assert fee.payout_in_progress_msats == 5_000
assert fee.total_paid_msats == 0
assert await db.complete_routstr_fee_payout(session, 5_000) is True
await session.refresh(fee)
assert fee.payout_in_progress_msats == 0
assert fee.total_paid_msats == 5_000
assert fee.last_paid_at is not None
await engine.dispose()
@pytest.mark.asyncio
async def test_fee_payout_checkpoints_before_sending() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
payout_wallet = Mock()
events: list[str] = []
async def checkpoint(*_args: object) -> bool:
events.append("checkpoint")
return True
async def send(*_args: object, **_kwargs: object) -> int:
events.append("send")
return 5
async def complete(*_args: object) -> bool:
events.append("complete")
return True
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
assert events == ["checkpoint", "send", "complete"]
@pytest.mark.asyncio
async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=10_000,
payout_in_progress_msats=5_000,
payout_started_at=123,
)
with (
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint,
patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet,
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock()) as send,
patch("routstr.wallet.logger.critical") as critical,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
checkpoint.assert_not_awaited()
get_wallet.assert_not_awaited()
send.assert_not_awaited()
critical.assert_called_once()
@pytest.mark.asyncio
async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
complete = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch(
"routstr.wallet.raw_send_to_lnurl",
AsyncMock(side_effect=TimeoutError("unknown outcome")),
),
patch("routstr.wallet.logger.critical") as critical,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
complete.assert_not_awaited()
critical.assert_called_once()

View File

@@ -0,0 +1,46 @@
import os
import sqlite3
import subprocess
import sys
from pathlib import Path
def _run_alembic(root: Path, database_url: str, revision: str) -> None:
env = os.environ.copy()
env["DATABASE_URL"] = database_url
subprocess.run(
[sys.executable, "-m", "alembic", "upgrade", revision],
cwd=root,
env=env,
check=True,
capture_output=True,
text=True,
)
def test_fee_payout_checkpoint_migration_preserves_existing_row(
tmp_path: Path,
) -> None:
root = Path(__file__).resolve().parents[2]
database_path = tmp_path / "migration.db"
database_url = f"sqlite+aiosqlite:///{database_path}"
_run_alembic(root, database_url, "c6d7e8f9a0b1")
with sqlite3.connect(database_path) as connection:
result = connection.execute(
"UPDATE routstr_fees SET accumulated_msats = 5000, "
"total_paid_msats = 1000, last_paid_at = 123 WHERE id = 1"
)
assert result.rowcount == 1
connection.commit()
_run_alembic(root, database_url, "head")
with sqlite3.connect(database_path) as connection:
row = connection.execute(
"SELECT accumulated_msats, total_paid_msats, last_paid_at, "
"payout_in_progress_msats, payout_started_at "
"FROM routstr_fees WHERE id = 1"
).fetchone()
assert row == (5000, 1000, 123, 0, None)

View File

@@ -0,0 +1,168 @@
from collections.abc import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import get_reservation_snapshot, release_reservation
from routstr.core.db import ApiKey
from routstr.upstream.base import BaseUpstreamProvider
@pytest.mark.asyncio
async def test_release_reservation_clears_reserved_balance() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
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, reserved_at=123
)
async with AsyncSession(engine, expire_on_commit=False) as session:
session.add(key)
await session.commit()
snapshot = await get_reservation_snapshot(key, session)
assert await release_reservation(snapshot, session, 500) is True
await session.refresh(key)
assert key.reserved_balance == 0
assert key.reserved_at is None
assert await release_reservation(snapshot, session, 500) is True
await engine.dispose()
@pytest.mark.asyncio
async def test_release_reservation_preserves_other_concurrent_reservations() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
key = ApiKey(
hashed_key="key", balance=1_000, reserved_balance=800, reserved_at=123
)
async with AsyncSession(engine, expire_on_commit=False) as session:
session.add(key)
await session.commit()
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 == 0
assert key.reserved_at is None
assert await release_reservation(first_snapshot, session, 400) is True
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()
snapshot = await get_reservation_snapshot(child, session)
assert await release_reservation(snapshot, 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()
snapshot = await get_reservation_snapshot(child, session)
assert await release_reservation(snapshot, session, 500) is False
await session.refresh(parent)
await session.refresh(child)
assert (parent.reserved_balance, child.reserved_balance) == (500, 100)
await engine.dispose()
@pytest.mark.asyncio
async def test_streaming_billing_error_releases_reservation_and_propagates() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
yield b"data: [DONE]\n\n"
upstream_response = MagicMock()
upstream_response.status_code = 200
upstream_response.headers = {"content-type": "text/event-stream"}
upstream_response.aiter_bytes = aiter_bytes
key = MagicMock(spec=ApiKey)
key.hashed_key = "test-key-hash"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.rollback = AsyncMock()
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
release = AsyncMock(return_value=True)
reservation_snapshot = MagicMock()
with (
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
AsyncMock(side_effect=SQLAlchemyError("database unavailable")),
),
patch(
"routstr.upstream.base.get_reservation_snapshot",
AsyncMock(return_value=reservation_snapshot),
),
patch("routstr.upstream.base.release_reservation", release),
patch("routstr.upstream.base.create_session", return_value=session_context),
):
response = await provider.handle_streaming_chat_completion(
response=upstream_response,
key=key,
max_cost_for_model=500,
background_tasks=MagicMock(),
)
with pytest.raises(SQLAlchemyError, match="database unavailable"):
async for _ in response.body_iterator:
pass
session.rollback.assert_awaited_once()
release.assert_awaited_once_with(reservation_snapshot, session, 500)