Merge main into fix-reset-reserve-balance

Resolve pay_for_request conflict: keep main's atomic child-key
balance_limit guard and post-rowcount-check ordering, and stamp
reserved_at on both billing and child reservations.
This commit is contained in:
9qeklajc
2026-06-12 21:08:53 +02:00
7 changed files with 330 additions and 48 deletions
@@ -16,8 +16,8 @@ depends_on = None
def upgrade() -> None: def upgrade() -> None:
# Nullable on purpose: existing keys keep NULL (no reservation recorded # existing keys keep NULL
# yet). New reservations populate it via pay_for_request. # New reservations populate it via pay_for_request.
op.add_column("api_keys", sa.Column("reserved_at", sa.Integer(), nullable=True)) op.add_column("api_keys", sa.Column("reserved_at", sa.Integer(), nullable=True))
+45 -29
View File
@@ -547,21 +547,6 @@ async def pay_for_request(
) )
result = await session.exec(stmt) # type: ignore[call-overload] result = await session.exec(stmt) # type: ignore[call-overload]
# Also increment total_requests and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_requests=col(ApiKey.total_requests) + 1,
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
reserved_at=reserved_at_now,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0: if result.rowcount == 0:
logger.error( logger.error(
"Concurrent request depleted balance", "Concurrent request depleted balance",
@@ -573,7 +558,6 @@ async def pay_for_request(
}, },
) )
# Another concurrent request spent the balance first
raise HTTPException( raise HTTPException(
status_code=402, status_code=402,
detail={ detail={
@@ -585,6 +569,44 @@ async def pay_for_request(
}, },
) )
# Also increment total_requests and reserved_balance on the child key if it's different.
# The balance_limit guard is enforced atomically here — the Python pre-check above
# is a fast-path rejection only and provides no concurrency guarantee.
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(
(col(ApiKey.balance_limit).is_(None))
| (
col(ApiKey.total_spent)
+ col(ApiKey.reserved_balance)
+ cost_per_request
<= col(ApiKey.balance_limit)
)
)
.values(
total_requests=col(ApiKey.total_requests) + 1,
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
reserved_at=reserved_at_now,
)
)
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if child_result.rowcount == 0:
raise HTTPException(
status_code=402,
detail={
"error": {
"message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.",
"type": "insufficient_quota",
"code": "balance_limit_exceeded",
}
},
)
await session.commit()
await session.refresh(billing_key) await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key: if billing_key.hashed_key != key.hashed_key:
await session.refresh(key) await session.refresh(key)
@@ -623,8 +645,7 @@ async def revert_pay_for_request(
False if the reservation was already released (prevents negative reserved_balance).""" False if the reservation was already released (prevents negative reserved_balance)."""
billing_key = await get_billing_key(key, session) billing_key = await get_billing_key(key, session)
# Keep reserved_at while other reservations remain; clear it once the # Keep reserved_at while other reservations remain
# reservation drains to zero so no stale-looking metadata lingers.
cleared_reserved_at = case( cleared_reserved_at = case(
(col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)), (col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)),
else_=None, else_=None,
@@ -1282,22 +1303,17 @@ STALE_RESERVATION_SWEEP_INTERVAL_SECONDS: int = 60
async def periodic_stale_reservation_sweep() -> None: async def periodic_stale_reservation_sweep() -> None:
"""Background task that releases reservations leaked by client disconnects, """Background task that releases reservations leaked by client disconnects,
crashes or abandoned streams. Without it, a single interrupted request can crashes or abandoned streams.
lock a key's balance (and block refunds) until the next process restart.""" """
from .core.db import create_session, release_stale_reservations from .core.db import create_session, release_stale_reservations
while True: while True:
try:
await asyncio.sleep(STALE_RESERVATION_SWEEP_INTERVAL_SECONDS)
except asyncio.CancelledError:
break
try: try:
async with create_session() as session: async with create_session() as session:
await release_stale_reservations( await release_stale_reservations(
session, settings.stale_reservation_timeout_seconds session, settings.stale_reservation_timeout_seconds
) )
except asyncio.CancelledError: except Exception:
break logger.exception("Error in periodic_stale_reservation_sweep")
except Exception as e:
logger.error(f"Error in periodic_stale_reservation_sweep: {e}") await asyncio.sleep(STALE_RESERVATION_SWEEP_INTERVAL_SECONDS)
+1 -4
View File
@@ -313,10 +313,7 @@ async def refund_wallet_endpoint(
) )
if key.reserved_balance > 0: if key.reserved_balance > 0:
# Self-heal: reservations leaked by client disconnects or crashes would # Release the reservation if it is stale
# otherwise lock the user out of refunding forever. Release the
# reservation if it is stale (or predates reserved_at tracking) and
# proceed; only reject when a reservation is genuinely recent.
cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds
stale_release_stmt = ( stale_release_stmt = (
update(ApiKey) update(ApiKey)
-5
View File
@@ -105,11 +105,6 @@ async def release_stale_reservations(
session: AsyncSession, max_age_seconds: int session: AsyncSession, max_age_seconds: int
) -> int: ) -> int:
"""Release reservations whose last reserve is older than max_age_seconds. """Release reservations whose last reserve is older than max_age_seconds.
Only rows with a known reservation time are touched — NULL `reserved_at`
rows are left alone here so reservations made by instances running older
code (rolling deploys) are never killed mid-flight. Those rows are healed
on demand by the refund endpoint instead.
""" """
cutoff = int(time.time()) - max_age_seconds cutoff = int(time.time()) - max_age_seconds
stmt = ( stmt = (
+1 -2
View File
@@ -173,8 +173,7 @@ def resolve_bootstrap() -> Settings:
pass pass
if not base.onion_url: if not base.onion_url:
try: try:
from ..nostr.listing import \ from ..nostr.listing import discover_onion_url_from_tor # type: ignore
discover_onion_url_from_tor # type: ignore
discovered = discover_onion_url_from_tor() discovered = discover_onion_url_from_tor()
if discovered: if discovered:
+14 -5
View File
@@ -9,14 +9,23 @@ from sqlmodel import select
from .algorithm import create_model_mappings from .algorithm import create_model_mappings
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger from .core import get_logger
from .core.db import (ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, from .core.db import (
create_session, get_session) ApiKey,
AsyncSession,
ModelRow,
UpstreamProviderRow,
create_session,
get_session,
)
from .core.exceptions import UpstreamError from .core.exceptions import UpstreamError
from .core.not_found import build_not_found_response from .core.not_found import build_not_found_response
from .core.settings import settings from .core.settings import settings
from .payment.helpers import (calculate_discounted_max_cost, from .payment.helpers import (
check_token_balance, create_error_response, calculate_discounted_max_cost,
get_max_cost_for_model) check_token_balance,
create_error_response,
get_max_cost_for_model,
)
from .payment.models import Model from .payment.models import Model
from .upstream import BaseUpstreamProvider from .upstream import BaseUpstreamProvider
from .upstream.helpers import init_upstreams from .upstream.helpers import init_upstreams
+267 -1
View File
@@ -1,12 +1,14 @@
import asyncio
import time import time
from datetime import datetime, timedelta from datetime import datetime, timedelta
import pytest import pytest
from fastapi import HTTPException
from sqlmodel import select from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import pay_for_request from routstr.auth import pay_for_request
from routstr.core.db import ApiKey from routstr.core.db import ApiKey, create_session
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -120,6 +122,270 @@ async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None
assert key2.total_spent == 0 assert key2.total_spent == 0
@pytest.mark.asyncio
async def test_balance_limit_enforced_atomically_under_concurrency(
patched_db_engine: None,
) -> None:
parent_hash = "parent_limit_atomic"
child_hash = "child_limit_atomic"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=cost,
total_spent=0,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_child = await session.get(ApiKey, child_hash)
assert final_child is not None
assert final_child.reserved_balance == cost, (
f"Child reserved_balance should equal one reservation, "
f"got {final_child.reserved_balance}"
)
@pytest.mark.asyncio
async def test_parallel_payments_with_parent_and_child_key(
patched_db_engine: None,
) -> None:
parent_hash = "parent_parallel_mixed"
child_hash = "child_parallel_mixed"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=2 * cost,
)
session.add(parent)
session.add(child)
await session.commit()
async def attempt(key_hash: str) -> str:
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
try:
await pay_for_request(fresh_key, cost, session)
return "success"
except HTTPException as exc:
assert exc.status_code == 402
return "blocked"
results = await asyncio.gather(attempt(parent_hash), attempt(child_hash))
assert results == ["success", "success"], (
f"Both parent and child payments should succeed, got: {results}"
)
async with create_session() as session:
final_parent = await session.get(ApiKey, parent_hash)
final_child = await session.get(ApiKey, child_hash)
assert final_parent is not None
assert final_child is not None
# Both requests bill the parent; only the child request reserves on the child.
assert final_parent.reserved_balance == 2 * cost
assert final_parent.total_requests == 2
assert final_child.reserved_balance == cost
assert final_child.total_requests == 1
@pytest.mark.asyncio
async def test_balance_limit_with_existing_total_spent_under_concurrency(
patched_db_engine: None,
) -> None:
parent_hash = "parent_total_spent"
child_hash = "child_total_spent"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
# 700 already spent against a 1000 limit: only one more 300 request fits.
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
total_spent=700,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_child = await session.get(ApiKey, child_hash)
assert final_child is not None
assert final_child.reserved_balance == cost
assert final_child.total_spent == 700
@pytest.mark.asyncio
async def test_balance_limit_with_existing_reserved_balance(
patched_db_engine: None,
) -> None:
parent_hash = "parent_reserved_set"
blocked_hash = "child_reserved_blocked"
allowed_hash = "child_reserved_allowed"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
# 800 already reserved against a 1000 limit: another 300 must be rejected.
blocked_child = ApiKey(
hashed_key=blocked_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
reserved_balance=800,
)
# 500 reserved against a 1000 limit: another 300 still fits.
allowed_child = ApiKey(
hashed_key=allowed_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
reserved_balance=500,
)
session.add(parent)
session.add(blocked_child)
session.add(allowed_child)
await session.commit()
async with create_session() as session:
fresh_blocked = await session.get(ApiKey, blocked_hash)
assert fresh_blocked is not None
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(fresh_blocked, cost, session)
assert exc_info.value.status_code == 402
async with create_session() as session:
fresh_allowed = await session.get(ApiKey, allowed_hash)
assert fresh_allowed is not None
await pay_for_request(fresh_allowed, cost, session)
async with create_session() as session:
final_blocked = await session.get(ApiKey, blocked_hash)
final_allowed = await session.get(ApiKey, allowed_hash)
final_parent = await session.get(ApiKey, parent_hash)
assert final_blocked is not None
assert final_allowed is not None
assert final_parent is not None
assert final_blocked.reserved_balance == 800, "Rejected request must not reserve"
assert final_blocked.total_requests == 0
assert final_allowed.reserved_balance == 500 + cost
assert final_allowed.total_requests == 1
# Only the allowed request should have billed the parent.
assert final_parent.reserved_balance == cost
assert final_parent.total_requests == 1
@pytest.mark.asyncio
async def test_child_reservation_discarded_when_parent_balance_depleted(
patched_db_engine: None,
) -> None:
parent_hash = "parent_depleted"
child_hash = "child_depleted"
cost = 300
async with create_session() as session:
# Parent can only afford one request; child has no balance_limit.
parent = ApiKey(hashed_key=parent_hash, balance=cost)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_parent = await session.get(ApiKey, parent_hash)
final_child = await session.get(ApiKey, child_hash)
assert final_parent is not None
assert final_child is not None
assert final_parent.reserved_balance == cost
# The failed request must not leave a committed reservation on the child.
assert final_child.reserved_balance == cost, (
f"Child reserved_balance should reflect only the successful request, "
f"got {final_child.reserved_balance}"
)
assert final_child.total_requests == 1
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None: async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None:
# This requires mocking the router call or testing the logic in balance.py # This requires mocking the router call or testing the logic in balance.py