mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge remote-tracking branch 'origin/main' into feat/postgresql-compatibility
# Conflicts: # routstr/core/db.py
This commit is contained in:
@@ -72,6 +72,12 @@ ROUTSTR_SECRET_KEY=
|
||||
# UPSTREAM_POOL_TIMEOUT=5
|
||||
# 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
|
||||
# LOG_LEVEL=INFO
|
||||
# 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
|
||||
|
||||
Revision ID: d3c8b21f7a04
|
||||
Revises: e4c7a1b9d520
|
||||
Revises: a73d19b6c204
|
||||
Create Date: 2026-09-29 00:00:00.000000
|
||||
|
||||
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
|
||||
|
||||
revision = "d3c8b21f7a04"
|
||||
down_revision = "e4c7a1b9d520"
|
||||
down_revision = "a73d19b6c204"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -60,6 +60,8 @@ WIDENED_COLUMNS: tuple[tuple[str, str, bool], ...] = (
|
||||
("refunds", "updated_at", False),
|
||||
("reservation_releases", "reserved_msats", False),
|
||||
("reservation_releases", "created_at", False),
|
||||
("reservation_releases", "started_at", True),
|
||||
("reservation_releases", "expires_at", True),
|
||||
("routstr_fees", "accumulated_msats", False),
|
||||
("routstr_fees", "total_paid_msats", False),
|
||||
("routstr_fees", "payout_in_progress_msats", False),
|
||||
|
||||
+3
-1
@@ -73,10 +73,12 @@ build-backend = "setuptools.build_meta"
|
||||
[tool.setuptools]
|
||||
packages = ["routstr"]
|
||||
|
||||
[tool.ruff]
|
||||
extend-exclude = ["examples"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "I"]
|
||||
ignore = ["E501"]
|
||||
exclude = ["examples"]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.11"
|
||||
|
||||
+33
-1
@@ -635,6 +635,14 @@ async def pay_for_request(
|
||||
)
|
||||
|
||||
# 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())
|
||||
stmt = (
|
||||
update(ApiKey)
|
||||
@@ -684,6 +692,13 @@ async def pay_for_request(
|
||||
billing_key_hash=reservation.billing_key_hash,
|
||||
reserved_msats=reservation.reserved_msats,
|
||||
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
|
||||
@@ -873,6 +888,10 @@ async def renew_reservation(
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == snapshot.release_id)
|
||||
.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()))
|
||||
)
|
||||
await session.commit()
|
||||
@@ -900,12 +919,21 @@ def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None:
|
||||
"""
|
||||
interval = max(1, settings.stale_reservation_timeout_seconds // 3)
|
||||
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:
|
||||
try:
|
||||
while True:
|
||||
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
|
||||
# sweeper can release the reservation if no terminal
|
||||
# transition ever ran.
|
||||
@@ -1089,6 +1117,10 @@ async def _claim_reservation_for_charge(
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == snapshot.release_id)
|
||||
.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.billing_key_hash) == snapshot.billing_key_hash)
|
||||
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
|
||||
|
||||
+10
-2
@@ -203,7 +203,10 @@ async def _transition_stale_reservation(
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == reservation_id)
|
||||
.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")
|
||||
)
|
||||
return bool(transition.rowcount == 1)
|
||||
@@ -250,7 +253,10 @@ async def release_stale_reservations(
|
||||
query = (
|
||||
select(ReservationRelease)
|
||||
.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:
|
||||
query = query.where(
|
||||
@@ -829,6 +835,8 @@ class ReservationRelease(SQLModel, table=True): # type: ignore
|
||||
key_hash: str = Field(index=True)
|
||||
billing_key_hash: str = Field(index=True)
|
||||
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")
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
|
||||
|
||||
@@ -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"}
|
||||
)
|
||||
@@ -45,6 +45,7 @@ from .exceptions import (
|
||||
http_exception_handler,
|
||||
validation_exception_handler,
|
||||
)
|
||||
from .lifecycle import RequestLifecycleMiddleware
|
||||
from .logging import get_logger, setup_logging
|
||||
from .middleware import LoggingMiddleware
|
||||
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
|
||||
@@ -315,6 +316,10 @@ app.add_middleware(
|
||||
# Add logging middleware
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
# Outermost: observe the actual downstream connection, not middleware streams.
|
||||
|
||||
app.add_middleware(RequestLifecycleMiddleware)
|
||||
|
||||
# Add exception handlers
|
||||
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
|
||||
app.add_exception_handler(RequestValidationError, validation_exception_handler)
|
||||
|
||||
@@ -117,6 +117,16 @@ class Settings(BaseSettings):
|
||||
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
|
||||
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
|
||||
# Comma-separated METHOD:path pairs adding to the proxy's canonical
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastapi.responses import StreamingResponse
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -74,10 +75,14 @@ class PersistentStreamFinalizer:
|
||||
self._task: asyncio.Future[None] | None = None
|
||||
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 with self._lock:
|
||||
if self._task is None:
|
||||
self._task = asyncio.ensure_future(self._finalize())
|
||||
self._task = asyncio.ensure_future(self._bounded_finalize())
|
||||
task = self._task
|
||||
await asyncio.shield(task)
|
||||
|
||||
|
||||
@@ -133,6 +133,8 @@ def test_money_columns_are_64_bit_on_postgresql(
|
||||
("lightning_invoices", "paid_at"),
|
||||
("refunds", "claimed_at"),
|
||||
("reservation_releases", "created_at"),
|
||||
("reservation_releases", "expires_at"),
|
||||
("reservation_releases", "started_at"),
|
||||
("routstr_fees", "payout_started_at"),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -9,6 +9,7 @@ Covers:
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import time
|
||||
from typing import AsyncGenerator
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
@@ -69,7 +70,6 @@ async def test_pay_for_request_sets_reserved_at(
|
||||
payments_info = MagicMock()
|
||||
monkeypatch.setattr(auth_module.logger, "info", logger_info)
|
||||
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
|
||||
|
||||
before = int(time.time())
|
||||
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",)
|
||||
|
||||
|
||||
@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
|
||||
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")
|
||||
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user