From 9fb6f54d12a0bdfd2bc30374a75905dc52e76361 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 3 Mar 2026 15:29:07 +0100 Subject: [PATCH] fix negative reserve balance --- routstr/auth.py | 52 +++-- routstr/proxy.py | 6 +- routstr/upstream/base.py | 28 +-- .../test_reserved_balance_negative.py | 221 +++++++++++++++++- 4 files changed, 253 insertions(+), 54 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 6e03d968..fa880f92 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -588,12 +588,15 @@ async def pay_for_request( async def revert_pay_for_request( key: ApiKey, session: AsyncSession, cost_per_request: int -) -> None: +) -> bool: + """Revert a previously reserved payment. Returns True if revert succeeded, + False if the reservation was already released (prevents negative reserved_balance).""" billing_key = await get_billing_key(key, session) stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.reserved_balance) >= cost_per_request) .values( reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, @@ -607,6 +610,7 @@ async def revert_pay_for_request( child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.reserved_balance) >= cost_per_request) .values( total_requests=col(ApiKey.total_requests) - 1, reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, @@ -616,8 +620,8 @@ async def revert_pay_for_request( await session.commit() if result.rowcount == 0: - logger.error( - "Failed to revert payment - insufficient reserved balance", + logger.warning( + "Revert skipped - reservation already released (no-op to prevent negative reserved_balance)", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", @@ -625,19 +629,11 @@ async def revert_pay_for_request( "current_reserved_balance": billing_key.reserved_balance, }, ) - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", - "type": "payment_error", - "code": "payment_error", - } - }, - ) + return False await session.refresh(billing_key) if billing_key.hashed_key != key.hashed_key: await session.refresh(key) + return True async def adjust_payment_for_tokens( @@ -669,17 +665,19 @@ async def adjust_payment_for_tokens( release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost ) ) - await session.exec(release_stmt) # type: ignore[call-overload] + result = await session.exec(release_stmt) # type: ignore[call-overload] # Also release on child key if it's different if billing_key.hashed_key != key.hashed_key: child_release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost @@ -688,14 +686,24 @@ async def adjust_payment_for_tokens( await session.exec(child_release_stmt) # type: ignore[call-overload] await session.commit() - logger.warning( - "Released reservation without charging (fallback)", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "deducted_max_cost": deducted_max_cost, - }, - ) + if result.rowcount == 0: # type: ignore[union-attr] + logger.warning( + "Release reservation skipped - already released (no-op to prevent negative reserved_balance)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, + ) + else: + logger.warning( + "Released reservation without charging (fallback)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, + ) except Exception as e: logger.error( "Failed to release reservation in fallback", diff --git a/routstr/proxy.py b/routstr/proxy.py index 4d8f666d..09b34549 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -300,9 +300,13 @@ async def proxy( session, model_obj, ) + except UpstreamError: + # Let the outer UpstreamError handler manage retry/revert + raise except Exception as e: + # Unexpected error (not an upstream failure) — revert and propagate logger.error( - "Upstream request failed, ensuring payment is reverted", + "Unexpected error in upstream request, reverting payment", extra={ "error": str(e), "error_type": type(e).__name__, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5efe5b0c..1be99b15 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -13,7 +13,7 @@ from fastapi.responses import Response, StreamingResponse from pydantic import BaseModel from sqlmodel import select -from ..auth import adjust_payment_for_tokens, revert_pay_for_request +from ..auth import adjust_payment_for_tokens from ..core import get_logger from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow, create_session from ..core.exceptions import UpstreamError @@ -1279,8 +1279,7 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - + # Don't revert here — proxy.py owns payment revert to avoid double-revert if isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" elif isinstance(exc, httpx.TimeoutException): @@ -1310,13 +1309,9 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + # Don't revert here — proxy.py owns payment revert to avoid double-revert + raise UpstreamError( + "An unexpected server error occurred", status_code=500 ) async def forward_responses_request( @@ -1501,8 +1496,7 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - + # Don't revert here — proxy.py owns payment revert to avoid double-revert if isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" elif isinstance(exc, httpx.TimeoutException): @@ -1532,13 +1526,9 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + # Don't revert here — proxy.py owns payment revert to avoid double-revert + raise UpstreamError( + "An unexpected server error occurred", status_code=500 ) async def forward_get_request( diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py index 21f9b6d0..b283d55b 100644 --- a/tests/integration/test_reserved_balance_negative.py +++ b/tests/integration/test_reserved_balance_negative.py @@ -133,13 +133,16 @@ async def test_reserved_balance_with_successful_requests( @pytest.mark.asyncio -async def test_insufficient_reserved_balance_for_revert( +async def test_revert_with_zero_reserved_balance_is_noop( integration_session: AsyncSession, ) -> None: - """Test revert_pay_for_request behavior with insufficient reserved balance.""" + """Test that revert_pay_for_request is a no-op when reserved_balance is 0. + + Previously this would drive reserved_balance negative. With the floor guard, + it should return False and leave reserved_balance at 0. + """ from routstr.auth import revert_pay_for_request - # Create key with zero reserved balance unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}" test_key = ApiKey( hashed_key=unique_key, @@ -149,17 +152,211 @@ async def test_insufficient_reserved_balance_for_revert( integration_session.add(test_key) await integration_session.commit() - # Try to revert more than available - # Note: Current implementation allows reserved_balance to go negative - await revert_pay_for_request(test_key, integration_session, 100) + # Try to revert more than available — should be a no-op + result = await revert_pay_for_request(test_key, integration_session, 100) - # Refresh to get updated values await integration_session.refresh(test_key) - # Current implementation allows negative reserved balance - assert test_key.reserved_balance == -100, ( - f"Expected reserved_balance to be -100, got: {test_key.reserved_balance}" + assert result is False, "Revert should return False when reservation already released" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should remain 0, got: {test_key.reserved_balance}" ) - assert test_key.total_requests == -1, ( - f"Expected total_requests to be -1, got: {test_key.total_requests}" + assert test_key.total_requests == 0, ( + f"Total requests should remain 0, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_revert_with_sufficient_reserved_balance_succeeds( + integration_session: AsyncSession, +) -> None: + """Test that revert_pay_for_request works correctly when there is enough reserved balance.""" + from routstr.auth import revert_pay_for_request + + unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=5000, + reserved_balance=500, + total_requests=3, + ) + integration_session.add(test_key) + await integration_session.commit() + + result = await revert_pay_for_request(test_key, integration_session, 500) + + await integration_session.refresh(test_key) + + assert result is True, "Revert should return True on success" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should be 0, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 2, ( + f"Total requests should be 2, got: {test_key.total_requests}" + ) + assert test_key.balance == 5000, "Balance should not change on revert" + + +@pytest.mark.asyncio +async def test_revert_partial_reserved_balance_is_noop( + integration_session: AsyncSession, +) -> None: + """Test that reverting more than the current reserved_balance is a no-op.""" + from routstr.auth import revert_pay_for_request + + unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=5000, + reserved_balance=50, + total_requests=1, + ) + integration_session.add(test_key) + await integration_session.commit() + + # Try to revert 500 when only 50 is reserved — should be no-op + result = await revert_pay_for_request(test_key, integration_session, 500) + + await integration_session.refresh(test_key) + + assert result is False, "Revert should fail when cost > reserved_balance" + assert test_key.reserved_balance == 50, ( + f"Reserved balance should stay at 50, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 1, ( + f"Total requests should stay at 1, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_double_revert_prevented( + integration_session: AsyncSession, +) -> None: + """Test that calling revert twice doesn't drive reserved_balance negative. + + This simulates the double-revert scenario where both upstream/base.py + and proxy.py attempt to revert the same reservation. + """ + from routstr.auth import revert_pay_for_request + + unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=10000, + reserved_balance=500, + total_requests=5, + ) + integration_session.add(test_key) + await integration_session.commit() + + # First revert — should succeed + result1 = await revert_pay_for_request(test_key, integration_session, 500) + await integration_session.refresh(test_key) + + assert result1 is True + assert test_key.reserved_balance == 0 + assert test_key.total_requests == 4 + + # Second revert of the same amount — should be no-op + result2 = await revert_pay_for_request(test_key, integration_session, 500) + await integration_session.refresh(test_key) + + assert result2 is False, "Second revert should be a no-op" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should stay 0, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 4, ( + f"Total requests should stay 4, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_sequential_reverts_never_go_negative( + integration_session: AsyncSession, +) -> None: + """Test that multiple reverts don't cause negative reserved_balance. + + Simulates the double-revert scenario where multiple code paths + attempt to revert the same reservation. + """ + from routstr.auth import revert_pay_for_request + + unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=10000, + reserved_balance=500, + total_requests=5, + ) + integration_session.add(test_key) + await integration_session.commit() + + # Run 5 sequential reverts for the same 500 reservation + results = [] + for _ in range(5): + r = await revert_pay_for_request(test_key, integration_session, 500) + results.append(r) + + await integration_session.refresh(test_key) + + # Exactly one should succeed, rest should be no-ops + success_count = sum(1 for r in results if r is True) + assert success_count == 1, ( + f"Exactly one revert should succeed, got {success_count} successes" + ) + assert test_key.reserved_balance == 0, ( + f"Reserved balance should be 0, got: {test_key.reserved_balance}" + ) + assert test_key.reserved_balance >= 0, ( + f"Reserved balance went negative: {test_key.reserved_balance}" + ) + + +@pytest.mark.asyncio +async def test_child_key_revert_floor_guard( + integration_session: AsyncSession, +) -> None: + """Test that child key reserved_balance also has floor guard on revert.""" + from routstr.auth import revert_pay_for_request + + parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}" + child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}" + + parent_key = ApiKey( + hashed_key=parent_key_hash, + balance=10000, + reserved_balance=500, + total_requests=3, + ) + child_key = ApiKey( + hashed_key=child_key_hash, + balance=0, + reserved_balance=500, + total_requests=3, + parent_key_hash=parent_key_hash, + ) + integration_session.add(parent_key) + integration_session.add(child_key) + await integration_session.commit() + + # First revert succeeds + result1 = await revert_pay_for_request(child_key, integration_session, 500) + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key) + + assert result1 is True + assert parent_key.reserved_balance == 0 + assert child_key.reserved_balance == 0 + + # Second revert is a no-op for both parent and child + result2 = await revert_pay_for_request(child_key, integration_session, 500) + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key) + + assert result2 is False + assert parent_key.reserved_balance == 0, ( + f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}" + ) + assert child_key.reserved_balance == 0, ( + f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}" )