mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
518 lines
15 KiB
Python
518 lines
15 KiB
Python
"""Tests for stale reserved_balance handling (issue #551).
|
|
|
|
Covers:
|
|
- pay_for_request stamping reserved_at on charged keys
|
|
- release_stale_reservations sweeper semantics
|
|
- reset_all_reserved_balances clearing reserved_at
|
|
- refund endpoint self-healing stale/legacy reservations
|
|
- proxy reverting the reservation when the client disconnects (CancelledError)
|
|
"""
|
|
|
|
import asyncio
|
|
import math
|
|
import time
|
|
from typing import AsyncGenerator
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
|
from sqlalchemy.pool import StaticPool
|
|
from sqlmodel import SQLModel, select
|
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
|
|
|
import routstr.auth as auth_module
|
|
from routstr.auth import pay_for_request
|
|
from routstr.balance import refund_wallet_endpoint
|
|
from routstr.core.db import (
|
|
ApiKey,
|
|
ReservationRelease,
|
|
release_stale_reservations,
|
|
reset_all_reserved_balances,
|
|
)
|
|
|
|
from .proxy_test_utils import mock_request_stream, patch_proxy_session
|
|
|
|
|
|
def _make_engine() -> AsyncEngine:
|
|
return create_async_engine(
|
|
"sqlite+aiosqlite://",
|
|
poolclass=StaticPool,
|
|
connect_args={"check_same_thread": False},
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
async def session() -> "AsyncGenerator[AsyncSession, None]":
|
|
engine = _make_engine()
|
|
async with engine.begin() as conn:
|
|
await conn.run_sync(SQLModel.metadata.create_all)
|
|
db_session = AsyncSession(engine, expire_on_commit=False)
|
|
try:
|
|
yield db_session
|
|
finally:
|
|
await db_session.close()
|
|
await engine.dispose()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# pay_for_request stamps reserved_at
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pay_for_request_sets_reserved_at(
|
|
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
|
|
) -> None:
|
|
key = ApiKey(hashed_key="paykey", balance=10_000)
|
|
session.add(key)
|
|
await session.commit()
|
|
logger_info = MagicMock()
|
|
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)
|
|
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 1_000
|
|
assert key.reserved_at is not None
|
|
assert key.reserved_at >= before
|
|
success_logs = [
|
|
call
|
|
for call in logger_info.call_args_list
|
|
if call.args == ("Payment processed successfully",)
|
|
]
|
|
assert len(success_logs) == 1
|
|
payments_info.assert_called_once()
|
|
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(
|
|
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
|
|
) -> None:
|
|
key = ApiKey(hashed_key="invalid-reservation", balance=10_000)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
async def reject_reservation(*_args: object, **_kwargs: object) -> None:
|
|
raise RuntimeError("reservation identity changed")
|
|
|
|
logger_info = MagicMock()
|
|
payments_info = MagicMock()
|
|
monkeypatch.setattr(
|
|
auth_module, "_validate_reservation_snapshot", reject_reservation
|
|
)
|
|
monkeypatch.setattr(auth_module.logger, "info", logger_info)
|
|
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
|
|
|
|
with pytest.raises(RuntimeError, match="identity changed"):
|
|
await pay_for_request(key, 1_000, session)
|
|
|
|
assert not any(
|
|
call.args == ("Payment processed successfully",)
|
|
for call in logger_info.call_args_list
|
|
)
|
|
payments_info.assert_not_called()
|
|
|
|
await session.refresh(key)
|
|
release = (
|
|
await session.exec(
|
|
select(ReservationRelease).where(
|
|
ReservationRelease.key_hash == key.hashed_key
|
|
)
|
|
)
|
|
).one()
|
|
assert key.reserved_balance == 0
|
|
assert key.total_requests == 0
|
|
assert release.status == "released"
|
|
assert release.id not in auth_module._reservation_heartbeats
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_revert_clears_reserved_at_when_fully_released(
|
|
session: AsyncSession,
|
|
) -> None:
|
|
from routstr.auth import revert_pay_for_request
|
|
|
|
key = ApiKey(hashed_key="revertkey", balance=10_000)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
await pay_for_request(key, 1_000, session)
|
|
reverted = await revert_pay_for_request(key, session, 1_000)
|
|
|
|
assert reverted is True
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 0
|
|
assert key.reserved_at is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_revert_keeps_reserved_at_while_other_reservations_remain(
|
|
session: AsyncSession,
|
|
) -> None:
|
|
from routstr.auth import revert_pay_for_request
|
|
|
|
key = ApiKey(hashed_key="partialrevert", balance=10_000)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
await pay_for_request(key, 1_000, session)
|
|
await pay_for_request(key, 1_000, session)
|
|
reverted = await revert_pay_for_request(key, session, 1_000)
|
|
|
|
assert reverted is True
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 1_000
|
|
assert key.reserved_at is not None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# release_stale_reservations sweeper
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_release_stale_reservations_releases_old(session: AsyncSession) -> None:
|
|
key = ApiKey(
|
|
hashed_key="stalekey",
|
|
balance=5_000,
|
|
reserved_balance=1_000,
|
|
reserved_at=int(time.time()) - 1_000,
|
|
)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
released = await release_stale_reservations(session, max_age_seconds=300)
|
|
|
|
assert released == 1
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 0
|
|
assert key.reserved_at is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> None:
|
|
key = ApiKey(
|
|
hashed_key="freshkey",
|
|
balance=5_000,
|
|
reserved_balance=1_000,
|
|
reserved_at=int(time.time()),
|
|
)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
released = await release_stale_reservations(session, max_age_seconds=300)
|
|
|
|
assert released == 0
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 1_000
|
|
assert key.reserved_at is not None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_release_stale_reservations_skips_null_reserved_at(
|
|
session: AsyncSession,
|
|
) -> None:
|
|
# Reservations without a timestamp may belong to instances running older
|
|
# code (rolling deploy) — the background sweeper must not touch them.
|
|
key = ApiKey(
|
|
hashed_key="legacykey",
|
|
balance=5_000,
|
|
reserved_balance=1_000,
|
|
reserved_at=None,
|
|
)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
released = await release_stale_reservations(session, max_age_seconds=300)
|
|
|
|
assert released == 0
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 1_000
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reset_all_reserved_balances_clears_reserved_at(
|
|
session: AsyncSession,
|
|
) -> None:
|
|
key = ApiKey(
|
|
hashed_key="resetkey",
|
|
balance=5_000,
|
|
reserved_balance=1_000,
|
|
reserved_at=int(time.time()),
|
|
)
|
|
session.add(key)
|
|
await session.commit()
|
|
|
|
await reset_all_reserved_balances(session)
|
|
|
|
await session.refresh(key)
|
|
assert key.reserved_balance == 0
|
|
assert key.reserved_at is None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Refund endpoint self-healing
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def]
|
|
return (
|
|
patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)),
|
|
patch("routstr.refund.store_cashu_transaction", AsyncMock()),
|
|
)
|
|
|
|
|
|
async def _add_key(session: AsyncSession, **kwargs) -> ApiKey: # type: ignore[no-untyped-def]
|
|
key = ApiKey(refund_currency="sat", **kwargs)
|
|
session.add(key)
|
|
await session.commit()
|
|
return key
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_refund_self_heals_stale_reservation(session: AsyncSession) -> None:
|
|
key = await _add_key(
|
|
session,
|
|
hashed_key="stalerefund",
|
|
balance=5_000,
|
|
reserved_balance=2_000,
|
|
reserved_at=int(time.time()) - 10_000,
|
|
)
|
|
|
|
p1, p2 = _refund_patches()
|
|
with p1, p2:
|
|
result = await refund_wallet_endpoint(
|
|
authorization="Bearer sk-stalerefund",
|
|
x_cashu=None,
|
|
session=session,
|
|
)
|
|
|
|
assert isinstance(result, dict)
|
|
assert result["token"] == "cashuArefund"
|
|
# Full balance refunded (5000 msats -> 5 sats), reservation healed
|
|
assert result["sats"] == "5"
|
|
await session.refresh(key)
|
|
assert key.balance == 0
|
|
assert key.reserved_balance == 0
|
|
assert key.reserved_at is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_refund_self_heals_legacy_null_reserved_at(session: AsyncSession) -> None:
|
|
# Keys stuck from before reserved_at existed must be refundable.
|
|
key = await _add_key(
|
|
session,
|
|
hashed_key="legacyrefund",
|
|
balance=5_000,
|
|
reserved_balance=2_000,
|
|
reserved_at=None,
|
|
)
|
|
|
|
p1, p2 = _refund_patches()
|
|
with p1, p2:
|
|
result = await refund_wallet_endpoint(
|
|
authorization="Bearer sk-legacyrefund",
|
|
x_cashu=None,
|
|
session=session,
|
|
)
|
|
|
|
assert isinstance(result, dict)
|
|
assert result["token"] == "cashuArefund"
|
|
await session.refresh(key)
|
|
assert key.balance == 0
|
|
assert key.reserved_balance == 0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_refund_rejects_recent_reservation(session: AsyncSession) -> None:
|
|
from fastapi import HTTPException
|
|
|
|
await _add_key(
|
|
session,
|
|
hashed_key="activerefund",
|
|
balance=5_000,
|
|
reserved_balance=2_000,
|
|
reserved_at=int(time.time()),
|
|
)
|
|
|
|
p1, p2 = _refund_patches()
|
|
with p1, p2:
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await refund_wallet_endpoint(
|
|
authorization="Bearer sk-activerefund",
|
|
x_cashu=None,
|
|
session=session,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 400
|
|
assert "ongoing requests" in exc_info.value.detail
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_refund_without_reservation_still_works(session: AsyncSession) -> None:
|
|
key = await _add_key(
|
|
session,
|
|
hashed_key="plainrefund",
|
|
balance=5_000,
|
|
reserved_balance=0,
|
|
)
|
|
|
|
p1, p2 = _refund_patches()
|
|
with p1, p2:
|
|
result = await refund_wallet_endpoint(
|
|
authorization="Bearer sk-plainrefund",
|
|
x_cashu=None,
|
|
session=session,
|
|
)
|
|
|
|
assert isinstance(result, dict)
|
|
assert result["token"] == "cashuArefund"
|
|
await session.refresh(key)
|
|
assert key.balance == 0
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Proxy reverts reservation on client disconnect
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
|
from routstr import proxy as proxy_module
|
|
|
|
key = ApiKey(hashed_key="cancelkey", balance=10_000)
|
|
|
|
request = MagicMock()
|
|
request.method = "POST"
|
|
request.headers = {"authorization": "Bearer sk-cancelkey"}
|
|
mock_request_stream(request, b'{"model": "test-model"}')
|
|
|
|
upstream = MagicMock()
|
|
upstream.provider_type = "test"
|
|
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
|
|
upstream.forward_request = AsyncMock(side_effect=asyncio.CancelledError())
|
|
|
|
session = MagicMock()
|
|
reservation_snapshot = MagicMock()
|
|
revert_mock = AsyncMock(return_value=True)
|
|
|
|
with (
|
|
patch.object(
|
|
proxy_module,
|
|
"get_candidates",
|
|
return_value=[(MagicMock(), upstream)],
|
|
),
|
|
patch.object(
|
|
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
|
|
),
|
|
patch.object(
|
|
proxy_module,
|
|
"calculate_discounted_max_cost",
|
|
AsyncMock(return_value=1_000),
|
|
),
|
|
patch.object(proxy_module, "check_token_balance", MagicMock()),
|
|
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
|
|
patch.object(
|
|
proxy_module,
|
|
"pay_for_request",
|
|
AsyncMock(return_value=reservation_snapshot),
|
|
),
|
|
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
|
patch_proxy_session(session),
|
|
):
|
|
with pytest.raises(asyncio.CancelledError):
|
|
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)
|