diff --git a/migrations/versions/7f2843d3f4e4_add_reservation_release_idempotency_.py b/migrations/versions/7f2843d3f4e4_add_reservation_release_idempotency_.py new file mode 100644 index 00000000..309e35b7 --- /dev/null +++ b/migrations/versions/7f2843d3f4e4_add_reservation_release_idempotency_.py @@ -0,0 +1,62 @@ +"""add reservation release idempotency records + +Revision ID: 7f2843d3f4e4 +Revises: fc4fa29630d2 +Create Date: 2026-07-24 02:06:06.066726 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "7f2843d3f4e4" +down_revision = "fc4fa29630d2" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "reservation_releases", + sa.Column("id", sa.String(), nullable=False), + sa.Column("key_hash", sa.String(), nullable=False), + sa.Column("billing_key_hash", sa.String(), nullable=False), + sa.Column("reserved_msats", sa.Integer(), nullable=False), + sa.Column( + "status", sa.String(), nullable=False, server_default="active" + ), + sa.Column("created_at", sa.Integer(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index( + "ix_reservation_releases_key_hash", + "reservation_releases", + ["key_hash"], + ) + op.create_index( + "ix_reservation_releases_billing_key_hash", + "reservation_releases", + ["billing_key_hash"], + ) + op.create_index( + "ix_reservation_releases_status_created_at", + "reservation_releases", + ["status", "created_at"], + ) + + +def downgrade() -> None: + op.drop_index( + "ix_reservation_releases_status_created_at", + table_name="reservation_releases", + ) + op.drop_index( + "ix_reservation_releases_billing_key_hash", + table_name="reservation_releases", + ) + op.drop_index( + "ix_reservation_releases_key_hash", + table_name="reservation_releases", + ) + op.drop_table("reservation_releases") diff --git a/routstr/auth.py b/routstr/auth.py index 84a57b07..42c7eed2 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -3,16 +3,25 @@ import hashlib import math import random import time +import uuid +from contextvars import ContextVar +from dataclasses import dataclass from datetime import datetime from typing import TYPE_CHECKING, Optional from fastapi import HTTPException -from sqlalchemy import case +from sqlalchemy import case, inspect from sqlalchemy.exc import IntegrityError from sqlmodel import col, select, update from .core import get_logger -from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee +from .core.db import ( + ApiKey, + AsyncSession, + ReservationRelease, + accumulate_routstr_fee, + create_session, +) from .core.settings import settings from .payment.cost_calculation import ( CostData, @@ -34,10 +43,32 @@ payments_logger = get_logger("routstr.payments") # Routstr platform fee constants ROUTSTR_FEE_PERCENT: float = 2.1 -ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash" +ROUTSTR_LN_ADDRESS: str = ( + "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash" +) ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900 ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200 + +@dataclass(frozen=True) +class ReservationSnapshot: + release_id: str + key_hash: str + billing_key_hash: str + reserved_msats: int + + +_current_reservation: ContextVar[ReservationSnapshot | None] = ContextVar( + "current_billing_reservation", default=None +) + + +def _clear_current_reservation(snapshot: ReservationSnapshot) -> None: + current = _current_reservation.get() + if current is not None and current.release_id == snapshot.release_id: + _current_reservation.set(None) + + # TODO: implement prepaid api key (not like it was before) # PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None) # PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats @@ -595,6 +626,16 @@ async def pay_for_request( }, ) + # Create the durable reservation identity before changing aggregate balances. + # The row and balance updates commit together, so every reserved amount has one + # owner that can reach exactly one terminal state. + reservation = ReservationSnapshot( + release_id=uuid.uuid4().hex, + key_hash=key.hashed_key, + billing_key_hash=billing_key.hashed_key, + reserved_msats=cost_per_request, + ) + # Charge the base cost for the request atomically to avoid race conditions reserved_at_now = int(time.time()) stmt = ( @@ -656,22 +697,83 @@ async def pay_for_request( child_result = await session.exec(child_stmt) # type: ignore[call-overload] if child_result.rowcount == 0: + # Build the error before rollback expires ORM attributes. + limit_message = ( + f"Balance limit exceeded: {key.balance_limit} mSats limit. " + f"{key.total_spent} already spent ({key.reserved_balance} reserved), " + f"{cost_per_request} required for this request." + ) + # The parent reservation update already ran in this transaction. + # Roll it back before failover code attempts to restore the previous + # reservation; otherwise that later commit can persist both updates. + await session.rollback() 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.", + "message": limit_message, "type": "insufficient_quota", "code": "balance_limit_exceeded", } }, ) - await session.commit() + session.add( + ReservationRelease( + id=reservation.release_id, + key_hash=reservation.key_hash, + billing_key_hash=reservation.billing_key_hash, + reserved_msats=reservation.reserved_msats, + status="active", + ) + ) + # Publish the identity before commit. If the commit succeeds but its + # acknowledgement is interrupted, exact cleanup can still recover the + # durable row. A definitely failed commit is harmless because every + # terminal transition validates that row before touching balances. + _current_reservation.set(reservation) + try: + await session.commit() + except BaseException: + # The database may have committed even if acknowledgement was cancelled + # or the connection failed. Reconcile using a fresh transaction and the + # exact durable identity; no upstream request has started yet. + try: + await session.rollback() + except Exception: + pass + try: + async with create_session() as cleanup_session: + record = await cleanup_session.get( + ReservationRelease, reservation.release_id + ) + if record is not None and record.status == "active": + await _transition_reservation_to_released( + reservation, + cleanup_session, + decrement_requests=True, + idempotent_success=True, + ) + except Exception: + logger.exception( + "Failed to reconcile ambiguous reservation commit", + extra={"reservation_id": reservation.release_id}, + ) + finally: + _clear_current_reservation(reservation) + raise - await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) + try: + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) + except Exception: + # The reservation transaction is already committed and durable. Logging + # refresh failures must not make the caller treat it as unreserved. + logger.exception( + "Reservation committed but post-commit refresh failed", + extra={"reservation_id": reservation.release_id}, + ) logger.info( "Payment processed successfully", @@ -701,81 +803,185 @@ async def pay_for_request( async def revert_pay_for_request( - key: ApiKey, session: AsyncSession, cost_per_request: int + key: ApiKey, + session: AsyncSession, + cost_per_request: int, + reservation_snapshot: ReservationSnapshot | None = 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) - - # Keep reserved_at while other reservations remain - cleared_reserved_at = case( - (col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)), - else_=None, + """Revert the current request's durable reservation exactly once.""" + snapshot = reservation_snapshot or await get_reservation_snapshot(key, session) + await _validate_reservation_snapshot(key, snapshot, session, require_active=False) + if cost_per_request != snapshot.reserved_msats: + return False + return await _transition_reservation_to_released( + snapshot, + session, + decrement_requests=True, + idempotent_success=False, ) - stmt = ( + +async def _validate_reservation_snapshot( + key: ApiKey, + snapshot: ReservationSnapshot, + session: AsyncSession, + *, + require_active: bool = True, +) -> None: + """Reject cross-request or forged reservation handles before any mutation.""" + state = inspect(key) + identity = state.identity if state is not None else None + key_hash = str(identity[0]) if identity else key.__dict__.get("hashed_key") + if snapshot.key_hash != key_hash: + raise RuntimeError("Billing reservation does not belong to this key") + + persisted_key = await session.get(ApiKey, snapshot.key_hash) + if persisted_key is None: + raise RuntimeError("Billing reservation key no longer exists") + expected_billing_hash = persisted_key.parent_key_hash or persisted_key.hashed_key + if snapshot.billing_key_hash != expected_billing_hash: + raise RuntimeError("Billing reservation does not belong to this billing key") + + record = await session.get(ReservationRelease, snapshot.release_id) + if ( + record is None + or (require_active and record.status != "active") + or record.key_hash != snapshot.key_hash + or record.billing_key_hash != snapshot.billing_key_hash + or record.reserved_msats != snapshot.reserved_msats + ): + raise RuntimeError("Billing reservation record does not match the request") + + +async def get_reservation_snapshot( + key: ApiKey, session: AsyncSession +) -> ReservationSnapshot: + """Return the durable reservation created for the current request.""" + snapshot = _current_reservation.get() + if snapshot is None: + raise RuntimeError("No billing reservation is associated with this request") + await _validate_reservation_snapshot(key, snapshot, session) + return snapshot + + +async def _transition_reservation_to_released( + snapshot: ReservationSnapshot, + session: AsyncSession, + *, + decrement_requests: bool, + idempotent_success: bool, +) -> bool: + transition = ( + update(ReservationRelease) + .where(col(ReservationRelease.id) == snapshot.release_id) + .where(col(ReservationRelease.status) == "active") + .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) + .values(status="released") + ) + transition_result = await session.exec(transition) # type: ignore[call-overload] + if transition_result.rowcount != 1: + await session.rollback() + existing = await session.get(ReservationRelease, snapshot.release_id) + return bool( + idempotent_success + and existing is not None + and existing.status == "released" + and existing.key_hash == snapshot.key_hash + and existing.billing_key_hash == snapshot.billing_key_hash + and existing.reserved_msats == snapshot.reserved_msats + ) + + values: dict[str, object] = { + "reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats, + "reserved_at": case( + ( + col(ApiKey.reserved_balance) - snapshot.reserved_msats > 0, + col(ApiKey.reserved_at), + ), + else_=None, + ), + } + if decrement_requests: + values["total_requests"] = col(ApiKey.total_requests) - 1 + + release_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, - reserved_at=cleared_reserved_at, - total_requests=col(ApiKey.total_requests) - 1, - ) + .where(col(ApiKey.hashed_key) == snapshot.billing_key_hash) + .where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats) + .values(**values) ) + result = await session.exec(release_stmt) # type: ignore[call-overload] + if result.rowcount != 1: + await session.rollback() + return False - result = await session.exec(stmt) # type: ignore[call-overload] - - # Also decrement total_requests and reserved_balance on the child key if it's different - if billing_key.hashed_key != key.hashed_key: - child_stmt = ( + if snapshot.billing_key_hash != snapshot.key_hash: + child_release_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, - reserved_at=cleared_reserved_at, - ) + .where(col(ApiKey.hashed_key) == snapshot.key_hash) + .where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats) + .values(**values) ) - await session.exec(child_stmt) # type: ignore[call-overload] + child_result = await session.exec( # type: ignore[call-overload] + child_release_stmt + ) + if child_result.rowcount != 1: + await session.rollback() + return False await session.commit() - if result.rowcount == 0: - 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] + "...", - "cost_to_revert": cost_per_request, - "current_reserved_balance": billing_key.reserved_balance, - }, - ) - return False - await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) - payments_logger.info( - "REVERT", - extra={ - "event": "revert", - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "cost_reverted": cost_per_request, - "balance": billing_key.balance, - "reserved_balance": billing_key.reserved_balance, - }, - ) + _clear_current_reservation(snapshot) return True +async def release_reservation( + snapshot: ReservationSnapshot, + session: AsyncSession, + reserved_msats: int, +) -> bool: + """Release one durable reservation exactly once without charging.""" + if reserved_msats <= 0 or reserved_msats != snapshot.reserved_msats: + return False + return await _transition_reservation_to_released( + snapshot, + session, + decrement_requests=False, + idempotent_success=True, + ) + + +async def _claim_reservation_for_charge( + snapshot: ReservationSnapshot, session: AsyncSession +) -> bool: + """Claim an active reservation in the caller's charge transaction.""" + statement = ( + update(ReservationRelease) + .where(col(ReservationRelease.id) == snapshot.release_id) + .where(col(ReservationRelease.status) == "active") + .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) + .values(status="charged") + ) + result = await session.exec(statement) # type: ignore[call-overload] + if result.rowcount == 1: + _clear_current_reservation(snapshot) + return True + + await session.rollback() + return False + + async def adjust_payment_for_tokens( key: ApiKey, response_data: dict, session: AsyncSession, deducted_max_cost: int, - model_obj: "Model | None", - provider_fee: float | None, + model_obj: "Model | None" = None, + provider_fee: float | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> dict: """ Adjusts the payment based on token usage in the response. @@ -790,6 +996,13 @@ async def adjust_payment_for_tokens( ``calculate_cost``. """ billing_key = await get_billing_key(key, session) + reservation = reservation_snapshot or await get_reservation_snapshot(key, session) + await _validate_reservation_snapshot( + key, reservation, session, require_active=False + ) + # The persisted amount is authoritative if request-level minimum pricing + # changed the caller's original estimate. + deducted_max_cost = reservation.reserved_msats model = response_data.get("model", "unknown") logger.debug( @@ -805,50 +1018,21 @@ async def adjust_payment_for_tokens( ) async def release_reservation_only() -> None: - """Fallback to release reservation without charging when main update fails.""" + """Fallback to release this request's reservation without charging.""" try: - 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 - ) + released = await release_reservation( + reservation, session, reservation.reserved_msats + ) + logger.warning( + "Released reservation without charging (fallback)" + if released + else "Reservation was already finalized; fallback skipped", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, ) - 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 - ) - ) - await session.exec(child_release_stmt) # type: ignore[call-overload] - - await session.commit() - 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", @@ -870,9 +1054,17 @@ async def adjust_payment_for_tokens( extra={"error": str(e), "fee_msats": fee_msats}, ) - match await calculate_cost( + calculated_cost = await calculate_cost( response_data, deducted_max_cost, model_obj, provider_fee - ): + ) + if not isinstance(calculated_cost, CostDataError): + if not await _claim_reservation_for_charge(reservation, session): + # A prior charge or release already owns this reservation. Returning + # the calculated metadata is safe; the aggregate balances must not + # be modified a second time. + return calculated_cost.dict() + + match calculated_cost: case MaxCostData() as cost: logger.debug( "Using max cost data (no token adjustment)", @@ -900,8 +1092,10 @@ async def adjust_payment_for_tokens( ) safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) @@ -919,8 +1113,10 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: child_safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) child_stmt = ( @@ -1030,8 +1226,10 @@ async def adjust_payment_for_tokens( ) exact_safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) @@ -1049,8 +1247,10 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: child_exact_safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) child_stmt = ( @@ -1089,31 +1289,45 @@ async def adjust_payment_for_tokens( # actual cost exceeded discounted reservation (due to tolerance_percentage) if cost_difference > 0: - # Always release the reservation and charge min(actual_cost, balance). - # CASE expressions keep this atomic and safe even when the - # stale-reservation sweeper has already released the reservation. - chargeable = case( - (col(ApiKey.balance) >= total_cost_msats, total_cost_msats), - else_=col(ApiKey.balance), - ) - overrun_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=overrun_safe_reserved, - balance=col(ApiKey.balance) - chargeable, - total_spent=col(ApiKey.total_spent) + chargeable, + # Lock the billing row so the parent and child record the same + # database-determined charge under concurrent finalizations. + actual_charge_msats = 0 + for attempt in range(5): + locked_billing_key = ( + await session.exec( + select(ApiKey) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .with_for_update() + .execution_options(populate_existing=True) + ) + ).one() + observed_balance = locked_billing_key.balance + actual_charge_msats = min(observed_balance, total_cost_msats) + overrun_safe_reserved = case( + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), + else_=0, ) - ) - await session.exec(finalize_stmt) # type: ignore[call-overload] + finalize_result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.balance) == observed_balance) + .values( + reserved_balance=overrun_safe_reserved, + balance=col(ApiKey.balance) - actual_charge_msats, + total_spent=col(ApiKey.total_spent) + actual_charge_msats, + ) + ) + if finalize_result.rowcount == 1: + break + await session.rollback() + if not await _claim_reservation_for_charge(reservation, session): + return cost.dict() + else: + await session.rollback() + raise RuntimeError("Could not atomically finalize cost overrun") if billing_key.hashed_key != key.hashed_key: child_stmt = ( @@ -1121,7 +1335,7 @@ async def adjust_payment_for_tokens( .where(col(ApiKey.hashed_key) == key.hashed_key) .values( reserved_balance=overrun_safe_reserved, - total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats), + total_spent=col(ApiKey.total_spent) + actual_charge_msats, ) ) await session.exec(child_stmt) # type: ignore[call-overload] @@ -1131,18 +1345,18 @@ async def adjust_payment_for_tokens( await session.refresh(billing_key) if billing_key.hashed_key != key.hashed_key: await session.refresh(key) - cost.total_msats = total_cost_msats + cost.total_msats = actual_charge_msats logger.info( "Finalized payment with additional charge", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", - "charged_amount": total_cost_msats, + "charged_amount": actual_charge_msats, "new_balance": billing_key.balance, "model": model, }, ) - await _accumulate_fee(total_cost_msats) + await _accumulate_fee(actual_charge_msats) payments_logger.info( "FINALIZE", extra={ @@ -1151,7 +1365,7 @@ async def adjust_payment_for_tokens( "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "cost_reserved": deducted_max_cost, - "cost_charged": total_cost_msats, + "cost_charged": actual_charge_msats, "input_tokens": cost.input_tokens, "output_tokens": cost.output_tokens, "balance": billing_key.balance, @@ -1191,8 +1405,10 @@ async def adjust_payment_for_tokens( ) refund_safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) @@ -1210,8 +1426,10 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: child_refund_safe_reserved = case( - (col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost), + ( + col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost, + ), else_=0, ) child_stmt = ( @@ -1386,9 +1604,7 @@ async def periodic_dead_key_prune() -> None: try: async with create_session() as session: - await prune_dead_api_keys( - session, settings.dead_key_min_age_seconds - ) + await prune_dead_api_keys(session, settings.dead_key_min_age_seconds) except asyncio.CancelledError: break except Exception as e: diff --git a/routstr/balance.py b/routstr/balance.py index 4630c224..03ddf36c 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -7,7 +7,7 @@ from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException from fastapi.responses import JSONResponse from pydantic import BaseModel -from sqlmodel import col, or_, select, update +from sqlmodel import col, select, update from .auth import get_billing_key, validate_bearer_key from .core.db import ( @@ -15,6 +15,7 @@ from .core.db import ( AsyncSession, CashuTransaction, get_session, + release_stale_reservations, ) from .core.db import ( store_cashu_transaction_with_retry as store_cashu_transaction, @@ -323,30 +324,19 @@ async def refund_wallet_endpoint( ) if key.reserved_balance > 0: - # Release the reservation if it is stale - cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds - stale_release_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.reserved_balance) > 0) - .where( - or_( - col(ApiKey.reserved_at).is_(None), - col(ApiKey.reserved_at) < cutoff, - ) - ) - .values(reserved_balance=0, reserved_at=None) + # Release only durable reservations old enough to be stale. A newer + # request on the same aggregate balance must remain reserved. + await release_stale_reservations( + session, + settings.stale_reservation_timeout_seconds, + key_hash=key.hashed_key, ) - stale_result = await session.exec(stale_release_stmt) # type: ignore[call-overload] - await session.commit() - - if stale_result.rowcount == 0: + await session.refresh(key) + if key.reserved_balance > 0: raise HTTPException( status_code=400, detail="Cannot refund key. There are ongoing requests for this api key.", ) - - await session.refresh(key) logger.warning( "refund_wallet_endpoint: released stale reservation before refund", extra={ diff --git a/routstr/core/db.py b/routstr/core/db.py index b81cb56d..16cfedfd 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -12,7 +12,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from alembic.util.exc import CommandError -from sqlalchemy import UniqueConstraint, delete +from sqlalchemy import Index, UniqueConstraint, case, delete, or_ from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlalchemy.orm import aliased @@ -99,32 +99,130 @@ class ApiKey(SQLModel, table=True): # type: ignore async def reset_all_reserved_balances(session: AsyncSession) -> None: - stmt = update(ApiKey).values(reserved_balance=0, reserved_at=None) - await session.exec(stmt) # type: ignore[call-overload] + """Release every active durable reservation during explicit startup reset.""" + await session.exec( # type: ignore[call-overload] + update(ReservationRelease) + .where(col(ReservationRelease.status) == "active") + .values(status="released") + ) + await session.exec( # type: ignore[call-overload] + update(ApiKey).values(reserved_balance=0, reserved_at=None) + ) await session.commit() logger.info("Reset reserved balances on startup") async def release_stale_reservations( - session: AsyncSession, max_age_seconds: int + session: AsyncSession, + max_age_seconds: int, + *, + key_hash: str | None = None, ) -> int: - """Release reservations whose last reserve is older than max_age_seconds. - """ + """Release stale durable reservations without touching newer reservations.""" cutoff = int(time.time()) - max_age_seconds - stmt = ( - update(ApiKey) - .where(col(ApiKey.reserved_balance) > 0) - .where(col(ApiKey.reserved_at).is_not(None)) - .where(col(ApiKey.reserved_at) < cutoff) - .values(reserved_balance=0, reserved_at=None) + query = ( + select(ReservationRelease) + .where(col(ReservationRelease.status) == "active") + .where(col(ReservationRelease.created_at) < cutoff) ) - result = await session.exec(stmt) # type: ignore[call-overload] + if key_hash is not None: + query = query.where( + or_( + col(ReservationRelease.key_hash) == key_hash, + col(ReservationRelease.billing_key_hash) == key_hash, + ) + ) + reservations = (await session.exec(query)).all() + released = 0 + + for reservation in reservations: + transition = await session.exec( # type: ignore[call-overload] + update(ReservationRelease) + .where(col(ReservationRelease.id) == reservation.id) + .where(col(ReservationRelease.status) == "active") + .values(status="released") + ) + if transition.rowcount != 1: + continue + + values = { + "reserved_balance": col(ApiKey.reserved_balance) + - reservation.reserved_msats, + "reserved_at": case( + ( + col(ApiKey.reserved_balance) - reservation.reserved_msats > 0, + col(ApiKey.reserved_at), + ), + else_=None, + ), + } + parent_result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == reservation.billing_key_hash) + .where(col(ApiKey.reserved_balance) >= reservation.reserved_msats) + .values(**values) + ) + if parent_result.rowcount != 1: + await session.rollback() + return 0 + + if reservation.billing_key_hash != reservation.key_hash: + child_result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == reservation.key_hash) + .where(col(ApiKey.reserved_balance) >= reservation.reserved_msats) + .values(**values) + ) + if child_result.rowcount != 1: + await session.rollback() + return 0 + released += 1 + + # Rolling upgrades can leave aggregate reservations created before durable + # reservation rows existed. Release only stale aggregates that have no active + # durable owner; targeted refund cleanup also heals legacy NULL timestamps. + legacy_query = select(ApiKey).where(col(ApiKey.reserved_balance) > 0) + if key_hash is None: + legacy_query = legacy_query.where(col(ApiKey.reserved_at).is_not(None)).where( + col(ApiKey.reserved_at) < cutoff + ) + else: + legacy_query = legacy_query.where( + or_( + col(ApiKey.hashed_key) == key_hash, + col(ApiKey.parent_key_hash) == key_hash, + ) + ).where( + or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff) + ) + + for legacy_key in (await session.exec(legacy_query)).all(): + active_owner = ( + await session.exec( + select(ReservationRelease.id) + .where(col(ReservationRelease.status) == "active") + .where( + or_( + col(ReservationRelease.key_hash) == legacy_key.hashed_key, + col(ReservationRelease.billing_key_hash) + == legacy_key.hashed_key, + ) + ) + .limit(1) + ) + ).first() + if active_owner is not None: + continue + legacy_key.reserved_balance = 0 + legacy_key.reserved_at = None + session.add(legacy_key) + released += 1 + await session.commit() - released = int(result.rowcount or 0) if released: logger.warning( - "Released stale balance reservations", - extra={"released_keys": released, "max_age_seconds": max_age_seconds}, + "Released stale reservations", + extra={"released_reservations": released, "max_age_seconds": max_age_seconds}, ) return released @@ -434,6 +532,20 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore ) +class ReservationRelease(SQLModel, table=True): # type: ignore + __tablename__ = "reservation_releases" + __table_args__ = ( + Index("ix_reservation_releases_status_created_at", "status", "created_at"), + ) + + id: str = Field(primary_key=True) + key_hash: str = Field(index=True) + billing_key_hash: str = Field(index=True) + reserved_msats: int + status: str = Field(default="active") + created_at: int = Field(default_factory=lambda: int(time.time())) + + class RoutstrFee(SQLModel, table=True): # type: ignore __tablename__ = "routstr_fees" id: int = Field(default=1, primary_key=True) diff --git a/routstr/proxy.py b/routstr/proxy.py index d0228c72..c0b8794f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -7,7 +7,13 @@ from fastapi.responses import Response, StreamingResponse from sqlmodel import select from .algorithm import create_model_mappings -from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key +from .auth import ( + ReservationSnapshot, + get_reservation_snapshot, + pay_for_request, + revert_pay_for_request, + validate_bearer_key, +) from .core import get_logger from .core.db import ( ApiKey, @@ -440,8 +446,10 @@ async def proxy( "upstream_error", "All upstreams failed", 502, request=request ) + reservation_snapshot: ReservationSnapshot | None = None if is_ehbp or request_body_dict: await pay_for_request(key, max_cost_for_model, session) + reservation_snapshot = await get_reservation_snapshot(key, session) # Tracks request params already removed in response to upstream rejections, # shared across providers so a stripped param stays stripped on failover and @@ -463,14 +471,18 @@ async def proxy( ) candidate_max = max(candidate_max, settings.min_request_msat) if candidate_max > max_cost_for_model: - await revert_pay_for_request(key, session, max_cost_for_model) + await revert_pay_for_request( + key, session, max_cost_for_model, reservation_snapshot + ) try: await pay_for_request(key, candidate_max, session) except HTTPException: if i == len(candidates) - 1: raise await pay_for_request(key, max_cost_for_model, session) + reservation_snapshot = await get_reservation_snapshot(key, session) continue + reservation_snapshot = await get_reservation_snapshot(key, session) max_cost_for_model = candidate_max headers = upstream.prepare_headers(dict(request.headers)) @@ -499,6 +511,7 @@ async def proxy( max_cost_for_model=max_cost_for_model, session=session, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) elif is_responses_api: response = await upstream.forward_responses_request( @@ -510,6 +523,7 @@ async def proxy( max_cost_for_model, session, model_obj, + reservation_snapshot, ) else: response = await upstream.forward_request( @@ -521,6 +535,7 @@ async def proxy( max_cost_for_model, session, model_obj, + reservation_snapshot, ) except UpstreamError: # Let the outer UpstreamError handler manage retry/revert @@ -537,7 +552,9 @@ async def proxy( "max_cost_for_model": max_cost_for_model, }, ) - await revert_pay_for_request(key, session, max_cost_for_model) + await revert_pay_for_request( + key, session, max_cost_for_model, reservation_snapshot + ) raise # Reactive recovery: some models reject one specific request @@ -606,7 +623,9 @@ async def proxy( continue # 4xx error (user error), or other non-retryable error, or last provider failed - await revert_pay_for_request(key, session, max_cost_for_model) + await revert_pay_for_request( + key, session, max_cost_for_model, reservation_snapshot + ) logger.warning( "Upstream request failed, revert payment " "(provider=%s model=%s status=%s path=%s)", @@ -638,8 +657,10 @@ async def proxy( "max_cost_for_model": max_cost_for_model, }, ) - await asyncio.shield( - revert_pay_for_request(key, session, max_cost_for_model) + # The cancellation has been caught, so complete exact cleanup in + # this task before the request-scoped session can be torn down. + await revert_pay_for_request( + key, session, max_cost_for_model, reservation_snapshot ) raise @@ -659,7 +680,9 @@ async def proxy( # If this was the last provider if i == len(candidates) - 1: - await revert_pay_for_request(key, session, max_cost_for_model) + await revert_pay_for_request( + key, session, max_cost_for_model, reservation_snapshot + ) return create_upstream_error_response(e, request) # Otherwise loop continues to next provider diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 799d181b..a8dba7e3 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -14,7 +14,12 @@ from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel -from ..auth import adjust_payment_for_tokens +from ..auth import ( + ReservationSnapshot, + adjust_payment_for_tokens, + get_reservation_snapshot, + release_reservation, +) from ..core import get_logger from ..core.db import ( ApiKey, @@ -65,8 +70,7 @@ logger = get_logger(__name__) def _is_json_content_type(content_type: str | None) -> bool: - """Return True when the upstream response should be parsed as JSON. - """ + """Return True when the upstream response should be parsed as JSON.""" if not content_type: return False main = content_type.split(";", 1)[0].strip().lower() @@ -229,9 +233,7 @@ class BaseUpstreamProvider: pass if "prompt_tokens" in usage: try: - usage["prompt_tokens"] = ( - int(usage.get("prompt_tokens") or 0) + extra - ) + usage["prompt_tokens"] = int(usage.get("prompt_tokens") or 0) + extra except (TypeError, ValueError): pass @@ -733,7 +735,9 @@ class BaseUpstreamProvider: and rate_limit.retry_after_seconds is not None and "retry-after" not in {k.lower() for k in headers} ): - headers["Retry-After"] = str(max(1, math.ceil(rate_limit.retry_after_seconds))) + headers["Retry-After"] = str( + max(1, math.ceil(rate_limit.retry_after_seconds)) + ) if is_json_body: if not content_type: @@ -792,6 +796,56 @@ class BaseUpstreamProvider: media_type="application/json", ) + async def _release_failed_streaming_reservation( + self, + key: ApiKey, + session: AsyncSession, + reservation_snapshot: ReservationSnapshot | None, + ) -> bool: + """Attempt exact release and suppress unsafe settlement retries.""" + try: + await session.rollback() + snapshot = reservation_snapshot + if snapshot is None: + snapshot = await get_reservation_snapshot(key, session) + released = await release_reservation( + snapshot, + session, + snapshot.reserved_msats, + ) + if not released: + logger.critical( + "Billing reservation could not be released", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "reserved_balance": snapshot.reserved_msats, + }, + ) + # A failed release remains recoverable by the stale-reservation + # sweep. Retrying settlement here could charge after an ambiguous + # database failure or replace the original stream exception. + return True + except asyncio.CancelledError: + # Preserve the exception that triggered billing cleanup. The stream + # propagates it immediately after this helper returns, and stale + # reservation cleanup can recover an interrupted release. + logger.critical( + "Billing reservation release was cancelled", + extra={"key_hash": key.hashed_key[:8] + "..."}, + exc_info=True, + ) + return True + except Exception as release_error: + logger.critical( + "Billing reservation release failed", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(release_error), + }, + exc_info=True, + ) + return True + async def handle_streaming_chat_completion( self, response: httpx.Response, @@ -800,6 +854,7 @@ class BaseUpstreamProvider: background_tasks: BackgroundTasks, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -811,6 +866,15 @@ class BaseUpstreamProvider: Returns: StreamingResponse with cost data injected at the end """ + if reservation_snapshot is None: + async with create_session() as snapshot_session: + snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key) + if snapshot_key is None: + raise RuntimeError("Billing key disappeared before streaming") + reservation_snapshot = await get_reservation_snapshot( + snapshot_key, snapshot_session + ) + logger.debug( "Processing streaming chat completion", extra={ @@ -844,6 +908,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True except Exception: @@ -968,9 +1033,7 @@ class BaseUpstreamProvider: # line so multi-line ``data`` stays valid SSE framing - a bare # second line would otherwise reach the client without its # ``data:`` field and break naive parsers. - body = b"".join( - b"data: " + ln + b"\n" for ln in data.split(b"\n") - ) + body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n")) yield prefix + body + b"\n" try: @@ -1018,48 +1081,44 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True - except Exception as e: - logger.exception( - "Error during usage finalization", + except BaseException as e: + logger.critical( + "Error during usage finalization — CRITICAL", extra={ "key_hash": key.hashed_key[:8] + "...", "error": str(e), }, + exc_info=True, ) - - # Fall back so we still emit a non-zero sats cost downstream. - cost_data = { - "base_msats": 0, - "input_msats": 0, - "output_msats": 0, - "total_msats": 0, - "total_usd": 0.0, - "input_tokens": 0, - "output_tokens": 0, - } + # Release is a terminal billing state. Do not enqueue + # finalize_db_only from the generator's finally block + # and charge this request later. + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + session, + reservation_snapshot, + ) + ) + raise if usage_chunk_data is None: if not hasattr(self, "_current_stream_id"): - self._current_stream_id = ( - f"chatcmpl-{uuid.uuid4()}" - ) + self._current_stream_id = f"chatcmpl-{uuid.uuid4()}" usage_chunk_data = { "id": self._current_stream_id, "object": "chat.completion.chunk", "model": last_model_seen or "unknown", "choices": [], "usage": { - "prompt_tokens": cost_data.get( - "input_tokens", 0 - ), + "prompt_tokens": cost_data.get("input_tokens", 0), "completion_tokens": cost_data.get( "output_tokens", 0 ), - "total_tokens": cost_data.get( - "input_tokens", 0 - ) + "total_tokens": cost_data.get("input_tokens", 0) + cost_data.get("output_tokens", 0), }, } @@ -1115,6 +1174,7 @@ class BaseUpstreamProvider: deducted_max_cost: int, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response: """Handle non-streaming chat completion responses with token usage tracking and cost adjustment. @@ -1163,6 +1223,7 @@ class BaseUpstreamProvider: deducted_max_cost, model_obj, self.provider_fee, + reservation_snapshot, ) await session.refresh(key) @@ -1259,6 +1320,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1304,6 +1366,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True except Exception: @@ -1388,9 +1451,7 @@ class BaseUpstreamProvider: return # Re-prefix each line so multi-line ``data`` stays valid SSE # framing for the client. - body = b"".join( - b"data: " + ln + b"\n" for ln in data.split(b"\n") - ) + body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n")) yield prefix + body + b"\n" try: @@ -1435,25 +1496,26 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True - except Exception as e: - logger.exception( - "Error during Responses API usage finalization", + except BaseException as e: + logger.critical( + "Error during Responses API usage finalization — CRITICAL", extra={ "key_hash": key.hashed_key[:8] + "...", "error": str(e), }, + exc_info=True, ) - cost_data = { - "base_msats": 0, - "input_msats": 0, - "output_msats": 0, - "total_msats": 0, - "total_usd": 0.0, - "input_tokens": 0, - "output_tokens": 0, - } + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + session, + reservation_snapshot, + ) + ) + raise if usage_chunk_data is None: usage_chunk_data = { @@ -1467,22 +1529,14 @@ class BaseUpstreamProvider: "output_tokens": cost_data.get( "output_tokens", 0 ), - "total_tokens": cost_data.get( - "input_tokens", 0 - ) + "total_tokens": cost_data.get("input_tokens", 0) + cost_data.get("output_tokens", 0), }, }, "usage": { - "input_tokens": cost_data.get( - "input_tokens", 0 - ), - "output_tokens": cost_data.get( - "output_tokens", 0 - ), - "total_tokens": cost_data.get( - "input_tokens", 0 - ) + "input_tokens": cost_data.get("input_tokens", 0), + "output_tokens": cost_data.get("output_tokens", 0), + "total_tokens": cost_data.get("input_tokens", 0) + cost_data.get("output_tokens", 0), }, } @@ -1498,9 +1552,9 @@ class BaseUpstreamProvider: usage_chunk_data["response"]["usage"]["cost"] = ( cost_data.get("total_usd", 0.0) ) - usage_chunk_data["response"]["usage"][ - "cost_sats" - ] = sats_cost + usage_chunk_data["response"]["usage"]["cost_sats"] = ( + sats_cost + ) usage_chunk_data["response"]["usage"][ "remaining_balance_msats" ] = remaining_balance_msats @@ -1554,6 +1608,7 @@ class BaseUpstreamProvider: deducted_max_cost: int, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response: """Handle non-streaming Responses API responses with token usage tracking and cost adjustment. @@ -1605,6 +1660,7 @@ class BaseUpstreamProvider: deducted_max_cost, model_obj, self.provider_fee, + reservation_snapshot, ) await session.refresh(key) @@ -1695,7 +1751,13 @@ class BaseUpstreamProvider: raise async def _finalize_generic_streaming_payment( - self, key_hash: str, max_cost: int, path: str + self, + key_hash: str, + max_cost: int, + path: str, + model_obj: Model | None, + provider_fee: float | None, + reservation_snapshot: ReservationSnapshot, ) -> None: """Background task to finalize payment for generic streaming requests.""" async with create_session() as session: @@ -1716,8 +1778,9 @@ class BaseUpstreamProvider: {"model": "unknown", "usage": None}, session, max_cost, - model_obj=None, - provider_fee=None, + model_obj=model_obj, + provider_fee=provider_fee, + reservation_snapshot=reservation_snapshot, ) logger.debug( "Finalized generic streaming payment in background", @@ -1743,6 +1806,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> StreamingResponse: async def stream_with_cost( max_cost_for_model: int, @@ -1785,9 +1849,7 @@ class BaseUpstreamProvider: _coerce_usd(cd.get("output_cost")), ) for field in ("total_cost", "cost"): - total_cost = max( - total_cost, _coerce_usd(usage_or_root.get(field)) - ) + total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field))) async def finalize_without_usage() -> bytes | None: nonlocal usage_finalized @@ -1810,12 +1872,27 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() - except Exception: - usage_finalized = True - return None + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise try: async for chunk in response.aiter_bytes(): @@ -1833,9 +1910,7 @@ class BaseUpstreamProvider: if msg and msg.get("model"): last_model_seen = str(msg.get("model")) - provider_added = ( - "provider" not in data - ) + provider_added = "provider" not in data self._apply_provider_field(data) if requested_model: @@ -1963,6 +2038,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) self.inject_cost_metadata( @@ -1972,8 +2048,23 @@ class BaseUpstreamProvider: usage_finalized = True # Emit the full combined_data as the cost yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode() - except Exception: - pass + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise if not usage_finalized: maybe_cost_event = await finalize_without_usage() @@ -2011,6 +2102,7 @@ class BaseUpstreamProvider: path: str, requested_model: str | None = None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response: try: content = await response.aread() @@ -2037,6 +2129,7 @@ class BaseUpstreamProvider: deducted_max_cost, model_obj, self.provider_fee, + reservation_snapshot, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2084,9 +2177,7 @@ class BaseUpstreamProvider: async def _aggregate_anthropic_events_to_message( self, iterator: AsyncIterator[Any] ) -> dict: - return await messages_dispatch.aggregate_anthropic_events_to_message( - iterator - ) + return await messages_dispatch.aggregate_anthropic_events_to_message(iterator) async def _dispatch_anthropic_messages( self, @@ -2112,6 +2203,7 @@ class BaseUpstreamProvider: session: AsyncSession, max_cost_for_model: int, model_obj: Model, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response | StreamingResponse: """Translate /v1/messages to upstream chat/completions via litellm. @@ -2132,6 +2224,7 @@ class BaseUpstreamProvider: max_cost_for_model, requested_model, model_obj, + reservation_snapshot, ) response_json = messages_dispatch.coerce_litellm_payload(result) @@ -2145,6 +2238,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2196,8 +2290,10 @@ class BaseUpstreamProvider: response_json, max_cost_for_model, model_obj ) - if cost_data and "usage" in response_json and isinstance( - response_json["usage"], dict + if ( + cost_data + and "usage" in response_json + and isinstance(response_json["usage"], dict) ): response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 self._fold_cache_into_input_tokens(response_json["usage"]) @@ -2240,6 +2336,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, requested_model: str | None, model_obj: Model | None = None, + reservation_snapshot: ReservationSnapshot | None = None, ) -> StreamingResponse: """Re-emit a litellm Anthropic-event iterator as live SSE bytes with cost reconciliation appended at end of stream.""" @@ -2274,9 +2371,7 @@ class BaseUpstreamProvider: }, ) async with create_session() as new_session: - fresh_key = await new_session.get( - key.__class__, key.hashed_key - ) + fresh_key = await new_session.get(key.__class__, key.hashed_key) if not fresh_key: usage_finalized = True return None @@ -2292,15 +2387,29 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) usage_finalized = True return ( - f"event: cost\ndata: " - f"{json.dumps({'cost': cost_data})}\n\n" + f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" ).encode() - except Exception: - usage_finalized = True - return None + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise try: async for annotated in messages_dispatch.stream_annotated_events( @@ -2335,9 +2444,7 @@ class BaseUpstreamProvider: or total_cost > 0 ): async with create_session() as new_session: - fresh_key = await new_session.get( - key.__class__, key.hashed_key - ) + fresh_key = await new_session.get(key.__class__, key.hashed_key) if fresh_key: try: rebuilt_usage: dict = { @@ -2367,6 +2474,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + reservation_snapshot, ) self.inject_cost_metadata( combined_data, cost_data, fresh_key @@ -2376,8 +2484,23 @@ class BaseUpstreamProvider: f"event: cost\ndata: " f"{json.dumps({'cost': cost_data})}\n\n" ).encode() - except Exception: - pass + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise if not usage_finalized: cost_event = await finalize_without_usage() @@ -2388,6 +2511,9 @@ class BaseUpstreamProvider: if not usage_finalized: await finalize_without_usage() raise + finally: + if not usage_finalized: + await finalize_without_usage() return StreamingResponse( stream_with_cost(), @@ -2511,8 +2637,7 @@ class BaseUpstreamProvider: ) response_headers["X-Cashu"] = refund_token logger.info( - "Refund processed for streaming /v1/messages " - "via litellm", + "Refund processed for streaming /v1/messages via litellm", extra={ "refund_amount": refund_amount, "unit": unit, @@ -2550,6 +2675,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, session: AsyncSession, model_obj: Model, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response | StreamingResponse: """Forward authenticated request to upstream service with cost tracking. @@ -2584,6 +2710,7 @@ class BaseUpstreamProvider: session=session, max_cost_for_model=max_cost_for_model, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) url = self.build_request_url(path, model_obj) @@ -2714,6 +2841,7 @@ class BaseUpstreamProvider: max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2731,6 +2859,7 @@ class BaseUpstreamProvider: path, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) finally: await response.aclose() @@ -2747,6 +2876,7 @@ class BaseUpstreamProvider: path, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) finally: await response.aclose() @@ -2797,6 +2927,7 @@ class BaseUpstreamProvider: background_tasks, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) result.background = background_tasks return result @@ -2811,11 +2942,15 @@ class BaseUpstreamProvider: max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) finally: await response.aclose() await client.aclose() + if reservation_snapshot is None: + reservation_snapshot = await get_reservation_snapshot(key, session) + background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) @@ -2824,6 +2959,9 @@ class BaseUpstreamProvider: key.hashed_key, max_cost_for_model, path, + model_obj, + self.provider_fee, + reservation_snapshot, ) logger.debug( @@ -2898,7 +3036,9 @@ class BaseUpstreamProvider: supports_ehbp: bool = False - def get_confidential_inference_profile(self) -> "ConfidentialInferenceProfile | None": + def get_confidential_inference_profile( + self, + ) -> "ConfidentialInferenceProfile | None": """Return provider policy for encrypted/confidential inference forwarding.""" return None @@ -2926,6 +3066,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, session: AsyncSession, model_obj: Model, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response | StreamingResponse: """Forward authenticated Responses API request to upstream service with cost tracking. @@ -3064,6 +3205,7 @@ class BaseUpstreamProvider: max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -3080,11 +3222,15 @@ class BaseUpstreamProvider: max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, + reservation_snapshot=reservation_snapshot, ) finally: await response.aclose() await client.aclose() + if reservation_snapshot is None: + reservation_snapshot = await get_reservation_snapshot(key, session) + background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) @@ -3093,6 +3239,9 @@ class BaseUpstreamProvider: key.hashed_key, max_cost_for_model, path, + model_obj, + self.provider_fee, + reservation_snapshot, ) logger.debug( @@ -3573,14 +3722,8 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if ( - cost_data - and "usage" in data_json - and data_json["usage"] - ): - data_json["usage"]["cost_sats"] = ( - cost_data.total_msats // 1000 - ) + if cost_data and "usage" in data_json and data_json["usage"]: + data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -4560,14 +4703,8 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if ( - cost_data - and "usage" in data_json - and data_json["usage"] - ): - data_json["usage"]["cost_sats"] = ( - cost_data.total_msats // 1000 - ) + if cost_data and "usage" in data_json and data_json["usage"]: + data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -5059,7 +5196,9 @@ class BaseUpstreamProvider: except Exception: self._models_cache = models_with_fees - self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} + self._models_by_id = { + m.forwarded_model_id or m.id: m for m in self._models_cache + } except Exception as e: logger.error( diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 213be7af..96955492 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -15,7 +15,11 @@ from sqlmodel import col, update from ..auth import ( ROUTSTR_FEE_PERCENT, + ReservationSnapshot, + _claim_reservation_for_charge, + _validate_reservation_snapshot, get_billing_key, + get_reservation_snapshot, payments_logger, ) from ..core import get_logger @@ -502,8 +506,14 @@ async def finalize_ehbp_actual_cost_payment( reserved_cost_for_model: int, model_id: str, cost_info: dict, + reservation_snapshot: ReservationSnapshot | None = None, ) -> None: """Finalize an EHBP bearer request using clamped provider usage metrics.""" + reservation = reservation_snapshot or await get_reservation_snapshot(key, session) + await _validate_reservation_snapshot(key, reservation, session) + if not await _claim_reservation_for_charge(reservation, session): + return + reserved_cost_for_model = reservation.reserved_msats billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key @@ -606,6 +616,7 @@ async def finalize_ehbp_max_cost_payment( session: AsyncSession, max_cost_for_model: int, model_id: str, + reservation_snapshot: ReservationSnapshot | None = None, ) -> None: """Finalize an EHBP bearer request by charging the reserved max cost. @@ -613,6 +624,11 @@ async def finalize_ehbp_max_cost_payment( normal completion handlers, this intentionally charges the pre-reserved max cost and releases the reservation. """ + reservation = reservation_snapshot or await get_reservation_snapshot(key, session) + await _validate_reservation_snapshot(key, reservation, session) + if not await _claim_reservation_for_charge(reservation, session): + return + max_cost_for_model = reservation.reserved_msats billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key @@ -766,6 +782,7 @@ async def forward_ehbp_request( max_cost_for_model: int, session: AsyncSession, model_obj: Model, + reservation_snapshot: ReservationSnapshot | None = None, ) -> Response | StreamingResponse: """Forward an EHBP bearer-auth request and finalize billing. @@ -883,7 +900,12 @@ async def forward_ehbp_request( # the requested model. billing_model = cost_info.pop("actual_model", None) or model_obj.id await finalize_ehbp_actual_cost_payment( - key, session, max_cost_for_model, billing_model, cost_info + key, + session, + max_cost_for_model, + billing_model, + cost_info, + reservation_snapshot, ) cost_data = {**cost_info, "total_usd": 0.0} else: @@ -897,7 +919,11 @@ async def forward_ehbp_request( }, ) await finalize_ehbp_max_cost_payment( - key, session, max_cost_for_model, model_obj.id + key, + session, + max_cost_for_model, + model_obj.id, + reservation_snapshot, ) cost_data = { "total_msats": max_cost_for_model, diff --git a/tests/integration/test_balance_negative_on_cost_overrun.py b/tests/integration/test_balance_negative_on_cost_overrun.py index e7b8fab8..cda891d4 100644 --- a/tests/integration/test_balance_negative_on_cost_overrun.py +++ b/tests/integration/test_balance_negative_on_cost_overrun.py @@ -15,6 +15,7 @@ from unittest.mock import patch import pytest from sqlmodel.ext.asyncio.session import AsyncSession +from routstr.auth import ReservationSnapshot from routstr.core.db import ApiKey from routstr.payment.cost_calculation import CostData @@ -23,7 +24,7 @@ def _make_key(balance: int, reserved: int) -> ApiKey: return ApiKey( hashed_key=f"test_{uuid.uuid4().hex}", balance=balance, - reserved_balance=reserved, + reserved_balance=0, total_spent=0, total_requests=1, ) @@ -75,6 +76,8 @@ async def test_balance_never_negative_when_cost_exceeds_reservation( key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost) integration_session.add(key) await integration_session.commit() + from routstr.auth import pay_for_request + await pay_for_request(key, deducted_max_cost, integration_session) response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}} @@ -111,6 +114,8 @@ async def test_balance_floor_at_zero_on_overrun( key = _make_key(balance=500, reserved=500) integration_session.add(key) await integration_session.commit() + from routstr.auth import pay_for_request + await pay_for_request(key, deducted_max_cost, integration_session) response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}} @@ -152,6 +157,8 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun( key = _make_key(balance=2000, reserved=990) integration_session.add(key) await integration_session.commit() + from routstr.auth import pay_for_request + await pay_for_request(key, deducted_max_cost, integration_session) response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}} @@ -190,7 +197,11 @@ async def test_concurrent_cost_overruns_never_negative( """Concurrent finalization with cost overruns must never produce negative balance.""" import asyncio - from routstr.auth import adjust_payment_for_tokens, pay_for_request + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) from routstr.core.db import create_session deducted_max_cost = 990 @@ -216,12 +227,14 @@ async def test_concurrent_cost_overruns_never_negative( async with create_session() as session: key_to_reserve = await session.get(ApiKey, key_hash) assert key_to_reserve is not None + reservations = [] for _ in range(n_requests): await pay_for_request(key_to_reserve, deducted_max_cost, session) + reservations.append(await get_reservation_snapshot(key_to_reserve, session)) await session.refresh(key_to_reserve) # Now finalize all concurrently with cost overrun - async def finalize() -> None: + async def finalize(reservation: ReservationSnapshot) -> None: response_data = { "model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}, @@ -230,7 +243,11 @@ async def test_concurrent_cost_overruns_never_negative( fresh_key = await session.get(ApiKey, key_hash) assert fresh_key is not None await adjust_payment_for_tokens( - fresh_key, response_data, session, deducted_max_cost, None, None + fresh_key, + response_data, + session, + deducted_max_cost, + reservation_snapshot=reservation, ) # Patch once around the gather: entering the same patch target from @@ -240,7 +257,7 @@ async def test_concurrent_cost_overruns_never_negative( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await asyncio.gather(*[finalize() for _ in range(n_requests)]) + await asyncio.gather(*(finalize(r) for r in reservations)) async with create_session() as session: final_key = await session.get(ApiKey, key_hash) @@ -281,6 +298,8 @@ async def test_zero_free_balance_overrun_is_safe( key = _make_key(balance=1000, reserved=1000) integration_session.add(key) await integration_session.commit() + from routstr.auth import pay_for_request + await pay_for_request(key, deducted_max_cost, integration_session) response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}} @@ -319,7 +338,11 @@ async def test_parallel_requests_no_free_inference( """Second parallel finalization must be charged even when first depleted free balance.""" import asyncio - from routstr.auth import adjust_payment_for_tokens + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) from routstr.core.db import create_session deducted_max_cost = 100 @@ -340,14 +363,18 @@ async def test_parallel_requests_no_free_inference( key = ApiKey( hashed_key=key_hash, balance=starting_balance, - reserved_balance=deducted_max_cost * 2, # both slots pre-reserved + reserved_balance=0, total_spent=0, total_requests=2, ) session.add(key) await session.commit() + reservations = [] + for _ in range(2): + await pay_for_request(key, deducted_max_cost, session) + reservations.append(await get_reservation_snapshot(key, session)) - async def finalize() -> None: + async def finalize(reservation: ReservationSnapshot) -> None: response_data = { "model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}, @@ -356,7 +383,11 @@ async def test_parallel_requests_no_free_inference( fresh_key = await session.get(ApiKey, key_hash) assert fresh_key is not None await adjust_payment_for_tokens( - fresh_key, response_data, session, deducted_max_cost, None, None + fresh_key, + response_data, + session, + deducted_max_cost, + reservation_snapshot=reservation, ) # Patch once around the gather: entering the same patch target from two @@ -366,7 +397,7 @@ async def test_parallel_requests_no_free_inference( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await asyncio.gather(finalize(), finalize()) + await asyncio.gather(*(finalize(r) for r in reservations)) async with create_session() as session: final_key = await session.get(ApiKey, key_hash) diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index 8633e075..9d45d475 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -14,13 +14,17 @@ from unittest.mock import patch import httpx import pytest from httpx import AsyncClient +from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession +from routstr.core.db import ApiKey, ReservationRelease from routstr.payment.models import Architecture, Model, Pricing from routstr.proxy import refresh_model_maps from routstr.upstream.base import BaseUpstreamProvider CHEAP_BASE_URL = "https://cheap.example.com/v1" EXPENSIVE_BASE_URL = "https://expensive.example.com/v1" +THIRD_BASE_URL = "https://third.example.com/v1" def _make_model( @@ -55,9 +59,7 @@ def _make_model( class _StaticProvider(BaseUpstreamProvider): """Upstream provider with a fixed model catalog and no remote refresh.""" - def __init__( - self, base_url: str, api_key: str, fee: float, model: Model - ) -> None: + def __init__(self, base_url: str, api_key: str, fee: float, model: Model) -> None: super().__init__(base_url, api_key, fee) self.provider_type = "custom" self._static_model = model @@ -98,7 +100,7 @@ async def dual_provider_maps( EXPENSIVE_BASE_URL, "key-expensive", 1.0, - _make_model("provb/dual-model", 0.005, 0.010), + _make_model("provb/dual-model", 0.005, 0.010, max_cost=100.0), ) async for _ in _install_providers([cheap, expensive]): yield cheap, expensive @@ -142,6 +144,7 @@ def _upstream_response(request: httpx.Request) -> httpx.Response: async def test_failover_serve_billed_at_serving_providers_rate( authenticated_client: AsyncClient, dual_provider_maps: tuple[_StaticProvider, _StaticProvider], + integration_session: AsyncSession, ) -> None: """A fallback serve is billed at the fallback's price, not the winner's. @@ -199,6 +202,17 @@ async def test_failover_serve_billed_at_serving_providers_rate( # Billed at the serving provider's rate: 1000/1000*5000 + 500/1000*10000. assert payload["cost"]["total_msats"] == 10_000 + # The fallback's larger max-cost envelope requires a replacement + # reservation. The failed candidate is released, the serving candidate is + # charged, and no request-owned reservation remains active. + key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined] + records = ( + await integration_session.exec( + select(ReservationRelease).where(ReservationRelease.key_hash == key_hash) + ) + ).all() + assert sorted(record.status for record in records) == ["charged", "released"] + @pytest.fixture async def same_id_provider_maps( @@ -349,9 +363,7 @@ async def test_usd_cost_serve_carries_serving_providers_fee( if request.url.host == "cheap.example.com": return httpx.Response( 502, - content=json.dumps( - {"error": {"message": "bad gateway"}} - ).encode(), + content=json.dumps({"error": {"message": "bad gateway"}}).encode(), headers={"content-type": "application/json"}, ) body = { @@ -478,6 +490,101 @@ async def test_failover_beyond_balance_envelope_is_rejected( assert [r.url.host for r in sent_requests] == ["cheap.example.com"] +@pytest.fixture +async def three_candidate_child_maps( + patched_db_engine: None, +) -> AsyncGenerator[None, None]: + """Second candidate cannot fit the child limit; third restores and serves.""" + first = _StaticProvider( + CHEAP_BASE_URL, + "key-first", + 1.0, + _make_model("dual-model", 0.001, 0.002, max_cost=50.0), + ) + too_large = _StaticProvider( + EXPENSIVE_BASE_URL, + "key-too-large", + 1.0, + _make_model("dual-model", 0.002, 0.003, max_cost=100.0), + ) + third = _StaticProvider( + THIRD_BASE_URL, + "key-third", + 1.0, + _make_model("dual-model", 0.003, 0.004, max_cost=50.0), + ) + async for _ in _install_providers([first, too_large, third]): + yield + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring( + authenticated_client: AsyncClient, + three_candidate_child_maps: None, + integration_session: AsyncSession, +) -> None: + """A failed child guard cannot leak its parent update into restoration.""" + key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined] + child = await integration_session.get(ApiKey, key_hash) + assert child is not None + parent = ApiKey(hashed_key="failover-parent", balance=10_000_000) + child.parent_key_hash = parent.hashed_key + child.balance_limit = 75_000 + integration_session.add(parent) + integration_session.add(child) + await integration_session.commit() + + sent_requests: list[httpx.Request] = [] + + async def fake_transport( + request: httpx.Request, *args: Any, **kwargs: Any + ) -> httpx.Response: + sent_requests.append(request) + return _upstream_response(request) + + with ( + patch( + "httpx.AsyncHTTPTransport.handle_async_request", + side_effect=fake_transport, + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=0.0005, + ), + ): + response = await authenticated_client.post( + "/v1/chat/completions", + json={ + "model": "dual-model", + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert response.status_code == 200 + # The 100-sat candidate is rejected before forwarding; the third serves. + assert [request.url.host for request in sent_requests] == [ + "cheap.example.com", + "third.example.com", + ] + + await integration_session.refresh(parent) + await integration_session.refresh(child) + assert parent.reserved_balance == 0 + assert child.reserved_balance == 0 + assert parent.total_spent == response.json()["cost"]["total_msats"] + + records = ( + await integration_session.exec( + select(ReservationRelease).where(ReservationRelease.key_hash == key_hash) + ) + ).all() + assert len(records) == 2 + assert sorted(record.status for record in records) == ["charged", "released"] + assert len({record.reserved_msats for record in records}) == 1 + assert all(record.status != "active" for record in records) + + @pytest.fixture async def raised_envelope_provider_maps( patched_db_engine: None, @@ -504,6 +611,7 @@ async def raised_envelope_provider_maps( async def test_failover_reserves_serving_candidates_envelope( authenticated_client: AsyncClient, raised_envelope_provider_maps: None, + integration_session: AsyncSession, ) -> None: """An affordable pricier fallback is re-reserved, served, and billed. @@ -544,3 +652,15 @@ async def test_failover_reserves_serving_candidates_envelope( "expensive.example.com", ] assert response.json()["cost"]["total_msats"] == 10_000 + + key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined] + records = ( + await integration_session.exec( + select(ReservationRelease).where(ReservationRelease.key_hash == key_hash) + ) + ).all() + assert len(records) == 2 + released = next(record for record in records if record.status == "released") + charged = next(record for record in records if record.status == "charged") + assert charged.reserved_msats > released.reserved_msats + assert all(record.status != "active" for record in records) diff --git a/tests/integration/test_free_response_stale_reservation.py b/tests/integration/test_free_response_stale_reservation.py index 00f3c924..dadb6a6a 100644 --- a/tests/integration/test_free_response_stale_reservation.py +++ b/tests/integration/test_free_response_stale_reservation.py @@ -38,7 +38,7 @@ async def test_overrun_charges_after_reservation_swept( integration_session: AsyncSession, ) -> None: """Overrun finalize must charge even when the reservation was already released.""" - from routstr.auth import adjust_payment_for_tokens + from routstr.auth import adjust_payment_for_tokens, pay_for_request deducted_max_cost = 990 # discounted reservation actual_token_cost = 1000 # actual cost overruns the reservation @@ -47,6 +47,10 @@ async def test_overrun_charges_after_reservation_swept( key = _make_key(balance=1000, reserved=0) integration_session.add(key) await integration_session.commit() + await pay_for_request(key, deducted_max_cost, integration_session) + key.reserved_balance = 0 + integration_session.add(key) + await integration_session.commit() response_data = { "model": "test-model", @@ -79,8 +83,16 @@ async def test_free_response_path_closed_end_to_end( patched_db_engine: None, ) -> None: """A reservation released by the real sweeper must not yield a free response.""" - from routstr.auth import adjust_payment_for_tokens, pay_for_request - from routstr.core.db import create_session, release_stale_reservations + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import ( + ReservationRelease, + create_session, + release_stale_reservations, + ) deducted_max_cost = 990 actual_token_cost = 1000 @@ -104,10 +116,15 @@ async def test_free_response_path_closed_end_to_end( key = await session.get(ApiKey, key_hash) assert key is not None await pay_for_request(key, deducted_max_cost, session) + snapshot = await get_reservation_snapshot(key, session) await session.refresh(key) assert key.reserved_balance == deducted_max_cost key.reserved_at = int(time.time()) - 10_000 + record = await session.get(ReservationRelease, snapshot.release_id) + assert record is not None + record.created_at = int(time.time()) - 10_000 session.add(key) + session.add(record) await session.commit() # Sweeper releases the stale reservation without charging. @@ -129,18 +146,20 @@ async def test_free_response_path_closed_end_to_end( return_value=_cost_data(actual_token_cost), ): await adjust_payment_for_tokens( - key, response_data, session, deducted_max_cost, None, None + key, + response_data, + session, + deducted_max_cost, + reservation_snapshot=snapshot, ) async with create_session() as session: final = await session.get(ApiKey, key_hash) assert final is not None - assert final.total_spent == actual_token_cost, ( - f"Free response: total_spent={final.total_spent}, expected {actual_token_cost}" - ) - assert final.balance == 1000 - actual_token_cost, ( - f"Balance not charged after sweep: {final.balance}" - ) + # Stale release is terminal for this reservation. A late finalizer must not + # charge aggregate balance that may now belong to a newer request. + assert final.total_spent == 0 + assert final.balance == 1000 assert final.balance >= 0 assert final.reserved_balance == 0 diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py index b283d55b..8ffebc86 100644 --- a/tests/integration/test_reserved_balance_negative.py +++ b/tests/integration/test_reserved_balance_negative.py @@ -141,7 +141,7 @@ async def test_revert_with_zero_reserved_balance_is_noop( 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 + from routstr.auth import pay_for_request, revert_pay_for_request unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}" test_key = ApiKey( @@ -151,8 +151,12 @@ async def test_revert_with_zero_reserved_balance_is_noop( ) integration_session.add(test_key) await integration_session.commit() + await pay_for_request(test_key, 100, integration_session) + test_key.reserved_balance = 0 + integration_session.add(test_key) + await integration_session.commit() - # Try to revert more than available — should be a no-op + # A stale cleanup already released the aggregate reservation. result = await revert_pay_for_request(test_key, integration_session, 100) await integration_session.refresh(test_key) @@ -161,8 +165,8 @@ async def test_revert_with_zero_reserved_balance_is_noop( assert test_key.reserved_balance == 0, ( f"Reserved balance should remain 0, got: {test_key.reserved_balance}" ) - assert test_key.total_requests == 0, ( - f"Total requests should remain 0, got: {test_key.total_requests}" + assert test_key.total_requests == 1, ( + f"Total requests should remain 1, got: {test_key.total_requests}" ) @@ -171,17 +175,18 @@ 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 + from routstr.auth import pay_for_request, 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, + reserved_balance=0, + total_requests=2, ) integration_session.add(test_key) await integration_session.commit() + await pay_for_request(test_key, 500, integration_session) result = await revert_pay_for_request(test_key, integration_session, 500) @@ -202,17 +207,21 @@ 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 + from routstr.auth import pay_for_request, 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, + reserved_balance=0, + total_requests=0, ) integration_session.add(test_key) await integration_session.commit() + await pay_for_request(test_key, 500, integration_session) + test_key.reserved_balance = 50 + 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) @@ -237,20 +246,28 @@ async def test_double_revert_prevented( 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 + from routstr.auth import ( + get_reservation_snapshot, + pay_for_request, + 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, + reserved_balance=0, + total_requests=4, ) integration_session.add(test_key) await integration_session.commit() + await pay_for_request(test_key, 500, integration_session) + snapshot = await get_reservation_snapshot(test_key, integration_session) # First revert — should succeed - result1 = await revert_pay_for_request(test_key, integration_session, 500) + result1 = await revert_pay_for_request( + test_key, integration_session, 500, snapshot + ) await integration_session.refresh(test_key) assert result1 is True @@ -258,7 +275,9 @@ async def test_double_revert_prevented( 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) + result2 = await revert_pay_for_request( + test_key, integration_session, 500, snapshot + ) await integration_session.refresh(test_key) assert result2 is False, "Second revert should be a no-op" @@ -279,22 +298,30 @@ async def test_sequential_reverts_never_go_negative( Simulates the double-revert scenario where multiple code paths attempt to revert the same reservation. """ - from routstr.auth import revert_pay_for_request + from routstr.auth import ( + get_reservation_snapshot, + pay_for_request, + 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, + reserved_balance=0, + total_requests=4, ) integration_session.add(test_key) await integration_session.commit() + await pay_for_request(test_key, 500, integration_session) + snapshot = await get_reservation_snapshot(test_key, integration_session) # 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) + r = await revert_pay_for_request( + test_key, integration_session, 500, snapshot + ) results.append(r) await integration_session.refresh(test_key) @@ -317,7 +344,11 @@ 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 + from routstr.auth import ( + get_reservation_snapshot, + pay_for_request, + revert_pay_for_request, + ) parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}" child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}" @@ -325,22 +356,26 @@ async def test_child_key_revert_floor_guard( parent_key = ApiKey( hashed_key=parent_key_hash, balance=10000, - reserved_balance=500, - total_requests=3, + reserved_balance=0, + total_requests=2, ) child_key = ApiKey( hashed_key=child_key_hash, balance=0, - reserved_balance=500, - total_requests=3, + reserved_balance=0, + total_requests=2, parent_key_hash=parent_key_hash, ) integration_session.add(parent_key) integration_session.add(child_key) await integration_session.commit() + await pay_for_request(child_key, 500, integration_session) + snapshot = await get_reservation_snapshot(child_key, integration_session) # First revert succeeds - result1 = await revert_pay_for_request(child_key, integration_session, 500) + result1 = await revert_pay_for_request( + child_key, integration_session, 500, snapshot + ) await integration_session.refresh(parent_key) await integration_session.refresh(child_key) @@ -349,7 +384,9 @@ async def test_child_key_revert_floor_guard( 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) + result2 = await revert_pay_for_request( + child_key, integration_session, 500, snapshot + ) await integration_session.refresh(parent_key) await integration_session.refresh(child_key) diff --git a/tests/unit/test_ehbp_finalize_payment.py b/tests/unit/test_ehbp_finalize_payment.py index d105edfe..08cb12ac 100644 --- a/tests/unit/test_ehbp_finalize_payment.py +++ b/tests/unit/test_ehbp_finalize_payment.py @@ -1,6 +1,7 @@ from __future__ import annotations -from typing import AsyncGenerator +from typing import Any, AsyncGenerator +from unittest.mock import AsyncMock, MagicMock import pytest from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine @@ -8,7 +9,8 @@ from sqlalchemy.pool import StaticPool from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.db import ApiKey +from routstr.auth import get_reservation_snapshot, pay_for_request +from routstr.core.db import ApiKey, ReservationRelease from routstr.upstream.ehbp import ( finalize_ehbp_actual_cost_payment, finalize_ehbp_max_cost_payment, @@ -43,18 +45,38 @@ async def _api_key(session: AsyncSession, hashed_key: str) -> ApiKey | None: ).one_or_none() +def _fail_nth_api_key_update( + session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, + target_update: int, +) -> None: + """Return rowcount=0 for one API-key UPDATE without mutating the database.""" + original_exec = session.exec + api_key_updates = 0 + + async def exec_with_failure( + statement: Any, *args: Any, **kwargs: Any + ) -> Any: + nonlocal api_key_updates + table = getattr(statement, "table", None) + if getattr(table, "name", None) == "api_keys": + api_key_updates += 1 + if api_key_updates == target_update: + return MagicMock(rowcount=0) + return await original_exec(statement, *args, **kwargs) + + monkeypatch.setattr(session, "exec", exec_with_failure) + + @pytest.mark.asyncio async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve( session: AsyncSession, ) -> None: - key = ApiKey( - hashed_key="ehbp-actual", - balance=10_000, - reserved_balance=3_000, - reserved_at=123, - ) + key = ApiKey(hashed_key="ehbp-actual", balance=10_000) session.add(key) await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) await finalize_ehbp_actual_cost_payment( key, @@ -68,6 +90,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve "input_msats": 500, "output_msats": 700, }, + reservation_snapshot=reservation, ) updated = await _api_key(session, "ehbp-actual") @@ -82,28 +105,22 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve async def test_finalize_max_cost_payment_updates_parent_and_child_spend( session: AsyncSession, ) -> None: - parent = ApiKey( - hashed_key="ehbp-parent", - balance=10_000, - reserved_balance=3_000, - reserved_at=123, - ) + parent = ApiKey(hashed_key="ehbp-parent", balance=10_000) child = ApiKey( - hashed_key="ehbp-child", - balance=0, - reserved_balance=3_000, - reserved_at=123, - parent_key_hash="ehbp-parent", + hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent" ) session.add(parent) session.add(child) await session.commit() + await pay_for_request(child, 3_000, session) + reservation = await get_reservation_snapshot(child, session) await finalize_ehbp_max_cost_payment( child, session, max_cost_for_model=3_000, model_id="tinfoil/model", + reservation_snapshot=reservation, ) updated_parent = await _api_key(session, "ehbp-parent") @@ -123,17 +140,16 @@ async def test_finalize_max_cost_payment_updates_parent_and_child_spend( @pytest.mark.asyncio async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows( session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, ) -> None: - key = ApiKey( - hashed_key="ehbp-missing-parent", - balance=10_000, - reserved_balance=3_000, - reserved_at=123, - ) + key = ApiKey(hashed_key="ehbp-missing-parent", balance=10_000) session.add(key) await session.commit() - await session.delete(key) - await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + _fail_nth_api_key_update(session, monkeypatch, target_update=1) + rollback_spy = AsyncMock(wraps=session.rollback) + monkeypatch.setattr(session, "rollback", rollback_spy) await finalize_ehbp_actual_cost_payment( key, @@ -141,45 +157,52 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche reserved_cost_for_model=3_000, model_id="tinfoil/model", cost_info={"total_msats": 1_200}, + reservation_snapshot=reservation, ) - assert await _api_key(session, "ehbp-missing-parent") is None + rollback_spy.assert_awaited_once() + updated = await _api_key(session, "ehbp-missing-parent") + assert updated is not None + assert updated.balance == 10_000 + assert updated.reserved_balance == 3_000 + assert updated.total_spent == 0 + release = await session.get(ReservationRelease, reservation.release_id) + assert release is not None + assert release.status == "active" @pytest.mark.asyncio async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows( session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, ) -> None: - parent = ApiKey( - hashed_key="ehbp-rollback-parent", - balance=10_000, - reserved_balance=3_000, - reserved_at=123, - ) + parent = ApiKey(hashed_key="ehbp-rollback-parent", balance=10_000) child = ApiKey( hashed_key="ehbp-missing-child", balance=0, - reserved_balance=3_000, - reserved_at=123, parent_key_hash="ehbp-rollback-parent", ) session.add(parent) session.add(child) await session.commit() - await session.delete(child) - await session.commit() + await pay_for_request(child, 3_000, session) + reservation = await get_reservation_snapshot(child, session) + _fail_nth_api_key_update(session, monkeypatch, target_update=2) await finalize_ehbp_max_cost_payment( child, session, max_cost_for_model=3_000, model_id="tinfoil/model", + reservation_snapshot=reservation, ) updated_parent = await _api_key(session, "ehbp-rollback-parent") assert updated_parent is not None assert updated_parent.balance == 10_000 assert updated_parent.reserved_balance == 3_000 - assert updated_parent.reserved_at == 123 assert updated_parent.total_spent == 0 - assert await _api_key(session, "ehbp-missing-child") is None + updated_child = await _api_key(session, "ehbp-missing-child") + assert updated_child is not None + assert updated_child.reserved_balance == 3_000 + assert updated_child.total_spent == 0 diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index ab98c270..c41568ad 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -18,6 +18,7 @@ from fastapi.responses import Response, StreamingResponse os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") os.environ.setdefault("UPSTREAM_API_KEY", "test") +from routstr.auth import ReservationSnapshot # noqa: E402 from routstr.core.db import ApiKey # noqa: E402 from routstr.payment.cost_calculation import CostData # noqa: E402 from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 @@ -498,6 +499,12 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: yield {"type": "message_stop"} fake_cost = {"total_msats": 4321, "total_usd": 0.00015} + reservation = ReservationSnapshot( + release_id="messages-stream", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=10_000, + ) captured_cost_call: dict[str, Any] = {} @@ -508,9 +515,11 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: max_cost: int, model_obj: Any = None, provider_fee: Any = None, + reservation_snapshot: Any = None, ) -> dict: captured_cost_call["combined_data"] = combined_data captured_cost_call["max_cost"] = max_cost + captured_cost_call["reservation_snapshot"] = reservation_snapshot return fake_cost fake_session = MagicMock() @@ -544,6 +553,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: session=session, max_cost_for_model=10_000, model_obj=model, + reservation_snapshot=reservation, ) assert isinstance(result, StreamingResponse) @@ -566,6 +576,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: assert combined["usage"]["input_tokens"] == 5 assert combined["usage"]["output_tokens"] == 7 assert combined["model"] == "openai/gpt-4o-mini" + assert captured_cost_call["reservation_snapshot"] is reservation @pytest.mark.asyncio @@ -593,6 +604,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: yield b'event: message_stop\ndata: {"type":"message_stop"}\n\n' fake_cost = {"total_msats": 999, "total_usd": 0.0001} + reservation = ReservationSnapshot( + release_id="messages-byte-stream", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=10_000, + ) captured: dict[str, Any] = {} async def fake_adjust( @@ -602,8 +619,10 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: max_cost: int, model_obj: Any = None, provider_fee: Any = None, + reservation_snapshot: Any = None, ) -> dict: captured["combined_data"] = combined_data + captured["reservation_snapshot"] = reservation_snapshot return fake_cost fake_session = MagicMock() @@ -636,6 +655,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: session=session, max_cost_for_model=10_000, model_obj=model, + reservation_snapshot=reservation, ) assert isinstance(result, StreamingResponse) @@ -659,6 +679,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: assert combined["usage"]["input_tokens"] == 3 assert combined["usage"]["output_tokens"] == 4 assert combined["model"] == "openai/gpt-4o-mini" + assert captured["reservation_snapshot"] is reservation # --------------------------------------------------------------------------- diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index b0e4a88c..99abc17d 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -16,13 +16,14 @@ 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 +from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession 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, ) @@ -150,6 +151,39 @@ async def test_release_stale_reservations_releases_old(session: AsyncSession) -> assert key.reserved_at is None +@pytest.mark.asyncio +async def test_targeted_parent_cleanup_releases_child_owned_reservation( + session: AsyncSession, +) -> None: + parent = ApiKey(hashed_key="stale-parent", balance=5_000) + child = ApiKey( + hashed_key="stale-child", parent_key_hash=parent.hashed_key, balance=0 + ) + session.add_all([parent, child]) + await session.commit() + await pay_for_request(child, 1_000, session) + reservation = ( + await session.exec( + select(ReservationRelease).where( + ReservationRelease.key_hash == child.hashed_key + ) + ) + ).one() + reservation.created_at = int(time.time()) - 1_000 + session.add(reservation) + await session.commit() + + released = await release_stale_reservations( + session, max_age_seconds=300, key_hash=parent.hashed_key + ) + + assert released == 1 + await session.refresh(parent) + await session.refresh(child) + assert parent.reserved_balance == 0 + assert child.reserved_balance == 0 + + @pytest.mark.asyncio async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> None: key = ApiKey( @@ -352,6 +386,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: upstream.forward_request = AsyncMock(side_effect=asyncio.CancelledError()) session = MagicMock() + reservation_snapshot = MagicMock() revert_mock = AsyncMock(return_value=True) with ( @@ -373,9 +408,16 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: proxy_module, "get_bearer_token_key", AsyncMock(return_value=key) ), patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), + patch.object( + proxy_module, + "get_reservation_snapshot", + AsyncMock(return_value=reservation_snapshot), + ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), ): with pytest.raises(asyncio.CancelledError): await proxy_module.proxy(request, "v1/chat/completions", session=session) - revert_mock.assert_awaited_once_with(key, session, 1_000) + revert_mock.assert_awaited_once_with( + key, session, 1_000, reservation_snapshot + ) diff --git a/tests/unit/test_stream_id_injection.py b/tests/unit/test_stream_id_injection.py index e19a9d3e..2d682bc5 100644 --- a/tests/unit/test_stream_id_injection.py +++ b/tests/unit/test_stream_id_injection.py @@ -4,6 +4,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from routstr.auth import ReservationSnapshot from routstr.core.db import ApiKey from routstr.upstream.base import BaseUpstreamProvider @@ -67,6 +68,12 @@ async def test_stream_with_id_injection() -> None: max_cost_for_model=100, background_tasks=background_tasks, requested_model="test-model", + reservation_snapshot=ReservationSnapshot( + release_id="test-release", + key_hash="test_hash", + billing_key_hash="test_hash", + reserved_msats=100, + ), ) results = [] diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py new file mode 100644 index 00000000..2ae574ab --- /dev/null +++ b/tests/unit/test_streaming_billing_finalization.py @@ -0,0 +1,448 @@ +import asyncio +from collections.abc import AsyncGenerator +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +import routstr.auth as auth_module +from routstr.auth import ( + ReservationSnapshot, + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + release_reservation, +) +from routstr.core.db import ApiKey, ReservationRelease +from routstr.payment.cost_calculation import MaxCostData +from routstr.upstream.base import BaseUpstreamProvider + + +async def _engine() -> AsyncEngine: + engine = create_async_engine("sqlite+aiosqlite://") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + return engine + + +@pytest.mark.asyncio +async def test_release_reservation_is_durable_and_idempotent() -> None: + engine = await _engine() + key = ApiKey(hashed_key="key", balance=1_000) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add(key) + await session.commit() + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + + record = await session.get(ReservationRelease, snapshot.release_id) + assert record is not None and record.status == "active" + assert await release_reservation(snapshot, session, 500) is True + assert await release_reservation(snapshot, session, 500) is True + + await session.refresh(key) + await session.refresh(record) + assert key.reserved_balance == 0 + assert key.reserved_at is None + assert record.status == "released" + await engine.dispose() + + +@pytest.mark.asyncio +async def test_release_only_owns_its_concurrent_reservation() -> None: + engine = await _engine() + key = ApiKey(hashed_key="key", balance=1_000) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add(key) + await session.commit() + + await pay_for_request(key, 400, session) + first = await get_reservation_snapshot(key, session) + await pay_for_request(key, 400, session) + second = await get_reservation_snapshot(key, session) + + assert first.release_id != second.release_id + assert await release_reservation(first, session, 400) is True + assert await release_reservation(first, session, 400) is True + await session.refresh(key) + assert key.reserved_balance == 400 + + assert await release_reservation(second, session, 400) is True + await session.refresh(key) + assert key.reserved_balance == 0 + await engine.dispose() + + +@pytest.mark.asyncio +async def test_release_updates_parent_and_child_atomically() -> None: + engine = await _engine() + parent = ApiKey(hashed_key="parent", balance=1_000) + child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add_all([parent, child]) + await session.commit() + await pay_for_request(child, 500, session) + snapshot = await get_reservation_snapshot(child, session) + + assert await release_reservation(snapshot, session, 500) is True + await session.refresh(parent) + await session.refresh(child) + assert (parent.reserved_balance, child.reserved_balance) == (0, 0) + assert (parent.reserved_at, child.reserved_at) == (None, None) + await engine.dispose() + + +@pytest.mark.asyncio +async def test_release_rolls_back_partial_parent_child_update() -> None: + engine = await _engine() + parent = ApiKey(hashed_key="parent", balance=1_000) + child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add_all([parent, child]) + await session.commit() + await pay_for_request(child, 500, session) + snapshot = await get_reservation_snapshot(child, session) + child.reserved_balance = 100 + session.add(child) + await session.commit() + + assert await release_reservation(snapshot, session, 500) is False + await session.refresh(parent) + await session.refresh(child) + record = await session.get(ReservationRelease, snapshot.release_id) + assert (parent.reserved_balance, child.reserved_balance) == (500, 100) + assert record is not None and record.status == "active" + await engine.dispose() + + +@pytest.mark.asyncio +async def test_post_commit_failure_cannot_release_charged_reservation() -> None: + engine = await _engine() + key = ApiKey(hashed_key="key", balance=1_000) + cost = MaxCostData( + base_msats=500, + input_msats=0, + output_msats=0, + total_msats=500, + ) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add(key) + await session.commit() + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + + with ( + patch("routstr.auth.calculate_cost", AsyncMock(return_value=cost)), + patch.object( + session, + "refresh", + AsyncMock(side_effect=SQLAlchemyError("post-commit refresh failed")), + ), + ): + with pytest.raises(SQLAlchemyError, match="post-commit refresh failed"): + await adjust_payment_for_tokens(key, {}, session, 500) + + await session.rollback() + assert await release_reservation(snapshot, session, 500) is False + charged_key = await session.get(ApiKey, "key") + record = await session.get(ReservationRelease, snapshot.release_id) + assert charged_key is not None + assert (charged_key.balance, charged_key.reserved_balance) == (500, 0) + assert record is not None and record.status == "charged" + await engine.dispose() + + +@pytest.mark.asyncio +async def test_generic_background_settlement_uses_explicit_reservation() -> None: + engine = await _engine() + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key", provider_fee=1.0 + ) + key = ApiKey(hashed_key="generic-key", balance=1_000) + cost = MaxCostData( + base_msats=500, + input_msats=0, + output_msats=0, + total_msats=500, + ) + + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add(key) + await session.commit() + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + + context_token = auth_module._current_reservation.set(None) + try: + with ( + patch( + "routstr.upstream.base.create_session", + side_effect=lambda: AsyncSession(engine, expire_on_commit=False), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + patch("routstr.auth.calculate_cost", AsyncMock(return_value=cost)), + ): + await provider._finalize_generic_streaming_payment( + key.hashed_key, + 500, + "audio/speech", + model_obj=None, + provider_fee=provider.provider_fee, + reservation_snapshot=snapshot, + ) + finally: + auth_module._current_reservation.reset(context_token) + + async with AsyncSession(engine, expire_on_commit=False) as session: + settled_key = await session.get(ApiKey, key.hashed_key) + record = await session.get(ReservationRelease, snapshot.release_id) + assert settled_key is not None + assert (settled_key.balance, settled_key.reserved_balance) == (500, 0) + assert record is not None and record.status == "charged" + await engine.dispose() + + +@pytest.mark.asyncio +async def test_streaming_release_is_terminal_and_suppresses_background_charge() -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield b"data: [DONE]\n\n" + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {"content-type": "text/event-stream"} + upstream_response.aiter_bytes = aiter_bytes + + key = MagicMock(spec=ApiKey) + key.hashed_key = "test-key-hash" + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.rollback = AsyncMock() + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + release = AsyncMock(return_value=True) + reservation_snapshot = MagicMock() + reservation_snapshot.reserved_msats = 500 + background_tasks = MagicMock() + + with ( + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + AsyncMock(side_effect=SQLAlchemyError("database unavailable")), + ), + patch( + "routstr.upstream.base.get_reservation_snapshot", + AsyncMock(return_value=reservation_snapshot), + ), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await provider.handle_streaming_chat_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + background_tasks=background_tasks, + ) + + with pytest.raises(SQLAlchemyError, match="database unavailable"): + async for _ in response.body_iterator: + pass + + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(reservation_snapshot, session, 500) + background_tasks.add_task.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "release_outcome", + [True, False, RuntimeError("release failed"), asyncio.CancelledError()], +) +async def test_responses_streaming_releases_and_raises_on_billing_failure( + release_outcome: bool | BaseException, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield ( + b'data: {"type":"response.completed","response":{"model":"test",' + b'"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + ) + yield b"data: [DONE]\n\n" + + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = aiter_bytes + key = MagicMock(spec=ApiKey) + key.hashed_key = "responses-key" + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.rollback = AsyncMock() + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + snapshot = ReservationSnapshot( + release_id="responses-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = ( + AsyncMock(side_effect=release_outcome) + if isinstance(release_outcome, BaseException) + else AsyncMock(return_value=release_outcome) + ) + adjust = AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await provider.handle_streaming_responses_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + ) + emitted = bytearray() + with pytest.raises(SQLAlchemyError, match="database unavailable"): + async for chunk in response.body_iterator: + if isinstance(chunk, str): + emitted.extend(chunk.encode()) + else: + emitted.extend(bytes(chunk)) + + assert b'"total_msats": 0' not in emitted + adjust.assert_awaited_once() + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("via_litellm", [False, True]) +@pytest.mark.parametrize( + "release_outcome", + [True, False, RuntimeError("release failed"), asyncio.CancelledError()], +) +async def test_messages_streaming_releases_and_raises_on_billing_failure( + via_litellm: bool, + release_outcome: bool | BaseException, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + key = MagicMock(spec=ApiKey) + key.hashed_key = "messages-key" + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.rollback = AsyncMock() + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + snapshot = ReservationSnapshot( + release_id=f"messages-{'litellm' if via_litellm else 'native'}", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = ( + AsyncMock(side_effect=release_outcome) + if isinstance(release_outcome, BaseException) + else AsyncMock(return_value=release_outcome) + ) + adjust = AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + + async def native_chunks() -> AsyncGenerator[bytes, None]: + yield ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"model":"test","usage":{"input_tokens":1,"output_tokens":0}}}\n\n' + ) + yield b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + + async def litellm_chunks() -> AsyncGenerator[dict, None]: + yield { + "type": "message_start", + "message": { + "model": "test", + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + } + yield {"type": "message_stop"} + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + if via_litellm: + response = provider._stream_litellm_messages( + iterator=litellm_chunks(), + key=key, + max_cost_for_model=500, + requested_model=None, + reservation_snapshot=snapshot, + ) + else: + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = native_chunks + response = await provider.handle_streaming_messages_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + ) + + with pytest.raises(SQLAlchemyError, match="database unavailable"): + async for _ in response.body_iterator: + pass + + adjust.assert_awaited_once() + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + + +@pytest.mark.asyncio +async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None: + engine = await _engine() + first = ApiKey(hashed_key="first", balance=1_000) + second = ApiKey(hashed_key="second", balance=1_000) + async with AsyncSession(engine, expire_on_commit=False) as session: + session.add(first) + session.add(second) + await session.commit() + await pay_for_request(first, 500, session) + snapshot = await get_reservation_snapshot(first, session) + + with pytest.raises(RuntimeError, match="does not belong"): + await adjust_payment_for_tokens( + second, + {"model": "test", "usage": None}, + session, + 500, + reservation_snapshot=snapshot, + ) + + await session.refresh(first) + await session.refresh(second) + assert first.reserved_balance == 500 + assert second.reserved_balance == 0 + + await engine.dispose() diff --git a/tests/unit/test_streaming_sse_providers.py b/tests/unit/test_streaming_sse_providers.py index 18deb599..ffb7e266 100644 --- a/tests/unit/test_streaming_sse_providers.py +++ b/tests/unit/test_streaming_sse_providers.py @@ -24,6 +24,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from routstr.auth import ReservationSnapshot from routstr.core.db import ApiKey from routstr.upstream import base from routstr.upstream.base import BaseUpstreamProvider @@ -67,6 +68,12 @@ async def _drive(chunks: list[bytes], requested_model: str | None = None) -> lis max_cost_for_model=100, background_tasks=MagicMock(), requested_model=requested_model, + reservation_snapshot=ReservationSnapshot( + release_id="test-release", + key_hash="test_hash", + billing_key_hash="test_hash", + reserved_msats=100, + ), ) out: list[bytes] = [] diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index 4290c32b..216e1c3a 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -331,6 +331,7 @@ async def test_5xx_wrapped_rate_limit_is_classified( @pytest.mark.asyncio async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: from routstr import proxy as proxy_module + from routstr.auth import ReservationSnapshot from routstr.core.db import ApiKey from routstr.core.exceptions import UpstreamError @@ -359,6 +360,12 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: ) session = MagicMock() + reservation = ReservationSnapshot( + release_id="rate-limit-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=1_000, + ) revert_mock = AsyncMock(return_value=True) with ( @@ -380,6 +387,11 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: proxy_module, "get_bearer_token_key", AsyncMock(return_value=key) ), patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), + patch.object( + proxy_module, + "get_reservation_snapshot", + AsyncMock(return_value=reservation), + ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), ): response = await proxy_module.proxy( @@ -396,4 +408,4 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: assert RAW_ORG_ID not in serialized assert "org-[REDACTED]" in serialized # Single upstream failed -> reservation reverted exactly once (no double-charge). - revert_mock.assert_awaited_once_with(key, session, 1_000) + revert_mock.assert_awaited_once_with(key, session, 1_000, reservation)