Merge remote-tracking branch 'origin/main' into feat/postgresql-compatibility

# Conflicts:
#	routstr/core/db.py
This commit is contained in:
9qeklajc
2026-09-30 22:56:25 +02:00
13 changed files with 628 additions and 8 deletions
+6
View File
@@ -72,6 +72,12 @@ ROUTSTR_SECRET_KEY=
# UPSTREAM_POOL_TIMEOUT=5 # UPSTREAM_POOL_TIMEOUT=5
# UPSTREAM_READ_TIMEOUT=900 # UPSTREAM_READ_TIMEOUT=900
# Request and reservation lifetime limits (seconds)
# STALE_RESERVATION_TIMEOUT_SECONDS=300
# MAX_REQUEST_LIFETIME_SECONDS=1800
# DOWNSTREAM_SEND_TIMEOUT_SECONDS=60
# REQUEST_CLEANUP_TIMEOUT_SECONDS=30
# Logging # Logging
# LOG_LEVEL=INFO # LOG_LEVEL=INFO
# ENABLE_CONSOLE_LOGGING=true # ENABLE_CONSOLE_LOGGING=true
@@ -0,0 +1,42 @@
"""Immutable reservation start and absolute recovery deadline.
Revision ID: a73d19b6c204
Revises: e4c7a1b9d520
"""
import time
import sqlalchemy as sa
from alembic import op
revision = "a73d19b6c204"
down_revision = "e4c7a1b9d520"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"reservation_releases", sa.Column("started_at", sa.Integer(), nullable=True)
)
op.add_column(
"reservation_releases", sa.Column("expires_at", sa.Integer(), nullable=True)
)
op.create_index(
"ix_reservation_releases_expires_at", "reservation_releases", ["expires_at"]
)
# Original ages are unknowable for renewed legacy rows. Give them a finite
# migration grace period; deploy only after draining old workers.
op.execute(
sa.text(
"UPDATE reservation_releases SET expires_at = :expiry WHERE status = 'active'"
).bindparams(expiry=int(time.time()) + 1830)
)
def downgrade() -> None:
op.drop_index(
"ix_reservation_releases_expires_at", table_name="reservation_releases"
)
op.drop_column("reservation_releases", "expires_at")
op.drop_column("reservation_releases", "started_at")
@@ -1,7 +1,7 @@
"""widen money and timestamp columns to 64-bit """widen money and timestamp columns to 64-bit
Revision ID: d3c8b21f7a04 Revision ID: d3c8b21f7a04
Revises: e4c7a1b9d520 Revises: a73d19b6c204
Create Date: 2026-09-29 00:00:00.000000 Create Date: 2026-09-29 00:00:00.000000
Balances are millisatoshis and every clock column is a unix timestamp, so both Balances are millisatoshis and every clock column is a unix timestamp, so both
@@ -24,7 +24,7 @@ import sqlalchemy as sa
from alembic import op from alembic import op
revision = "d3c8b21f7a04" revision = "d3c8b21f7a04"
down_revision = "e4c7a1b9d520" down_revision = "a73d19b6c204"
branch_labels = None branch_labels = None
depends_on = None depends_on = None
@@ -60,6 +60,8 @@ WIDENED_COLUMNS: tuple[tuple[str, str, bool], ...] = (
("refunds", "updated_at", False), ("refunds", "updated_at", False),
("reservation_releases", "reserved_msats", False), ("reservation_releases", "reserved_msats", False),
("reservation_releases", "created_at", False), ("reservation_releases", "created_at", False),
("reservation_releases", "started_at", True),
("reservation_releases", "expires_at", True),
("routstr_fees", "accumulated_msats", False), ("routstr_fees", "accumulated_msats", False),
("routstr_fees", "total_paid_msats", False), ("routstr_fees", "total_paid_msats", False),
("routstr_fees", "payout_in_progress_msats", False), ("routstr_fees", "payout_in_progress_msats", False),
+3 -1
View File
@@ -73,10 +73,12 @@ build-backend = "setuptools.build_meta"
[tool.setuptools] [tool.setuptools]
packages = ["routstr"] packages = ["routstr"]
[tool.ruff]
extend-exclude = ["examples"]
[tool.ruff.lint] [tool.ruff.lint]
select = ["E", "F", "I"] select = ["E", "F", "I"]
ignore = ["E501"] ignore = ["E501"]
exclude = ["examples"]
[tool.mypy] [tool.mypy]
python_version = "3.11" python_version = "3.11"
+33 -1
View File
@@ -635,6 +635,14 @@ async def pay_for_request(
) )
# Charge the base cost for the request atomically to avoid race conditions # Charge the base cost for the request atomically to avoid race conditions
from .core.lifecycle import request_lifetime
lifetime = request_lifetime.get()
remaining_lifetime = (
max(0, lifetime.deadline - asyncio.get_running_loop().time())
if lifetime is not None
else settings.max_request_lifetime_seconds
)
reserved_at_now = int(time.time()) reserved_at_now = int(time.time())
stmt = ( stmt = (
update(ApiKey) update(ApiKey)
@@ -684,6 +692,13 @@ async def pay_for_request(
billing_key_hash=reservation.billing_key_hash, billing_key_hash=reservation.billing_key_hash,
reserved_msats=reservation.reserved_msats, reserved_msats=reservation.reserved_msats,
status="active", status="active",
started_at=reserved_at_now,
# reserved_at_now floors to the second; add 1s margin so a
# finalizer finishing right at the nominal deadline isn't fenced
# out by truncation.
expires_at=reserved_at_now
+ math.ceil(remaining_lifetime + settings.request_cleanup_timeout_seconds)
+ 1,
) )
) )
# Publish the identity before commit. If the commit succeeds but its # Publish the identity before commit. If the commit succeeds but its
@@ -873,6 +888,10 @@ async def renew_reservation(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id) .where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(
(col(ReservationRelease.expires_at).is_(None))
| (col(ReservationRelease.expires_at) > int(time.time()))
)
.values(created_at=int(time.time())) .values(created_at=int(time.time()))
) )
await session.commit() await session.commit()
@@ -900,12 +919,21 @@ def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None:
""" """
interval = max(1, settings.stale_reservation_timeout_seconds // 3) interval = max(1, settings.stale_reservation_timeout_seconds // 3)
owner = asyncio.current_task() owner = asyncio.current_task()
from .core.lifecycle import request_lifetime
lifetime = request_lifetime.get()
deadline = asyncio.get_running_loop().time() + settings.max_request_lifetime_seconds
async def beat() -> None: async def beat() -> None:
try: try:
while True: while True:
await asyncio.sleep(interval) await asyncio.sleep(interval)
if owner is None or owner.done(): if (
owner is None
or owner.done()
or (lifetime is not None and lifetime.stopped)
or asyncio.get_running_loop().time() >= deadline
):
# Request control is gone; let the lease expire so the # Request control is gone; let the lease expire so the
# sweeper can release the reservation if no terminal # sweeper can release the reservation if no terminal
# transition ever ran. # transition ever ran.
@@ -1089,6 +1117,10 @@ async def _claim_reservation_for_charge(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id) .where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(
col(ReservationRelease.expires_at).is_(None)
| (col(ReservationRelease.expires_at) > int(time.time()))
)
.where(col(ReservationRelease.key_hash) == snapshot.key_hash) .where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash) .where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats) .where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
+10 -2
View File
@@ -203,7 +203,10 @@ async def _transition_stale_reservation(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == reservation_id) .where(col(ReservationRelease.id) == reservation_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff) .where(
(col(ReservationRelease.created_at) < cutoff)
| (col(ReservationRelease.expires_at) <= int(time.time()))
)
.values(status="released") .values(status="released")
) )
return bool(transition.rowcount == 1) return bool(transition.rowcount == 1)
@@ -250,7 +253,10 @@ async def release_stale_reservations(
query = ( query = (
select(ReservationRelease) select(ReservationRelease)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff) .where(
(col(ReservationRelease.created_at) < cutoff)
| (col(ReservationRelease.expires_at) <= int(time.time()))
)
) )
if key_hash is not None: if key_hash is not None:
query = query.where( query = query.where(
@@ -829,6 +835,8 @@ class ReservationRelease(SQLModel, table=True): # type: ignore
key_hash: str = Field(index=True) key_hash: str = Field(index=True)
billing_key_hash: str = Field(index=True) billing_key_hash: str = Field(index=True)
reserved_msats: int = Field(sa_type=Msats) reserved_msats: int = Field(sa_type=Msats)
started_at: int | None = Field(default=None, sa_type=UnixTimestamp)
expires_at: int | None = Field(default=None, index=True, sa_type=UnixTimestamp)
status: str = Field(default="active") status: str = Field(default="active")
created_at: int = Field( created_at: int = Field(
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
+134
View File
@@ -0,0 +1,134 @@
"""Supervise the real downstream connection, outside HTTP middleware wrappers."""
from __future__ import annotations
import asyncio
from contextvars import ContextVar
from dataclasses import dataclass
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from . import get_logger
from .settings import settings
logger = get_logger(__name__)
class DownstreamTerminated(OSError):
"""Raised by downstream_send after disconnect; expected, not a server error."""
@dataclass
class RequestLifetime:
deadline: float = 0
stopped: bool = False
request_lifetime: ContextVar[RequestLifetime | None] = ContextVar(
"request_lifetime", default=None
)
class RequestLifecycleMiddleware:
def __init__(self, app: ASGIApp) -> None:
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
lifetime = RequestLifetime(
deadline=asyncio.get_running_loop().time()
+ settings.max_request_lifetime_seconds
)
token = request_lifetime.set(lifetime)
disconnected = asyncio.Event()
# One receive consumer. Backpressure uploads until consumed; after the
# final body message, continue listening independently of the app.
messages: asyncio.Queue[Message] = asyncio.Queue(maxsize=1)
response_started = False
async def pump() -> None:
while True:
message = await receive()
if message["type"] == "http.disconnect":
disconnected.set()
return
await messages.put(message)
async def downstream_receive() -> Message:
if disconnected.is_set():
return {"type": "http.disconnect"}
get = asyncio.create_task(messages.get())
gone = asyncio.create_task(disconnected.wait())
try:
await asyncio.wait((get, gone), return_when=asyncio.FIRST_COMPLETED)
if disconnected.is_set():
return {"type": "http.disconnect"}
return get.result()
finally:
for task in (get, gone):
task.cancel()
await asyncio.gather(get, gone, return_exceptions=True)
async def downstream_send(message: Message) -> None:
nonlocal response_started
if disconnected.is_set() or lifetime.stopped:
raise DownstreamTerminated("Downstream request terminated")
async with asyncio.timeout(settings.downstream_send_timeout_seconds):
await send(message)
if message["type"] == "http.response.start":
response_started = True
receiver = asyncio.create_task(pump())
work: asyncio.Future[None] = asyncio.ensure_future(
self.app(scope, downstream_receive, downstream_send)
)
gone = asyncio.create_task(disconnected.wait())
timed_out = False
try:
done, _ = await asyncio.wait(
(work, gone),
timeout=settings.max_request_lifetime_seconds,
return_when=asyncio.FIRST_COMPLETED,
)
if gone in done and not response_started and work not in done:
# A pre-response wallet or billing operation may have accepted
# funds already. Let it reach its own settlement before closing.
done, _ = await asyncio.wait(
(work,),
timeout=max(0, lifetime.deadline - asyncio.get_running_loop().time()),
)
if work in done:
try:
await work
except DownstreamTerminated:
if not disconnected.is_set():
raise
logger.debug("Client disconnected before response completed")
elif not disconnected.is_set():
timed_out = True
finally:
lifetime.stopped = True
for task in (receiver, gone, work):
task.cancel()
# Detached stream finalizers own settlement. The heartbeat stops
# with the request; durable expiry recovers any abandoned row.
done, pending = await asyncio.wait(
(receiver, gone, work), timeout=settings.request_cleanup_timeout_seconds
)
for task in done:
if not task.cancelled():
task.exception()
for task in pending:
task.cancel()
task.add_done_callback(
lambda t: t.exception() if not t.cancelled() else None
)
request_lifetime.reset(token)
if timed_out and work.done() and not disconnected.is_set() and not response_started:
async with asyncio.timeout(settings.downstream_send_timeout_seconds):
await send({"type": "http.response.start", "status": 504, "headers": []})
await send(
{"type": "http.response.body", "body": b"Request deadline exceeded"}
)
+5
View File
@@ -45,6 +45,7 @@ from .exceptions import (
http_exception_handler, http_exception_handler,
validation_exception_handler, validation_exception_handler,
) )
from .lifecycle import RequestLifecycleMiddleware
from .logging import get_logger, setup_logging from .logging import get_logger, setup_logging
from .middleware import LoggingMiddleware from .middleware import LoggingMiddleware
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401 from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
@@ -315,6 +316,10 @@ app.add_middleware(
# Add logging middleware # Add logging middleware
app.add_middleware(LoggingMiddleware) app.add_middleware(LoggingMiddleware)
# Outermost: observe the actual downstream connection, not middleware streams.
app.add_middleware(RequestLifecycleMiddleware)
# Add exception handlers # Add exception handlers
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
app.add_exception_handler(RequestValidationError, validation_exception_handler) app.add_exception_handler(RequestValidationError, validation_exception_handler)
+10
View File
@@ -117,6 +117,16 @@ class Settings(BaseSettings):
default=604_800, env="DEAD_KEY_MIN_AGE_SECONDS" default=604_800, env="DEAD_KEY_MIN_AGE_SECONDS"
) )
max_request_lifetime_seconds: float = Field(
default=1800, gt=0, env="MAX_REQUEST_LIFETIME_SECONDS"
)
downstream_send_timeout_seconds: float = Field(
default=60, gt=0, env="DOWNSTREAM_SEND_TIMEOUT_SECONDS"
)
request_cleanup_timeout_seconds: float = Field(
default=30, gt=0, env="REQUEST_CLEANUP_TIMEOUT_SECONDS"
)
# Network # Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
# Comma-separated METHOD:path pairs adding to the proxy's canonical # Comma-separated METHOD:path pairs adding to the proxy's canonical
+6 -1
View File
@@ -10,6 +10,7 @@ from fastapi.responses import StreamingResponse
from starlette.types import Receive, Scope, Send from starlette.types import Receive, Scope, Send
from ..core import get_logger from ..core import get_logger
from ..core.settings import settings
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -74,10 +75,14 @@ class PersistentStreamFinalizer:
self._task: asyncio.Future[None] | None = None self._task: asyncio.Future[None] | None = None
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
async def _bounded_finalize(self) -> None:
async with asyncio.timeout(settings.request_cleanup_timeout_seconds):
await self._finalize()
async def run(self) -> None: async def run(self) -> None:
async with self._lock: async with self._lock:
if self._task is None: if self._task is None:
self._task = asyncio.ensure_future(self._finalize()) self._task = asyncio.ensure_future(self._bounded_finalize())
task = self._task task = self._task
await asyncio.shield(task) await asyncio.shield(task)
@@ -133,6 +133,8 @@ def test_money_columns_are_64_bit_on_postgresql(
("lightning_invoices", "paid_at"), ("lightning_invoices", "paid_at"),
("refunds", "claimed_at"), ("refunds", "claimed_at"),
("reservation_releases", "created_at"), ("reservation_releases", "created_at"),
("reservation_releases", "expires_at"),
("reservation_releases", "started_at"),
("routstr_fees", "payout_started_at"), ("routstr_fees", "payout_started_at"),
], ],
) )
+285
View File
@@ -0,0 +1,285 @@
import asyncio
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from pathlib import Path
from unittest.mock import patch
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.pool import NullPool
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import PlainTextResponse
from starlette.routing import Route
from starlette.types import Message, Receive, Scope, Send
import routstr.core.db as db_module
from routstr.auth import (
ReservationSnapshot,
_claim_reservation_for_charge,
_stop_reservation_heartbeat,
pay_for_request,
)
from routstr.core.db import ApiKey, ReservationRelease
from routstr.core.lifecycle import RequestLifecycleMiddleware
from routstr.core.middleware import LoggingMiddleware
from routstr.core.settings import settings
@pytest.mark.asyncio
@pytest.mark.parametrize("reason", ["disconnect", "deadline", "send"])
async def test_lifecycle_stops_live_work(reason: str) -> None:
closed = asyncio.Event()
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
sent: list[Message] = []
async def app(scope: Scope, receive: Receive, send: Send) -> None:
try:
assert (await receive())["type"] == "http.request"
await send({"type": "http.response.start", "status": 200, "headers": []})
while True:
await send(
{"type": "http.response.body", "body": b"x", "more_body": True}
)
await asyncio.sleep(0.01)
finally:
closed.set()
async def send(message: Message) -> None:
sent.append(message)
if reason == "send" and message["type"] == "http.response.body":
await asyncio.sleep(100)
async def disconnect() -> None:
await asyncio.sleep(0.02)
await receive_queue.put({"type": "http.disconnect"})
task = asyncio.create_task(disconnect()) if reason == "disconnect" else None
with (
patch.object(settings, "max_request_lifetime_seconds", 0.08),
patch.object(settings, "downstream_send_timeout_seconds", 0.03),
patch.object(settings, "request_cleanup_timeout_seconds", 0.1),
):
try:
await asyncio.wait_for(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, send
),
1,
)
except TimeoutError:
assert reason == "send"
if task:
await task
assert closed.is_set()
assert sent
@pytest.mark.asyncio
async def test_unrelated_oserror_still_propagates() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
raise OSError("Connection reset by peer")
async def send(message: Message) -> None:
pass
with pytest.raises(OSError, match="Connection reset by peer"):
await asyncio.wait_for(
RequestLifecycleMiddleware(app)({"type": "http"}, receive_queue.get, send),
1,
)
@pytest.mark.asyncio
async def test_disconnect_before_headers_preserves_wallet_work() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
entered = asyncio.Event()
finish_wallet = asyncio.Event()
wallet_credited = asyncio.Event()
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
entered.set()
await finish_wallet.wait() # The mint accepted the token; credit is still pending.
wallet_credited.set()
await send({"type": "http.response.start", "status": 200, "headers": []})
run = asyncio.create_task(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, lambda message: asyncio.sleep(0)
)
)
await asyncio.wait_for(entered.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.sleep(0.02)
assert not run.done()
finish_wallet.set()
await asyncio.wait_for(run, 1) # No propagated exception for an expected disconnect.
assert wallet_credited.is_set()
@pytest.mark.asyncio
async def test_disconnect_before_headers_with_logging_middleware() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
entered = asyncio.Event()
finish_wallet = asyncio.Event()
wallet_credited = asyncio.Event()
async def wallet(request: Request) -> PlainTextResponse:
await request.body()
entered.set()
await finish_wallet.wait()
wallet_credited.set()
return PlainTextResponse("settled")
app = RequestLifecycleMiddleware(
LoggingMiddleware(Starlette(routes=[Route("/wallet", wallet, methods=["POST"])]))
)
scope: Scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"http_version": "1.1",
"method": "POST",
"scheme": "http",
"path": "/wallet",
"raw_path": b"/wallet",
"root_path": "",
"query_string": b"",
"headers": [],
"client": ("test", 1234),
"server": ("test", 80),
}
async def send(message: Message) -> None:
pass
run = asyncio.create_task(app(scope, receive_queue.get, send))
try:
await asyncio.wait_for(entered.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.sleep(0.02)
assert not run.done()
finish_wallet.set()
await asyncio.wait_for(run, 1) # No propagated exception for an expected disconnect.
assert wallet_credited.is_set()
finally:
finish_wallet.set()
if not run.done():
run.cancel()
await asyncio.gather(run, return_exceptions=True)
@pytest.mark.asyncio
async def test_deadline_cancels_app_before_sending_504() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
sent: list[Message] = []
app_stopped = asyncio.Event()
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
try:
await asyncio.sleep(100)
finally:
with pytest.raises(OSError, match="Downstream request terminated"):
await send({"type": "http.response.start", "status": 200, "headers": []})
app_stopped.set()
async def send(message: Message) -> None:
assert app_stopped.is_set()
sent.append(message)
with (
patch.object(settings, "max_request_lifetime_seconds", 0.02),
patch.object(settings, "request_cleanup_timeout_seconds", 0.1),
):
await asyncio.wait_for(
RequestLifecycleMiddleware(app)({"type": "http"}, receive_queue.get, send),
1,
)
assert [message["type"] for message in sent] == [
"http.response.start",
"http.response.body",
]
assert sent[0]["status"] == 504
@pytest.mark.asyncio
async def test_disconnect_does_not_release_before_stream_settles(tmp_path: Path) -> None:
engine = create_async_engine(
f"sqlite+aiosqlite:///{tmp_path / 'reservations.db'}", poolclass=NullPool
)
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
@asynccontextmanager
async def session() -> AsyncGenerator[AsyncSession, None]:
async with AsyncSession(engine, expire_on_commit=False) as db:
yield db
with patch.object(db_module, "create_session", session):
async with session() as db:
db.add(ApiKey(hashed_key="stream-key", balance=10_000))
await db.commit()
started = asyncio.Event()
finalizer_started = asyncio.Event()
settle = asyncio.Event()
result: asyncio.Future[bool] = asyncio.get_running_loop().create_future()
snapshot: ReservationSnapshot | None = None
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
async def app(scope: Scope, receive: Receive, send: Send) -> None:
nonlocal snapshot
async with session() as db:
key = await db.get(ApiKey, "stream-key")
assert key is not None
snapshot = await pay_for_request(key, 1000, db)
await receive()
await send({"type": "http.response.start", "status": 200, "headers": []})
started.set()
try:
await asyncio.sleep(100)
finally:
async def finalize() -> None:
assert snapshot is not None
finalizer_started.set()
await settle.wait()
async with session() as db:
claimed = await _claim_reservation_for_charge(snapshot, db)
await db.commit()
await _stop_reservation_heartbeat(snapshot.release_id)
result.set_result(claimed)
asyncio.create_task(finalize())
run = asyncio.create_task(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, lambda message: asyncio.sleep(0)
)
)
try:
await asyncio.wait_for(started.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.wait_for(finalizer_started.wait(), 1)
await asyncio.wait_for(run, 1)
settle.set()
assert await asyncio.wait_for(result, 1)
assert snapshot is not None
async with session() as db:
row = await db.get(ReservationRelease, snapshot.release_id)
assert row is not None and row.status == "charged"
finally:
settle.set()
if snapshot is not None:
await _stop_reservation_heartbeat(snapshot.release_id)
await engine.dispose()
+88 -1
View File
@@ -9,6 +9,7 @@ Covers:
""" """
import asyncio import asyncio
import math
import time import time
from typing import AsyncGenerator from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
@@ -69,7 +70,6 @@ async def test_pay_for_request_sets_reserved_at(
payments_info = MagicMock() payments_info = MagicMock()
monkeypatch.setattr(auth_module.logger, "info", logger_info) monkeypatch.setattr(auth_module.logger, "info", logger_info)
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info) monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
before = int(time.time()) before = int(time.time())
await pay_for_request(key, 1_000, session) await pay_for_request(key, 1_000, session)
@@ -87,6 +87,34 @@ async def test_pay_for_request_sets_reserved_at(
assert payments_info.call_args.args == ("RESERVE",) assert payments_info.call_args.args == ("RESERVE",)
@pytest.mark.asyncio
async def test_pay_for_request_expires_at_has_floor_margin(
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
) -> None:
"""reserved_at_now floors to the second; expires_at must add 1s so a
finalizer finishing exactly at the nominal deadline isn't fenced out."""
key = ApiKey(hashed_key="floorkey", balance=10_000)
session.add(key)
await session.commit()
fixed_time = 1_700_000_000.9 # fractional second, floors when int()'d
monkeypatch.setattr(auth_module.time, "time", lambda: fixed_time)
snapshot = await pay_for_request(key, 1_000, session)
row = await session.get(ReservationRelease, snapshot.release_id)
assert row is not None
expected = (
int(fixed_time)
+ math.ceil(
auth_module.settings.max_request_lifetime_seconds
+ auth_module.settings.request_cleanup_timeout_seconds
)
+ 1
)
assert row.expires_at == expected
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pay_for_request_releases_reservation_when_validation_fails( async def test_pay_for_request_releases_reservation_when_validation_fails(
@@ -428,3 +456,62 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
await proxy_module.proxy(request, "v1/chat/completions") await proxy_module.proxy(request, "v1/chat/completions")
revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot) revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot)
@pytest.mark.asyncio
async def test_absolute_expiry_releases_fresh_lease(session: AsyncSession) -> None:
now = int(time.time())
key = ApiKey(
hashed_key="expired-deadline",
balance=5000,
reserved_balance=1000,
reserved_at=now,
)
session.add(key)
session.add(
ReservationRelease(
id="expired",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=1000,
created_at=now,
started_at=now - 100,
expires_at=now - 1,
)
)
await session.commit()
assert await release_stale_reservations(session, 300) == 1
await session.refresh(key)
assert key.reserved_balance == 0
assert key.balance == 5000
@pytest.mark.asyncio
async def test_expired_reservation_cannot_renew_or_claim_charge(
session: AsyncSession,
) -> None:
from routstr.auth import (
ReservationSnapshot,
_claim_reservation_for_charge,
renew_reservation,
)
snapshot = ReservationSnapshot(
release_id="fenced",
key_hash="fenced-key",
billing_key_hash="fenced-key",
reserved_msats=1000,
)
session.add(ApiKey(hashed_key="fenced-key", balance=5000, reserved_balance=1000))
session.add(
ReservationRelease(
id="fenced",
key_hash="fenced-key",
billing_key_hash="fenced-key",
reserved_msats=1000,
expires_at=int(time.time()) - 1,
)
)
await session.commit()
assert not await renew_reservation(snapshot, session)
assert not await _claim_reservation_for_charge(snapshot, session)