diff --git a/.env.example b/.env.example index 5688f6c3..e79b7770 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/migrations/versions/a73d19b6c204_reservation_deadlines.py b/migrations/versions/a73d19b6c204_reservation_deadlines.py new file mode 100644 index 00000000..74702d61 --- /dev/null +++ b/migrations/versions/a73d19b6c204_reservation_deadlines.py @@ -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") diff --git a/pyproject.toml b/pyproject.toml index 296fab3f..7daf7489 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -72,10 +72,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" diff --git a/routstr/auth.py b/routstr/auth.py index 19b062e9..99c1d18e 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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) diff --git a/routstr/core/db.py b/routstr/core/db.py index 500420d5..84307a7e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -174,7 +174,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) @@ -221,7 +224,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( @@ -784,6 +790,8 @@ class ReservationRelease(SQLModel, table=True): # type: ignore key_hash: str = Field(index=True) billing_key_hash: str = Field(index=True) reserved_msats: int + started_at: int | None = Field(default=None) + expires_at: int | None = Field(default=None, index=True) status: str = Field(default="active") created_at: int = Field(default_factory=lambda: int(time.time())) diff --git a/routstr/core/lifecycle.py b/routstr/core/lifecycle.py new file mode 100644 index 00000000..039ac046 --- /dev/null +++ b/routstr/core/lifecycle.py @@ -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"} + ) diff --git a/routstr/core/main.py b/routstr/core/main.py index 584fa736..46aad1fc 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -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) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 5dd2d57d..d9324f13 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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 diff --git a/routstr/upstream/stream_ownership.py b/routstr/upstream/stream_ownership.py index e0e3e51d..ea06e400 100644 --- a/routstr/upstream/stream_ownership.py +++ b/routstr/upstream/stream_ownership.py @@ -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) diff --git a/tests/unit/test_request_lifecycle.py b/tests/unit/test_request_lifecycle.py new file mode 100644 index 00000000..1573f074 --- /dev/null +++ b/tests/unit/test_request_lifecycle.py @@ -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() diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 558fdda5..61a3da5a 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -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)