Files
routstr-core/tests/unit/test_stale_reservations.py
T
redshift b4632db79c fix(refund): give refund refusals machine-readable error codes
/v1/wallet/refund answered three distinct refusals with a bare `detail`
string: "No balance to refund", "Balance too small to refund" and "Cannot
refund key. There are ongoing requests for this api key." They are
indistinguishable to a client, and the consequences differ: the first proves
the key holds nothing and its stored copy can be dropped, the second is dust
no retry can pay out, and the third is a transient race whose balance is
still on the key and must be kept.

The SDK tried to tell them apart by matching the whole error string, but the
refund error it compares against is built as "API key refund failed:
<detail>", so its no-balance branch never matched. Every dead key stayed in
storage and was re-swept every five minutes, forever.

All three now carry the structured envelope the other refund errors already
use ({"error": {"type", "code", "message"}}), with codes
`no_balance_to_refund`, `balance_too_small_to_refund` and
`refund_ongoing_requests`. Messages and status codes are unchanged, so
message-matching clients (including SDK <= 0.4.8) behave exactly as before.
2026-10-04 15:29:21 +08:00

520 lines
16 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
error = exc_info.value.detail["error"]
assert error["code"] == "refund_ongoing_requests"
assert "ongoing requests" in error["message"]
@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)