Merge pull request #623 from Routstr/fix/streaming-billing-finalization

fix: fail safely on streaming billing errors
This commit is contained in:
9qeklajc
2026-07-24 20:19:56 +02:00
committed by GitHub
18 changed files with 1744 additions and 409 deletions
@@ -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")
+369 -153
View File
@@ -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:
+10 -20
View File
@@ -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={
+128 -16
View File
@@ -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)
+30 -7
View File
@@ -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
+254 -115
View File
@@ -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(
+28 -2
View File
@@ -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,
@@ -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)
+127 -7
View File
@@ -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)
@@ -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
@@ -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)
+63 -40
View File
@@ -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
@@ -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
# ---------------------------------------------------------------------------
+44 -2
View File
@@ -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
)
+7
View File
@@ -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 = []
@@ -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()
@@ -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] = []
+13 -1
View File
@@ -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)