mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
7 Commits
coverage-t
...
fix/stream
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4defe4f227 | ||
|
|
f8125a8a2d | ||
|
|
999a5634fa | ||
|
|
2c218cce49 | ||
|
|
fa0b366f9a | ||
|
|
be5e323c68 | ||
|
|
f6d1a41728 |
@@ -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")
|
||||
50
migrations/versions/f9a0b1c2d3e4_add_reservation_releases.py
Normal file
50
migrations/versions/f9a0b1c2d3e4_add_reservation_releases.py
Normal 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")
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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={
|
||||
|
||||
160
tests/unit/test_fee_payout_crash_safety.py
Normal file
160
tests/unit/test_fee_payout_crash_safety.py
Normal 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()
|
||||
46
tests/unit/test_fee_payout_migration.py
Normal file
46
tests/unit/test_fee_payout_migration.py
Normal 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)
|
||||
168
tests/unit/test_streaming_billing_finalization.py
Normal file
168
tests/unit/test_streaming_billing_finalization.py
Normal 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)
|
||||
Reference in New Issue
Block a user