mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
stablize reserved fee calculation
This commit is contained in:
@@ -438,13 +438,18 @@ Routstr returns cost info as response headers:
|
|||||||
|
|
||||||
| Header | Auth | Description |
|
| Header | Auth | Description |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Total msats charged for this request |
|
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Settled msats debited for this request |
|
||||||
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the charge |
|
| `X-Routstr-Computed-Cost-Msats` | Bearer, X-Cashu | Computed usage cost, emitted when it differs from the settled debit |
|
||||||
|
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the computed usage |
|
||||||
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
|
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
|
||||||
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
|
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
|
||||||
|
|
||||||
The client/Tinfoil SDK can read these headers from the HTTP response without
|
The client/Tinfoil SDK can read these headers from the HTTP response without
|
||||||
needing to decrypt the body.
|
needing to decrypt the body. A duplicate or rejected finalization can therefore
|
||||||
|
report a zero settled debit while preserving the non-zero computed usage cost.
|
||||||
|
Normal JSON responses use the same contract: `total_msats` and `charged_msats`
|
||||||
|
are settled values, while `computed_msats` retains the usage calculation when
|
||||||
|
it differs.
|
||||||
|
|
||||||
### Setup
|
### Setup
|
||||||
|
|
||||||
|
|||||||
+320
-200
@@ -4,6 +4,7 @@ import math
|
|||||||
import random
|
import random
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
|
from contextlib import suppress
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
@@ -867,6 +868,10 @@ async def pay_for_request(
|
|||||||
_clear_current_reservation(reservation)
|
_clear_current_reservation(reservation)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
# The reservation is durable; keep its lease fresh for the whole request
|
||||||
|
# lifetime (upstream header waits, non-streaming and streaming alike).
|
||||||
|
_start_reservation_heartbeat(reservation)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
@@ -957,6 +962,84 @@ async def _validate_reservation_snapshot(
|
|||||||
raise RuntimeError("Billing reservation record does not match the request")
|
raise RuntimeError("Billing reservation record does not match the request")
|
||||||
|
|
||||||
|
|
||||||
|
async def renew_reservation(
|
||||||
|
snapshot: ReservationSnapshot, session: AsyncSession
|
||||||
|
) -> bool:
|
||||||
|
"""Push an active reservation's lease forward so the sweeper skips it.
|
||||||
|
|
||||||
|
``ReservationRelease.created_at`` doubles as the lease timestamp: the
|
||||||
|
stale-reservation sweeper releases reservations whose ``created_at`` is
|
||||||
|
older than the timeout, so a long-lived stream must renew it periodically
|
||||||
|
or lose its reservation mid-flight (and finish uncharged, since release is
|
||||||
|
terminal). Returns False once the reservation reached a terminal state.
|
||||||
|
"""
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.id) == snapshot.release_id)
|
||||||
|
.where(col(ReservationRelease.status) == "active")
|
||||||
|
.values(created_at=int(time.time()))
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return bool(result.rowcount == 1)
|
||||||
|
|
||||||
|
|
||||||
|
# One heartbeat task per in-flight reservation, keyed by release id. Started
|
||||||
|
# when the reservation is created and stopped when it reaches a terminal
|
||||||
|
# state, so every request path — header waits, non-streaming, streaming — is
|
||||||
|
# covered for its whole lifetime.
|
||||||
|
_reservation_heartbeats: dict[str, "asyncio.Task[None]"] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None:
|
||||||
|
"""Keep an in-flight reservation's lease fresh until it is finalized.
|
||||||
|
|
||||||
|
Spawns a background task that renews the lease every third of the stale
|
||||||
|
timeout using its own session, so requests longer than
|
||||||
|
``STALE_RESERVATION_TIMEOUT_SECONDS`` are not swept and finish charged.
|
||||||
|
The task stops on its own once the reservation reaches a terminal state
|
||||||
|
or its owning request task finishes; terminal transitions also stop it
|
||||||
|
explicitly. Binding renewal to the owner's lifetime guarantees the sweeper
|
||||||
|
can always recover a reservation whose request died without finalizing —
|
||||||
|
a detached heartbeat would otherwise renew it forever and lock the funds.
|
||||||
|
"""
|
||||||
|
interval = max(1, settings.stale_reservation_timeout_seconds // 3)
|
||||||
|
owner = asyncio.current_task()
|
||||||
|
|
||||||
|
async def beat() -> None:
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(interval)
|
||||||
|
if owner is None or owner.done():
|
||||||
|
# Request control is gone; let the lease expire so the
|
||||||
|
# sweeper can release the reservation if no terminal
|
||||||
|
# transition ever ran.
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
if not await renew_reservation(snapshot, session):
|
||||||
|
return
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Failed to renew billing reservation lease",
|
||||||
|
extra={"release_id": snapshot.release_id},
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
_reservation_heartbeats.pop(snapshot.release_id, None)
|
||||||
|
|
||||||
|
_reservation_heartbeats[snapshot.release_id] = asyncio.create_task(beat())
|
||||||
|
|
||||||
|
|
||||||
|
async def _stop_reservation_heartbeat(release_id: str) -> None:
|
||||||
|
"""Cancel and await a reservation's heartbeat so no renewal overlaps
|
||||||
|
finalization."""
|
||||||
|
task = _reservation_heartbeats.pop(release_id, None)
|
||||||
|
if task is None:
|
||||||
|
return
|
||||||
|
task.cancel()
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
|
||||||
async def get_reservation_snapshot(
|
async def get_reservation_snapshot(
|
||||||
key: ApiKey, session: AsyncSession
|
key: ApiKey, session: AsyncSession
|
||||||
) -> ReservationSnapshot:
|
) -> ReservationSnapshot:
|
||||||
@@ -968,6 +1051,60 @@ async def get_reservation_snapshot(
|
|||||||
return snapshot
|
return snapshot
|
||||||
|
|
||||||
|
|
||||||
|
async def _repair_corrupt_reservation(
|
||||||
|
snapshot: ReservationSnapshot,
|
||||||
|
session: AsyncSession,
|
||||||
|
*,
|
||||||
|
decrement_requests: bool,
|
||||||
|
) -> bool:
|
||||||
|
"""Terminalize a reservation without subtracting uncertain aggregates."""
|
||||||
|
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")
|
||||||
|
)
|
||||||
|
result = await session.exec(transition) # type: ignore[call-overload]
|
||||||
|
if result.rowcount != 1:
|
||||||
|
await session.rollback()
|
||||||
|
return False
|
||||||
|
|
||||||
|
if decrement_requests:
|
||||||
|
for key_hash in {snapshot.billing_key_hash, snapshot.key_hash}:
|
||||||
|
request_result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(
|
||||||
|
total_requests=case(
|
||||||
|
(
|
||||||
|
col(ApiKey.total_requests) > 0,
|
||||||
|
col(ApiKey.total_requests) - 1,
|
||||||
|
),
|
||||||
|
else_=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if request_result.rowcount != 1:
|
||||||
|
await session.rollback()
|
||||||
|
return False
|
||||||
|
|
||||||
|
await session.commit()
|
||||||
|
logger.error(
|
||||||
|
"Released corrupt reservation without aggregate subtraction",
|
||||||
|
extra={
|
||||||
|
"reservation_id": snapshot.release_id,
|
||||||
|
"billing_key_hash": snapshot.billing_key_hash[:8] + "...",
|
||||||
|
"reserved_msats": snapshot.reserved_msats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await _stop_reservation_heartbeat(snapshot.release_id)
|
||||||
|
_clear_current_reservation(snapshot)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
async def _transition_reservation_to_released(
|
async def _transition_reservation_to_released(
|
||||||
snapshot: ReservationSnapshot,
|
snapshot: ReservationSnapshot,
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
@@ -988,7 +1125,7 @@ async def _transition_reservation_to_released(
|
|||||||
if transition_result.rowcount != 1:
|
if transition_result.rowcount != 1:
|
||||||
await session.rollback()
|
await session.rollback()
|
||||||
existing = await session.get(ReservationRelease, snapshot.release_id)
|
existing = await session.get(ReservationRelease, snapshot.release_id)
|
||||||
return bool(
|
already_released = bool(
|
||||||
idempotent_success
|
idempotent_success
|
||||||
and existing is not None
|
and existing is not None
|
||||||
and existing.status == "released"
|
and existing.status == "released"
|
||||||
@@ -996,6 +1133,9 @@ async def _transition_reservation_to_released(
|
|||||||
and existing.billing_key_hash == snapshot.billing_key_hash
|
and existing.billing_key_hash == snapshot.billing_key_hash
|
||||||
and existing.reserved_msats == snapshot.reserved_msats
|
and existing.reserved_msats == snapshot.reserved_msats
|
||||||
)
|
)
|
||||||
|
if already_released:
|
||||||
|
await _stop_reservation_heartbeat(snapshot.release_id)
|
||||||
|
return already_released
|
||||||
|
|
||||||
values: dict[str, object] = {
|
values: dict[str, object] = {
|
||||||
"reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats,
|
"reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats,
|
||||||
@@ -1019,7 +1159,9 @@ async def _transition_reservation_to_released(
|
|||||||
result = await session.exec(release_stmt) # type: ignore[call-overload]
|
result = await session.exec(release_stmt) # type: ignore[call-overload]
|
||||||
if result.rowcount != 1:
|
if result.rowcount != 1:
|
||||||
await session.rollback()
|
await session.rollback()
|
||||||
return False
|
return await _repair_corrupt_reservation(
|
||||||
|
snapshot, session, decrement_requests=decrement_requests
|
||||||
|
)
|
||||||
|
|
||||||
if snapshot.billing_key_hash != snapshot.key_hash:
|
if snapshot.billing_key_hash != snapshot.key_hash:
|
||||||
child_release_stmt = (
|
child_release_stmt = (
|
||||||
@@ -1033,9 +1175,12 @@ async def _transition_reservation_to_released(
|
|||||||
)
|
)
|
||||||
if child_result.rowcount != 1:
|
if child_result.rowcount != 1:
|
||||||
await session.rollback()
|
await session.rollback()
|
||||||
return False
|
return await _repair_corrupt_reservation(
|
||||||
|
snapshot, session, decrement_requests=decrement_requests
|
||||||
|
)
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(snapshot.release_id)
|
||||||
_clear_current_reservation(snapshot)
|
_clear_current_reservation(snapshot)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -1071,6 +1216,9 @@ async def _claim_reservation_for_charge(
|
|||||||
)
|
)
|
||||||
result = await session.exec(statement) # type: ignore[call-overload]
|
result = await session.exec(statement) # type: ignore[call-overload]
|
||||||
if result.rowcount == 1:
|
if result.rowcount == 1:
|
||||||
|
# The claim is not committed yet — the heartbeat must keep running
|
||||||
|
# until the surrounding charge transaction commits, or a rollback
|
||||||
|
# would restore an active reservation with no lease renewal.
|
||||||
_clear_current_reservation(snapshot)
|
_clear_current_reservation(snapshot)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -1078,6 +1226,70 @@ async def _claim_reservation_for_charge(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def _charge_reservation_rows(
|
||||||
|
session: AsyncSession,
|
||||||
|
*,
|
||||||
|
billing_key_hash: str,
|
||||||
|
key_hash: str,
|
||||||
|
reserved_msats: int,
|
||||||
|
charge_msats: int,
|
||||||
|
extra_billing_guards: tuple = (),
|
||||||
|
) -> bool:
|
||||||
|
"""Release the reserved amount and record the charge on the billing row
|
||||||
|
(and the child row when different) inside the caller's transaction.
|
||||||
|
|
||||||
|
Guarded subtraction replaces defensive clamping: every row must still hold
|
||||||
|
the full reserved amount, otherwise the whole transaction rolls back and
|
||||||
|
nothing is charged. A violated invariant must never silently erase the
|
||||||
|
aggregate reservations of sibling requests. Returns False after rollback.
|
||||||
|
"""
|
||||||
|
billing_stmt = (
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == billing_key_hash)
|
||||||
|
.where(col(ApiKey.reserved_balance) >= reserved_msats)
|
||||||
|
.values(
|
||||||
|
reserved_balance=col(ApiKey.reserved_balance) - reserved_msats,
|
||||||
|
reserved_at=case(
|
||||||
|
(
|
||||||
|
col(ApiKey.reserved_balance) - reserved_msats > 0,
|
||||||
|
col(ApiKey.reserved_at),
|
||||||
|
),
|
||||||
|
else_=None,
|
||||||
|
),
|
||||||
|
balance=col(ApiKey.balance) - charge_msats,
|
||||||
|
total_spent=col(ApiKey.total_spent) + charge_msats,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for guard in extra_billing_guards:
|
||||||
|
billing_stmt = billing_stmt.where(guard)
|
||||||
|
result = await session.exec(billing_stmt) # type: ignore[call-overload]
|
||||||
|
if result.rowcount != 1:
|
||||||
|
await session.rollback()
|
||||||
|
return False
|
||||||
|
|
||||||
|
if key_hash != billing_key_hash:
|
||||||
|
child_result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.where(col(ApiKey.reserved_balance) >= reserved_msats)
|
||||||
|
.values(
|
||||||
|
reserved_balance=col(ApiKey.reserved_balance) - reserved_msats,
|
||||||
|
reserved_at=case(
|
||||||
|
(
|
||||||
|
col(ApiKey.reserved_balance) - reserved_msats > 0,
|
||||||
|
col(ApiKey.reserved_at),
|
||||||
|
),
|
||||||
|
else_=None,
|
||||||
|
),
|
||||||
|
total_spent=col(ApiKey.total_spent) + charge_msats,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if child_result.rowcount != 1:
|
||||||
|
await session.rollback()
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
async def adjust_payment_for_tokens(
|
async def adjust_payment_for_tokens(
|
||||||
key: ApiKey,
|
key: ApiKey,
|
||||||
response_data: dict,
|
response_data: dict,
|
||||||
@@ -1108,6 +1320,10 @@ async def adjust_payment_for_tokens(
|
|||||||
# changed the caller's original estimate.
|
# changed the caller's original estimate.
|
||||||
deducted_max_cost = reservation.reserved_msats
|
deducted_max_cost = reservation.reserved_msats
|
||||||
model = response_data.get("model", "unknown")
|
model = response_data.get("model", "unknown")
|
||||||
|
# Failure paths log after a rollback has expired the ORM instances, so
|
||||||
|
# capture the identifiers as plain strings up front.
|
||||||
|
key_log_hash = key.hashed_key[:8] + "..."
|
||||||
|
billing_log_hash = billing_key.hashed_key[:8] + "..."
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Starting payment adjustment for tokens",
|
"Starting payment adjustment for tokens",
|
||||||
@@ -1132,8 +1348,8 @@ async def adjust_payment_for_tokens(
|
|||||||
if released
|
if released
|
||||||
else "Reservation was already finalized; fallback skipped",
|
else "Reservation was already finalized; fallback skipped",
|
||||||
extra={
|
extra={
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": billing_log_hash,
|
||||||
"deducted_max_cost": deducted_max_cost,
|
"deducted_max_cost": deducted_max_cost,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -1142,8 +1358,8 @@ async def adjust_payment_for_tokens(
|
|||||||
"Failed to release reservation in fallback",
|
"Failed to release reservation in fallback",
|
||||||
extra={
|
extra={
|
||||||
"error": str(e),
|
"error": str(e),
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": billing_log_hash,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1166,6 +1382,7 @@ async def adjust_payment_for_tokens(
|
|||||||
# A prior charge or release already owns this reservation. Returning
|
# A prior charge or release already owns this reservation. Returning
|
||||||
# the calculated metadata is safe; the aggregate balances must not
|
# the calculated metadata is safe; the aggregate balances must not
|
||||||
# be modified a second time.
|
# be modified a second time.
|
||||||
|
calculated_cost.charged_msats = 0
|
||||||
return calculated_cost.dict()
|
return calculated_cost.dict()
|
||||||
|
|
||||||
match calculated_cost:
|
match calculated_cost:
|
||||||
@@ -1179,75 +1396,32 @@ async def adjust_payment_for_tokens(
|
|||||||
"max_cost": cost.total_msats,
|
"max_cost": cost.total_msats,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
# Finalize by releasing reservation and charging max cost
|
# Finalize by releasing the reservation and charging max cost.
|
||||||
if billing_key.reserved_balance < deducted_max_cost:
|
charged = await _charge_reservation_rows(
|
||||||
logger.error(
|
session,
|
||||||
"reserved_balance below deducted_max_cost before MaxCost finalization — clamping to 0",
|
billing_key_hash=billing_key.hashed_key,
|
||||||
extra={
|
key_hash=key.hashed_key,
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
reserved_msats=deducted_max_cost,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
charge_msats=cost.total_msats,
|
||||||
"reserved_balance": billing_key.reserved_balance,
|
|
||||||
"deducted_max_cost": deducted_max_cost,
|
|
||||||
"total_cost_msats": cost.total_msats,
|
|
||||||
"balance": billing_key.balance,
|
|
||||||
"total_spent": billing_key.total_spent,
|
|
||||||
"model": model,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
safe_reserved = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
|
||||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
)
|
||||||
|
if charged:
|
||||||
finalize_stmt = (
|
await session.commit()
|
||||||
update(ApiKey)
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
if not charged:
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
balance=col(ApiKey.balance) - cost.total_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + cost.total_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(finalize_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
# 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,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
total_spent=col(ApiKey.total_spent) + cost.total_msats,
|
|
||||||
reserved_balance=child_safe_reserved,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
await session.commit()
|
|
||||||
if result.rowcount == 0:
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to finalize max-cost payment - retrying reservation release",
|
"Failed to finalize max-cost payment - retrying reservation release",
|
||||||
extra={
|
extra={
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": billing_log_hash,
|
||||||
"deducted_max_cost": deducted_max_cost,
|
"deducted_max_cost": deducted_max_cost,
|
||||||
"current_reserved_balance": billing_key.reserved_balance,
|
|
||||||
"total_cost": cost.total_msats,
|
"total_cost": cost.total_msats,
|
||||||
"model": model,
|
"model": model,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
cost.charged_msats = 0
|
||||||
await release_reservation_only()
|
await release_reservation_only()
|
||||||
else:
|
else:
|
||||||
|
cost.charged_msats = cost.total_msats
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -1314,60 +1488,30 @@ async def adjust_payment_for_tokens(
|
|||||||
"model": model,
|
"model": model,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
if billing_key.reserved_balance < deducted_max_cost:
|
if not await _charge_reservation_rows(
|
||||||
|
session,
|
||||||
|
billing_key_hash=billing_key.hashed_key,
|
||||||
|
key_hash=key.hashed_key,
|
||||||
|
reserved_msats=deducted_max_cost,
|
||||||
|
charge_msats=total_cost_msats,
|
||||||
|
):
|
||||||
logger.error(
|
logger.error(
|
||||||
"reserved_balance below deducted_max_cost on exact-cost finalization — clamping to 0",
|
"Failed to finalize exact-cost payment - releasing reservation",
|
||||||
extra={
|
extra={
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": billing_log_hash,
|
||||||
"reserved_balance": billing_key.reserved_balance,
|
|
||||||
"deducted_max_cost": deducted_max_cost,
|
"deducted_max_cost": deducted_max_cost,
|
||||||
"total_cost_msats": total_cost_msats,
|
"total_cost": total_cost_msats,
|
||||||
"balance": billing_key.balance,
|
|
||||||
"total_spent": billing_key.total_spent,
|
|
||||||
"model": model,
|
"model": model,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
cost.charged_msats = 0
|
||||||
exact_safe_reserved = case(
|
await release_reservation_only()
|
||||||
(
|
return cost.dict()
|
||||||
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=exact_safe_reserved,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(finalize_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
# 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,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
reserved_balance=child_exact_safe_reserved,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
cost.charged_msats = total_cost_msats
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -1406,50 +1550,70 @@ async def adjust_payment_for_tokens(
|
|||||||
)
|
)
|
||||||
).one()
|
).one()
|
||||||
observed_balance = locked_billing_key.balance
|
observed_balance = locked_billing_key.balance
|
||||||
actual_charge_msats = min(observed_balance, total_cost_msats)
|
observed_reserved = locked_billing_key.reserved_balance
|
||||||
overrun_safe_reserved = case(
|
# An overrun may only spend this request's own reservation
|
||||||
(
|
# plus funds no other in-flight request has reserved.
|
||||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
# Charging against the raw balance would consume sibling
|
||||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
# reservations and drive the available balance negative.
|
||||||
),
|
if observed_reserved < deducted_max_cost:
|
||||||
else_=0,
|
# Invariant violated — never clamp and charge anyway,
|
||||||
)
|
# that would erase sibling reservations. Release only.
|
||||||
finalize_result = await session.exec( # type: ignore[call-overload]
|
logger.error(
|
||||||
update(ApiKey)
|
"reserved_balance below reservation on overrun finalization — releasing without charge",
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
extra={
|
||||||
.where(col(ApiKey.balance) == observed_balance)
|
"key_hash": key_log_hash,
|
||||||
.values(
|
"billing_key_hash": billing_log_hash,
|
||||||
reserved_balance=overrun_safe_reserved,
|
"reserved_balance": observed_reserved,
|
||||||
balance=col(ApiKey.balance) - actual_charge_msats,
|
"deducted_max_cost": deducted_max_cost,
|
||||||
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
|
"total_cost_msats": total_cost_msats,
|
||||||
|
"model": model,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
)
|
await session.rollback()
|
||||||
if finalize_result.rowcount == 1:
|
cost.charged_msats = 0
|
||||||
|
await release_reservation_only()
|
||||||
|
return cost.dict()
|
||||||
|
sibling_reserved = observed_reserved - deducted_max_cost
|
||||||
|
chargeable_msats = max(0, observed_balance - sibling_reserved)
|
||||||
|
actual_charge_msats = min(chargeable_msats, total_cost_msats)
|
||||||
|
if await _charge_reservation_rows(
|
||||||
|
session,
|
||||||
|
billing_key_hash=billing_key.hashed_key,
|
||||||
|
key_hash=key.hashed_key,
|
||||||
|
reserved_msats=deducted_max_cost,
|
||||||
|
charge_msats=actual_charge_msats,
|
||||||
|
extra_billing_guards=(
|
||||||
|
col(ApiKey.balance) == observed_balance,
|
||||||
|
col(ApiKey.reserved_balance) == observed_reserved,
|
||||||
|
),
|
||||||
|
):
|
||||||
break
|
break
|
||||||
await session.rollback()
|
|
||||||
if not await _claim_reservation_for_charge(reservation, session):
|
if not await _claim_reservation_for_charge(reservation, session):
|
||||||
|
cost.charged_msats = 0
|
||||||
return cost.dict()
|
return cost.dict()
|
||||||
else:
|
else:
|
||||||
await session.rollback()
|
await session.rollback()
|
||||||
raise RuntimeError("Could not atomically finalize cost overrun")
|
raise RuntimeError("Could not atomically finalize cost overrun")
|
||||||
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=overrun_safe_reserved,
|
|
||||||
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
cost.total_msats = actual_charge_msats
|
cost.charged_msats = actual_charge_msats
|
||||||
|
if actual_charge_msats < total_cost_msats:
|
||||||
|
logger.warning(
|
||||||
|
"Cost overrun exceeded chargeable funds — shortfall written off",
|
||||||
|
extra={
|
||||||
|
"key_hash": key.hashed_key[:8] + "...",
|
||||||
|
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||||
|
"actual_cost_msats": total_cost_msats,
|
||||||
|
"charged_msats": actual_charge_msats,
|
||||||
|
"shortfall_msats": total_cost_msats - actual_charge_msats,
|
||||||
|
"model": model,
|
||||||
|
},
|
||||||
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Finalized payment with additional charge",
|
"Finalized payment with additional charge",
|
||||||
extra={
|
extra={
|
||||||
@@ -1492,77 +1656,33 @@ async def adjust_payment_for_tokens(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
if billing_key.reserved_balance < deducted_max_cost:
|
charged = await _charge_reservation_rows(
|
||||||
logger.error(
|
session,
|
||||||
"reserved_balance below deducted_max_cost on refund finalization — clamping to 0",
|
billing_key_hash=billing_key.hashed_key,
|
||||||
extra={
|
key_hash=key.hashed_key,
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
reserved_msats=deducted_max_cost,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
charge_msats=total_cost_msats,
|
||||||
"reserved_balance": billing_key.reserved_balance,
|
|
||||||
"deducted_max_cost": deducted_max_cost,
|
|
||||||
"total_cost_msats": total_cost_msats,
|
|
||||||
"refund_amount": refund,
|
|
||||||
"balance": billing_key.balance,
|
|
||||||
"total_spent": billing_key.total_spent,
|
|
||||||
"model": model,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
refund_safe_reserved = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
|
||||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
)
|
||||||
|
if charged:
|
||||||
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
refund_stmt = (
|
if not charged:
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=refund_safe_reserved,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(refund_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
# 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,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
reserved_balance=child_refund_safe_reserved,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
if result.rowcount == 0:
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to finalize payment - releasing reservation",
|
"Failed to finalize payment - releasing reservation",
|
||||||
extra={
|
extra={
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": billing_log_hash,
|
||||||
"deducted_max_cost": deducted_max_cost,
|
"deducted_max_cost": deducted_max_cost,
|
||||||
"current_reserved_balance": billing_key.reserved_balance,
|
|
||||||
"total_cost": total_cost_msats,
|
"total_cost": total_cost_msats,
|
||||||
"model": model,
|
"model": model,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
cost.charged_msats = 0
|
||||||
await release_reservation_only()
|
await release_reservation_only()
|
||||||
else:
|
else:
|
||||||
cost.total_msats = total_cost_msats
|
cost.total_msats = total_cost_msats
|
||||||
|
cost.charged_msats = total_cost_msats
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
|
|||||||
+50
-21
@@ -69,7 +69,9 @@ async def require_admin_api(request: Request) -> None:
|
|||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
result = await session.exec(select(CliToken).where(CliToken.token == token))
|
result = await session.exec(select(CliToken).where(CliToken.token == token))
|
||||||
cli_token = result.first()
|
cli_token = result.first()
|
||||||
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
|
if cli_token and (
|
||||||
|
cli_token.expires_at is None or cli_token.expires_at > now_ts
|
||||||
|
):
|
||||||
cli_token.last_used_at = now_ts
|
cli_token.last_used_at = now_ts
|
||||||
session.add(cli_token)
|
session.add(cli_token)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -107,7 +109,7 @@ async def get_temporary_balances_api(
|
|||||||
# Aggregate totals across the whole (search-filtered) set, not just the
|
# Aggregate totals across the whole (search-filtered) set, not just the
|
||||||
# current page. Balance counts only parent (non-child) keys to avoid
|
# current page. Balance counts only parent (non-child) keys to avoid
|
||||||
# double-counting, since child keys draw from their parent's balance.
|
# double-counting, since child keys draw from their parent's balance.
|
||||||
totals_result = await session.exec(
|
balance_totals_result = await session.exec(
|
||||||
select(
|
select(
|
||||||
func.coalesce(
|
func.coalesce(
|
||||||
func.sum(
|
func.sum(
|
||||||
@@ -118,11 +120,44 @@ async def get_temporary_balances_api(
|
|||||||
),
|
),
|
||||||
0,
|
0,
|
||||||
),
|
),
|
||||||
|
func.coalesce(
|
||||||
|
func.sum(
|
||||||
|
case(
|
||||||
|
(
|
||||||
|
col(ApiKey.parent_key_hash).is_(None),
|
||||||
|
ApiKey.reserved_balance,
|
||||||
|
),
|
||||||
|
else_=0,
|
||||||
|
)
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
),
|
||||||
|
func.coalesce(
|
||||||
|
func.sum(
|
||||||
|
case(
|
||||||
|
(
|
||||||
|
col(ApiKey.parent_key_hash).is_(None),
|
||||||
|
col(ApiKey.balance) - col(ApiKey.reserved_balance),
|
||||||
|
),
|
||||||
|
else_=0,
|
||||||
|
)
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
),
|
||||||
|
).where(*filters)
|
||||||
|
)
|
||||||
|
(
|
||||||
|
total_balance,
|
||||||
|
total_reserved_balance,
|
||||||
|
total_available_balance,
|
||||||
|
) = balance_totals_result.one()
|
||||||
|
usage_totals_result = await session.exec(
|
||||||
|
select(
|
||||||
func.coalesce(func.sum(ApiKey.total_spent), 0),
|
func.coalesce(func.sum(ApiKey.total_spent), 0),
|
||||||
func.coalesce(func.sum(ApiKey.total_requests), 0),
|
func.coalesce(func.sum(ApiKey.total_requests), 0),
|
||||||
).where(*filters)
|
).where(*filters)
|
||||||
)
|
)
|
||||||
total_balance, total_spent, total_requests = totals_result.one()
|
total_spent, total_requests = usage_totals_result.one()
|
||||||
|
|
||||||
# Latest created first; keys with no created_at (legacy rows) sort last.
|
# Latest created first; keys with no created_at (legacy rows) sort last.
|
||||||
# Use an explicit CASE rather than relying on dialect NULL-ordering so
|
# Use an explicit CASE rather than relying on dialect NULL-ordering so
|
||||||
@@ -143,6 +178,10 @@ async def get_temporary_balances_api(
|
|||||||
{
|
{
|
||||||
"hashed_key": key.hashed_key,
|
"hashed_key": key.hashed_key,
|
||||||
"balance": key.balance,
|
"balance": key.balance,
|
||||||
|
"reserved_balance": key.reserved_balance,
|
||||||
|
"available_balance": (
|
||||||
|
key.total_balance if key.parent_key_hash is None else None
|
||||||
|
),
|
||||||
"total_spent": key.total_spent,
|
"total_spent": key.total_spent,
|
||||||
"total_requests": key.total_requests,
|
"total_requests": key.total_requests,
|
||||||
"refund_address": key.refund_address,
|
"refund_address": key.refund_address,
|
||||||
@@ -158,6 +197,8 @@ async def get_temporary_balances_api(
|
|||||||
"total": total,
|
"total": total,
|
||||||
"totals": {
|
"totals": {
|
||||||
"total_balance": total_balance,
|
"total_balance": total_balance,
|
||||||
|
"total_reserved_balance": total_reserved_balance,
|
||||||
|
"total_available_balance": total_available_balance,
|
||||||
"total_spent": total_spent,
|
"total_spent": total_spent,
|
||||||
"total_requests": total_requests,
|
"total_requests": total_requests,
|
||||||
},
|
},
|
||||||
@@ -256,16 +297,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) ->
|
|||||||
secret = await get_secret(session)
|
secret = await get_secret(session)
|
||||||
|
|
||||||
if not secret.admin_password_hash:
|
if not secret.admin_password_hash:
|
||||||
raise HTTPException(
|
raise HTTPException(status_code=500, detail="Admin password not configured")
|
||||||
status_code=500, detail="Admin password not configured"
|
|
||||||
)
|
|
||||||
|
|
||||||
if not vault.verify_password(
|
if not vault.verify_password(
|
||||||
password_update.current_password, secret.admin_password_hash
|
password_update.current_password, secret.admin_password_hash
|
||||||
):
|
):
|
||||||
raise HTTPException(
|
raise HTTPException(status_code=401, detail="Current password is incorrect")
|
||||||
status_code=401, detail="Current password is incorrect"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Validate new password
|
# Validate new password
|
||||||
new_password = password_update.new_password.strip()
|
new_password = password_update.new_password.strip()
|
||||||
@@ -883,9 +920,7 @@ async def _active_ppq_claim_in_session(session: AsyncSession, provider_id: int)
|
|||||||
return claim is not None and not claim.collected and not claim.swept
|
return claim is not None and not claim.collected and not claim.swept
|
||||||
|
|
||||||
|
|
||||||
def _require_valid_ppq_auto_topup(
|
def _require_valid_ppq_auto_topup(provider_type: str, settings: dict | None) -> None:
|
||||||
provider_type: str, settings: dict | None
|
|
||||||
) -> None:
|
|
||||||
"""Reject PPQ auto top-up settings the worker would later refuse."""
|
"""Reject PPQ auto top-up settings the worker would later refuse."""
|
||||||
if provider_type != "ppqai":
|
if provider_type != "ppqai":
|
||||||
return
|
return
|
||||||
@@ -1017,9 +1052,7 @@ async def create_upstream_provider(
|
|||||||
else:
|
else:
|
||||||
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
||||||
|
|
||||||
_require_valid_ppq_auto_topup(
|
_require_valid_ppq_auto_topup(payload.provider_type, payload.provider_settings)
|
||||||
payload.provider_type, payload.provider_settings
|
|
||||||
)
|
|
||||||
|
|
||||||
provider = UpstreamProviderRow(
|
provider = UpstreamProviderRow(
|
||||||
slug=slug,
|
slug=slug,
|
||||||
@@ -1078,9 +1111,7 @@ async def update_upstream_provider_by_slug(
|
|||||||
lookup = _validate_slug(payload.slug)
|
lookup = _validate_slug(payload.slug)
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
result = await session.exec(
|
result = await session.exec(
|
||||||
select(UpstreamProviderRow).where(
|
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup)
|
||||||
UpstreamProviderRow.slug == lookup
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
provider = result.first()
|
provider = result.first()
|
||||||
if not provider:
|
if not provider:
|
||||||
@@ -1910,9 +1941,7 @@ async def get_transactions_api(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@admin_router.get(
|
@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)])
|
||||||
"/api/lightning-invoices", dependencies=[Depends(require_admin_api)]
|
|
||||||
)
|
|
||||||
async def get_lightning_invoices_api(
|
async def get_lightning_invoices_api(
|
||||||
status: str | None = None,
|
status: str | None = None,
|
||||||
purpose: str | None = None,
|
purpose: str | None = None,
|
||||||
|
|||||||
+94
-28
@@ -171,6 +171,55 @@ async def reset_all_reserved_balances(session: AsyncSession) -> None:
|
|||||||
logger.info("Reset reserved balances on startup")
|
logger.info("Reset reserved balances on startup")
|
||||||
|
|
||||||
|
|
||||||
|
async def _transition_stale_reservation(
|
||||||
|
session: AsyncSession, reservation_id: str, cutoff: int
|
||||||
|
) -> bool:
|
||||||
|
"""Mark one reservation released iff its lease is still older than cutoff.
|
||||||
|
|
||||||
|
``created_at`` doubles as the heartbeat lease timestamp, so the guard must
|
||||||
|
be part of this update: a reservation renewed between the sweeper's select
|
||||||
|
and this transition is in flight and must survive.
|
||||||
|
"""
|
||||||
|
transition = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.id) == reservation_id)
|
||||||
|
.where(col(ReservationRelease.status) == "active")
|
||||||
|
.where(col(ReservationRelease.created_at) < cutoff)
|
||||||
|
.values(status="released")
|
||||||
|
)
|
||||||
|
return bool(transition.rowcount == 1)
|
||||||
|
|
||||||
|
|
||||||
|
async def _release_legacy_aggregate(
|
||||||
|
session: AsyncSession,
|
||||||
|
key_hash: str,
|
||||||
|
observed_reserved: int,
|
||||||
|
observed_reserved_at: int | None,
|
||||||
|
) -> bool:
|
||||||
|
"""Zero one legacy aggregate reservation iff it is exactly as observed.
|
||||||
|
|
||||||
|
A new reservation committing between the sweeper's read and this update
|
||||||
|
changes ``reserved_balance``/``reserved_at`` in the same transaction that
|
||||||
|
creates its durable row, so this compare-and-swap fails instead of erasing
|
||||||
|
the newcomer's reserved funds.
|
||||||
|
"""
|
||||||
|
if observed_reserved <= 0:
|
||||||
|
return False
|
||||||
|
reserved_at_guard = (
|
||||||
|
col(ApiKey.reserved_at).is_(None)
|
||||||
|
if observed_reserved_at is None
|
||||||
|
else col(ApiKey.reserved_at) == observed_reserved_at
|
||||||
|
)
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.where(col(ApiKey.reserved_balance) == observed_reserved)
|
||||||
|
.where(reserved_at_guard)
|
||||||
|
.values(reserved_balance=0, reserved_at=None)
|
||||||
|
)
|
||||||
|
return bool(result.rowcount == 1)
|
||||||
|
|
||||||
|
|
||||||
async def release_stale_reservations(
|
async def release_stale_reservations(
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
max_age_seconds: int,
|
max_age_seconds: int,
|
||||||
@@ -191,25 +240,22 @@ async def release_stale_reservations(
|
|||||||
col(ReservationRelease.billing_key_hash) == key_hash,
|
col(ReservationRelease.billing_key_hash) == key_hash,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
reservations = (await session.exec(query)).all()
|
# Capture primitives: a repair rollback below would expire ORM instances.
|
||||||
|
reservation_rows = [
|
||||||
|
(r.id, r.key_hash, r.billing_key_hash, r.reserved_msats)
|
||||||
|
for r in (await session.exec(query)).all()
|
||||||
|
]
|
||||||
released = 0
|
released = 0
|
||||||
|
|
||||||
for reservation in reservations:
|
for res_id, res_key_hash, res_billing_hash, res_msats in reservation_rows:
|
||||||
transition = await session.exec( # type: ignore[call-overload]
|
if not await _transition_stale_reservation(session, res_id, cutoff):
|
||||||
update(ReservationRelease)
|
|
||||||
.where(col(ReservationRelease.id) == reservation.id)
|
|
||||||
.where(col(ReservationRelease.status) == "active")
|
|
||||||
.values(status="released")
|
|
||||||
)
|
|
||||||
if transition.rowcount != 1:
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
values = {
|
values = {
|
||||||
"reserved_balance": col(ApiKey.reserved_balance)
|
"reserved_balance": col(ApiKey.reserved_balance) - res_msats,
|
||||||
- reservation.reserved_msats,
|
|
||||||
"reserved_at": case(
|
"reserved_at": case(
|
||||||
(
|
(
|
||||||
col(ApiKey.reserved_balance) - reservation.reserved_msats > 0,
|
col(ApiKey.reserved_balance) - res_msats > 0,
|
||||||
col(ApiKey.reserved_at),
|
col(ApiKey.reserved_at),
|
||||||
),
|
),
|
||||||
else_=None,
|
else_=None,
|
||||||
@@ -217,24 +263,42 @@ async def release_stale_reservations(
|
|||||||
}
|
}
|
||||||
parent_result = await session.exec( # type: ignore[call-overload]
|
parent_result = await session.exec( # type: ignore[call-overload]
|
||||||
update(ApiKey)
|
update(ApiKey)
|
||||||
.where(col(ApiKey.hashed_key) == reservation.billing_key_hash)
|
.where(col(ApiKey.hashed_key) == res_billing_hash)
|
||||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
.where(col(ApiKey.reserved_balance) >= res_msats)
|
||||||
.values(**values)
|
.values(**values)
|
||||||
)
|
)
|
||||||
if parent_result.rowcount != 1:
|
aggregates_ok = parent_result.rowcount == 1
|
||||||
await session.rollback()
|
if aggregates_ok and res_billing_hash != res_key_hash:
|
||||||
return 0
|
|
||||||
|
|
||||||
if reservation.billing_key_hash != reservation.key_hash:
|
|
||||||
child_result = await session.exec( # type: ignore[call-overload]
|
child_result = await session.exec( # type: ignore[call-overload]
|
||||||
update(ApiKey)
|
update(ApiKey)
|
||||||
.where(col(ApiKey.hashed_key) == reservation.key_hash)
|
.where(col(ApiKey.hashed_key) == res_key_hash)
|
||||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
.where(col(ApiKey.reserved_balance) >= res_msats)
|
||||||
.values(**values)
|
.values(**values)
|
||||||
)
|
)
|
||||||
if child_result.rowcount != 1:
|
aggregates_ok = child_result.rowcount == 1
|
||||||
await session.rollback()
|
|
||||||
return 0
|
if not aggregates_ok:
|
||||||
|
# The aggregates no longer hold this reservation's msats — the
|
||||||
|
# durable row is corrupt. Repair by terminalizing it WITHOUT
|
||||||
|
# subtracting uncertain aggregates (legacy cleanup below reconciles
|
||||||
|
# any stale remainder) and keep sweeping the rest of the batch:
|
||||||
|
# one corrupt row must not poison all stale cleanup.
|
||||||
|
await session.rollback()
|
||||||
|
if await _transition_stale_reservation(session, res_id, cutoff):
|
||||||
|
await session.commit()
|
||||||
|
released += 1
|
||||||
|
logger.error(
|
||||||
|
"Released corrupt stale reservation without aggregate subtraction",
|
||||||
|
extra={
|
||||||
|
"reservation_id": res_id,
|
||||||
|
"billing_key_hash": res_billing_hash[:8] + "...",
|
||||||
|
"reserved_msats": res_msats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
# Commit each release on its own so a later corrupt record's rollback
|
||||||
|
# cannot discard the healthy releases already processed in this batch.
|
||||||
|
await session.commit()
|
||||||
released += 1
|
released += 1
|
||||||
|
|
||||||
# Rolling upgrades can leave aggregate reservations created before durable
|
# Rolling upgrades can leave aggregate reservations created before durable
|
||||||
@@ -256,6 +320,8 @@ async def release_stale_reservations(
|
|||||||
)
|
)
|
||||||
|
|
||||||
for legacy_key in (await session.exec(legacy_query)).all():
|
for legacy_key in (await session.exec(legacy_query)).all():
|
||||||
|
observed_reserved = legacy_key.reserved_balance
|
||||||
|
observed_reserved_at = legacy_key.reserved_at
|
||||||
active_owner = (
|
active_owner = (
|
||||||
await session.exec(
|
await session.exec(
|
||||||
select(ReservationRelease.id)
|
select(ReservationRelease.id)
|
||||||
@@ -272,10 +338,10 @@ async def release_stale_reservations(
|
|||||||
).first()
|
).first()
|
||||||
if active_owner is not None:
|
if active_owner is not None:
|
||||||
continue
|
continue
|
||||||
legacy_key.reserved_balance = 0
|
if await _release_legacy_aggregate(
|
||||||
legacy_key.reserved_at = None
|
session, legacy_key.hashed_key, observed_reserved, observed_reserved_at
|
||||||
session.add(legacy_key)
|
):
|
||||||
released += 1
|
released += 1
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
if released:
|
if released:
|
||||||
|
|||||||
@@ -34,6 +34,8 @@ class CostData(BaseModel):
|
|||||||
cache_creation_input_tokens: int = 0
|
cache_creation_input_tokens: int = 0
|
||||||
cache_read_msats: int = 0
|
cache_read_msats: int = 0
|
||||||
cache_creation_msats: int = 0
|
cache_creation_msats: int = 0
|
||||||
|
# Actual debit after finalization; None means settlement has not run yet.
|
||||||
|
charged_msats: int | None = None
|
||||||
|
|
||||||
|
|
||||||
class MaxCostData(CostData):
|
class MaxCostData(CostData):
|
||||||
@@ -247,9 +249,7 @@ async def calculate_cost(
|
|||||||
"Token counts %s in the upstream response but cannot be "
|
"Token counts %s in the upstream response but cannot be "
|
||||||
"priced; the request will appear in dashboards with the "
|
"priced; the request will appear in dashboards with the "
|
||||||
"raw counts and a fixed max-cost charge.",
|
"raw counts and a fixed max-cost charge.",
|
||||||
"are present"
|
"are present" if (input_tokens > 0 or output_tokens > 0) else "are zero",
|
||||||
if (input_tokens > 0 or output_tokens > 0)
|
|
||||||
else "are zero",
|
|
||||||
extra={
|
extra={
|
||||||
"base_cost_msats": max_cost,
|
"base_cost_msats": max_cost,
|
||||||
"model": response_data.get("model", "unknown"),
|
"model": response_data.get("model", "unknown"),
|
||||||
@@ -326,9 +326,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
|||||||
# actually deducts from the balance. For non-BYOK providers (e.g.
|
# actually deducts from the balance. For non-BYOK providers (e.g.
|
||||||
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
|
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
|
||||||
# fall through to the normal ``cost`` lookup below.
|
# fall through to the normal ``cost`` lookup below.
|
||||||
upstream_cost = _coerce_usd(
|
upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost"))
|
||||||
cost_details.get("upstream_inference_cost")
|
|
||||||
)
|
|
||||||
if upstream_cost > 0 and usage_data.get("is_byok"):
|
if upstream_cost > 0 and usage_data.get("is_byok"):
|
||||||
byok_fee = _coerce_usd(usage_data.get("cost"))
|
byok_fee = _coerce_usd(usage_data.get("cost"))
|
||||||
return upstream_cost + byok_fee
|
return upstream_cost + byok_fee
|
||||||
@@ -359,8 +357,7 @@ def _get_pricing_rates(
|
|||||||
``None`` means configured fixed pricing should be used by the caller.
|
``None`` means configured fixed pricing should be used by the caller.
|
||||||
"""
|
"""
|
||||||
if settings.fixed_pricing and (
|
if settings.fixed_pricing and (
|
||||||
settings.fixed_per_1k_input_tokens
|
settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens
|
||||||
or settings.fixed_per_1k_output_tokens
|
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -416,12 +413,8 @@ def _get_pricing_rates(
|
|||||||
usd_per_sat = sats_usd_price()
|
usd_per_sat = sats_usd_price()
|
||||||
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||||
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||||
cache_read_usd = _coerce_usd(
|
cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost"))
|
||||||
pricing.get("cache_read_input_token_cost")
|
cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost"))
|
||||||
)
|
|
||||||
cache_write_usd = _coerce_usd(
|
|
||||||
pricing.get("cache_creation_input_token_cost")
|
|
||||||
)
|
|
||||||
mscr_1k = (
|
mscr_1k = (
|
||||||
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||||
if cache_read_usd > 0
|
if cache_read_usd > 0
|
||||||
@@ -525,9 +518,7 @@ def _calculate_from_usd_cost(
|
|||||||
regular_weight = input_tokens * input_rate
|
regular_weight = input_tokens * input_rate
|
||||||
cache_read_weight = cache_read_tokens * cache_read_rate
|
cache_read_weight = cache_read_tokens * cache_read_rate
|
||||||
cache_creation_weight = cache_creation_tokens * cache_creation_rate
|
cache_creation_weight = cache_creation_tokens * cache_creation_rate
|
||||||
total_input_weight = (
|
total_input_weight = regular_weight + cache_read_weight + cache_creation_weight
|
||||||
regular_weight + cache_read_weight + cache_creation_weight
|
|
||||||
)
|
|
||||||
if total_input_weight > 0:
|
if total_input_weight > 0:
|
||||||
cache_read_msats = int(
|
cache_read_msats = int(
|
||||||
round(
|
round(
|
||||||
|
|||||||
+54
-69
@@ -83,6 +83,24 @@ def _cost_field(
|
|||||||
return value if isinstance(value, (int, float)) else default
|
return value if isinstance(value, (int, float)) else default
|
||||||
|
|
||||||
|
|
||||||
|
def _settled_cost_msats(cost_data: CostMetadata) -> int:
|
||||||
|
charged = _cost_field(cost_data, "charged_msats", -1)
|
||||||
|
if charged >= 0:
|
||||||
|
return int(charged)
|
||||||
|
return int(_cost_field(cost_data, "total_msats"))
|
||||||
|
|
||||||
|
|
||||||
|
def _published_cost(cost_data: CostMetadata) -> dict[str, Any]:
|
||||||
|
cost = dict(cost_data) if isinstance(cost_data, dict) else cost_data.dict()
|
||||||
|
computed_msats = int(_cost_field(cost_data, "total_msats"))
|
||||||
|
settled_msats = _settled_cost_msats(cost_data)
|
||||||
|
if computed_msats != settled_msats:
|
||||||
|
cost["computed_msats"] = computed_msats
|
||||||
|
cost["total_msats"] = settled_msats
|
||||||
|
cost["charged_msats"] = settled_msats
|
||||||
|
return cost
|
||||||
|
|
||||||
|
|
||||||
def _inject_cost_response_headers(
|
def _inject_cost_response_headers(
|
||||||
headers: dict[str, str], cost_data: CostMetadata
|
headers: dict[str, str], cost_data: CostMetadata
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -93,9 +111,11 @@ def _inject_cost_response_headers(
|
|||||||
usage tracking entry — without them, x-cashu requests show 0.0 for all
|
usage tracking entry — without them, x-cashu requests show 0.0 for all
|
||||||
sat cost fields.
|
sat cost fields.
|
||||||
"""
|
"""
|
||||||
headers["X-Routstr-Cost-Msats"] = str(
|
settled_msats = _settled_cost_msats(cost_data)
|
||||||
int(_cost_field(cost_data, "total_msats"))
|
computed_msats = int(_cost_field(cost_data, "total_msats"))
|
||||||
)
|
headers["X-Routstr-Cost-Msats"] = str(settled_msats)
|
||||||
|
if computed_msats != settled_msats:
|
||||||
|
headers["X-Routstr-Computed-Cost-Msats"] = str(computed_msats)
|
||||||
headers["X-Routstr-Input-Cost-Msats"] = str(
|
headers["X-Routstr-Input-Cost-Msats"] = str(
|
||||||
int(_cost_field(cost_data, "input_msats"))
|
int(_cost_field(cost_data, "input_msats"))
|
||||||
)
|
)
|
||||||
@@ -122,11 +142,14 @@ def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> Non
|
|||||||
# data always overwrites any upstream-provided cost values. Using
|
# data always overwrites any upstream-provided cost values. Using
|
||||||
# setdefault would silently keep stale upstream values and drop our
|
# setdefault would silently keep stale upstream values and drop our
|
||||||
# calculated msats breakdown.
|
# calculated msats breakdown.
|
||||||
|
computed_msats = int(_cost_field(cost_data, "total_msats"))
|
||||||
|
settled_msats = _settled_cost_msats(cost_data)
|
||||||
cost_obj: dict[str, int | float] = {
|
cost_obj: dict[str, int | float] = {
|
||||||
"base_msats": int(_cost_field(cost_data, "base_msats")),
|
"base_msats": int(_cost_field(cost_data, "base_msats")),
|
||||||
"input_msats": int(_cost_field(cost_data, "input_msats")),
|
"input_msats": int(_cost_field(cost_data, "input_msats")),
|
||||||
"output_msats": int(_cost_field(cost_data, "output_msats")),
|
"output_msats": int(_cost_field(cost_data, "output_msats")),
|
||||||
"total_msats": int(_cost_field(cost_data, "total_msats")),
|
"total_msats": settled_msats,
|
||||||
|
"charged_msats": settled_msats,
|
||||||
"cache_read_input_tokens": int(
|
"cache_read_input_tokens": int(
|
||||||
_cost_field(cost_data, "cache_read_input_tokens")
|
_cost_field(cost_data, "cache_read_input_tokens")
|
||||||
),
|
),
|
||||||
@@ -134,15 +157,15 @@ def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> Non
|
|||||||
_cost_field(cost_data, "cache_creation_input_tokens")
|
_cost_field(cost_data, "cache_creation_input_tokens")
|
||||||
),
|
),
|
||||||
"cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")),
|
"cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")),
|
||||||
"cache_creation_msats": int(
|
"cache_creation_msats": int(_cost_field(cost_data, "cache_creation_msats")),
|
||||||
_cost_field(cost_data, "cache_creation_msats")
|
|
||||||
),
|
|
||||||
}
|
}
|
||||||
|
if computed_msats != settled_msats:
|
||||||
|
cost_obj["computed_msats"] = computed_msats
|
||||||
total_usd = float(_cost_field(cost_data, "total_usd", 0.0))
|
total_usd = float(_cost_field(cost_data, "total_usd", 0.0))
|
||||||
if total_usd:
|
if total_usd:
|
||||||
cost_obj["total_usd"] = total_usd
|
cost_obj["total_usd"] = total_usd
|
||||||
usage["cost"] = cost_obj
|
usage["cost"] = cost_obj
|
||||||
usage["cost_sats"] = int(_cost_field(cost_data, "total_msats")) // 1000
|
usage["cost_sats"] = settled_msats // 1000
|
||||||
|
|
||||||
|
|
||||||
def _is_json_content_type(content_type: str | None) -> bool:
|
def _is_json_content_type(content_type: str | None) -> bool:
|
||||||
@@ -349,14 +372,8 @@ class BaseUpstreamProvider:
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Unifies the injection of cost and usage metadata across all completion types."""
|
"""Unifies the injection of cost and usage metadata across all completion types."""
|
||||||
self._apply_provider_field(response_json)
|
self._apply_provider_field(response_json)
|
||||||
if isinstance(cost_data, dict):
|
cost_dict = _published_cost(cost_data)
|
||||||
total_msats = cost_data.get("total_msats", 0)
|
sats_cost = cost_dict["total_msats"] // 1000
|
||||||
cost_dict = cost_data
|
|
||||||
else:
|
|
||||||
total_msats = cost_data.total_msats
|
|
||||||
cost_dict = cost_data.dict()
|
|
||||||
|
|
||||||
sats_cost = total_msats // 1000
|
|
||||||
|
|
||||||
# Inject the shared SDK cost contract into every usage shape.
|
# Inject the shared SDK cost contract into every usage shape.
|
||||||
if isinstance(response_json.get("usage"), dict):
|
if isinstance(response_json.get("usage"), dict):
|
||||||
@@ -370,6 +387,14 @@ class BaseUpstreamProvider:
|
|||||||
message["usage"]["remaining_balance_msats"] = key.balance
|
message["usage"]["remaining_balance_msats"] = key.balance
|
||||||
self._fold_cache_into_input_tokens(message["usage"])
|
self._fold_cache_into_input_tokens(message["usage"])
|
||||||
|
|
||||||
|
nested_response = response_json.get("response")
|
||||||
|
if isinstance(nested_response, dict) and isinstance(
|
||||||
|
nested_response.get("usage"), dict
|
||||||
|
):
|
||||||
|
_inject_cost_into_usage(nested_response, cost_data)
|
||||||
|
nested_response["usage"]["remaining_balance_msats"] = key.balance
|
||||||
|
self._fold_cache_into_input_tokens(nested_response["usage"])
|
||||||
|
|
||||||
# Unified Routstr metadata
|
# Unified Routstr metadata
|
||||||
response_json["metadata"] = response_json.get("metadata", {})
|
response_json["metadata"] = response_json.get("metadata", {})
|
||||||
response_json["metadata"]["routstr"] = {
|
response_json["metadata"]["routstr"] = {
|
||||||
@@ -1307,18 +1332,12 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
self._fold_cache_into_input_tokens(response_json["usage"])
|
self._fold_cache_into_input_tokens(response_json["usage"])
|
||||||
|
|
||||||
# Keep detailed cost
|
published_cost = _published_cost(cost_data)
|
||||||
|
published_cost["sats_cost"] = published_cost["total_msats"] // 1000
|
||||||
|
published_cost["remaining_balance_msats"] = remaining_balance_msats
|
||||||
response_json["metadata"] = response_json.get("metadata", {})
|
response_json["metadata"] = response_json.get("metadata", {})
|
||||||
response_json["metadata"]["routstr"] = {"cost": cost_data}
|
response_json["metadata"]["routstr"] = {"cost": published_cost.copy()}
|
||||||
response_json["metadata"]["routstr"]["cost"]["sats_cost"] = (
|
response_json["cost"] = published_cost
|
||||||
cost_data.get("total_msats", 0) // 1000
|
|
||||||
)
|
|
||||||
response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = (
|
|
||||||
remaining_balance_msats
|
|
||||||
)
|
|
||||||
response_json["cost"] = cost_data
|
|
||||||
response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000
|
|
||||||
response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats
|
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Payment adjustment completed for non-streaming",
|
"Payment adjustment completed for non-streaming",
|
||||||
@@ -1609,24 +1628,6 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
remaining_balance_msats = fresh_key.balance
|
|
||||||
sats_cost = cost_data.get("total_msats", 0) // 1000
|
|
||||||
|
|
||||||
if (
|
|
||||||
"response" in usage_chunk_data
|
|
||||||
and isinstance(usage_chunk_data["response"], dict)
|
|
||||||
and "usage" in usage_chunk_data["response"]
|
|
||||||
):
|
|
||||||
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"][
|
|
||||||
"remaining_balance_msats"
|
|
||||||
] = remaining_balance_msats
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
self.inject_cost_metadata(
|
self.inject_cost_metadata(
|
||||||
usage_chunk_data, cost_data, fresh_key
|
usage_chunk_data, cost_data, fresh_key
|
||||||
@@ -1742,18 +1743,12 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
self._fold_cache_into_input_tokens(response_json["usage"])
|
self._fold_cache_into_input_tokens(response_json["usage"])
|
||||||
|
|
||||||
# Keep detailed cost
|
published_cost = _published_cost(cost_data)
|
||||||
|
published_cost["sats_cost"] = published_cost["total_msats"] // 1000
|
||||||
|
published_cost["remaining_balance_msats"] = remaining_balance_msats
|
||||||
response_json["metadata"] = response_json.get("metadata", {})
|
response_json["metadata"] = response_json.get("metadata", {})
|
||||||
response_json["metadata"]["routstr"] = {"cost": cost_data}
|
response_json["metadata"]["routstr"] = {"cost": published_cost.copy()}
|
||||||
response_json["metadata"]["routstr"]["cost"]["sats_cost"] = (
|
response_json["cost"] = published_cost
|
||||||
cost_data.get("total_msats", 0) // 1000
|
|
||||||
)
|
|
||||||
response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = (
|
|
||||||
remaining_balance_msats
|
|
||||||
)
|
|
||||||
response_json["cost"] = cost_data
|
|
||||||
response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000
|
|
||||||
response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats
|
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Payment adjustment completed for non-streaming Responses API",
|
"Payment adjustment completed for non-streaming Responses API",
|
||||||
@@ -2747,9 +2742,7 @@ class BaseUpstreamProvider:
|
|||||||
event_type = str(event.get("type") or "")
|
event_type = str(event.get("type") or "")
|
||||||
prefix = f"event: {event_type}\n" if event_type else ""
|
prefix = f"event: {event_type}\n" if event_type else ""
|
||||||
buffered[index] = annotated._replace(
|
buffered[index] = annotated._replace(
|
||||||
sse_bytes=(
|
sse_bytes=(f"{prefix}data: {json.dumps(event)}\n\n".encode())
|
||||||
f"{prefix}data: {json.dumps(event)}\n\n".encode()
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
async def replay() -> AsyncGenerator[bytes, None]:
|
async def replay() -> AsyncGenerator[bytes, None]:
|
||||||
@@ -3822,11 +3815,7 @@ class BaseUpstreamProvider:
|
|||||||
if "provider" not in data_json:
|
if "provider" not in data_json:
|
||||||
self._apply_provider_field(data_json)
|
self._apply_provider_field(data_json)
|
||||||
changed = True
|
changed = True
|
||||||
if (
|
if cost_data and "usage" in data_json and data_json["usage"]:
|
||||||
cost_data
|
|
||||||
and "usage" in data_json
|
|
||||||
and data_json["usage"]
|
|
||||||
):
|
|
||||||
_inject_cost_into_usage(data_json, cost_data)
|
_inject_cost_into_usage(data_json, cost_data)
|
||||||
changed = True
|
changed = True
|
||||||
if changed:
|
if changed:
|
||||||
@@ -4811,11 +4800,7 @@ class BaseUpstreamProvider:
|
|||||||
if "provider" not in data_json:
|
if "provider" not in data_json:
|
||||||
self._apply_provider_field(data_json)
|
self._apply_provider_field(data_json)
|
||||||
changed = True
|
changed = True
|
||||||
if (
|
if cost_data and "usage" in data_json and data_json["usage"]:
|
||||||
cost_data
|
|
||||||
and "usage" in data_json
|
|
||||||
and data_json["usage"]
|
|
||||||
):
|
|
||||||
_inject_cost_into_usage(data_json, cost_data)
|
_inject_cost_into_usage(data_json, cost_data)
|
||||||
changed = True
|
changed = True
|
||||||
if changed:
|
if changed:
|
||||||
|
|||||||
+72
-128
@@ -10,17 +10,18 @@ from urllib.parse import urlsplit, urlunsplit
|
|||||||
|
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
from fastapi.responses import Response, StreamingResponse
|
from fastapi.responses import Response, StreamingResponse
|
||||||
from sqlalchemy import case
|
|
||||||
from sqlmodel import col, update
|
|
||||||
|
|
||||||
from ..auth import (
|
from ..auth import (
|
||||||
ROUTSTR_FEE_PERCENT,
|
ROUTSTR_FEE_PERCENT,
|
||||||
ReservationSnapshot,
|
ReservationSnapshot,
|
||||||
|
_charge_reservation_rows,
|
||||||
_claim_reservation_for_charge,
|
_claim_reservation_for_charge,
|
||||||
|
_stop_reservation_heartbeat,
|
||||||
_validate_reservation_snapshot,
|
_validate_reservation_snapshot,
|
||||||
get_billing_key,
|
get_billing_key,
|
||||||
get_reservation_snapshot,
|
get_reservation_snapshot,
|
||||||
payments_logger,
|
payments_logger,
|
||||||
|
release_reservation,
|
||||||
)
|
)
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.db import (
|
from ..core.db import (
|
||||||
@@ -191,7 +192,9 @@ def _resolve_ehbp_target_url(
|
|||||||
otherwise the header is ignored so callers cannot redirect other providers
|
otherwise the header is ignored so callers cannot redirect other providers
|
||||||
or leak upstream API keys.
|
or leak upstream API keys.
|
||||||
"""
|
"""
|
||||||
override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
|
override_header = (
|
||||||
|
profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
|
||||||
|
)
|
||||||
if not override_header:
|
if not override_header:
|
||||||
return target_url
|
return target_url
|
||||||
enclave_url = _get_header_case_insensitive(headers, override_header)
|
enclave_url = _get_header_case_insensitive(headers, override_header)
|
||||||
@@ -295,9 +298,7 @@ def _build_cost_info(
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _inject_cost_response_headers(
|
def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None:
|
||||||
headers: dict[str, str], cost_info: dict
|
|
||||||
) -> None:
|
|
||||||
"""Add per-request cost headers to an EHBP response.
|
"""Add per-request cost headers to an EHBP response.
|
||||||
|
|
||||||
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
|
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
|
||||||
@@ -305,6 +306,8 @@ def _inject_cost_response_headers(
|
|||||||
the client/Tinfoil SDK can read without decrypting.
|
the client/Tinfoil SDK can read without decrypting.
|
||||||
"""
|
"""
|
||||||
headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"])
|
headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"])
|
||||||
|
if "computed_msats" in cost_info:
|
||||||
|
headers["X-Routstr-Computed-Cost-Msats"] = str(cost_info["computed_msats"])
|
||||||
headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"])
|
headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"])
|
||||||
headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"])
|
headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"])
|
||||||
|
|
||||||
@@ -375,9 +378,7 @@ async def _compute_ehbp_actual_cost(
|
|||||||
resolved_upstream_model = (
|
resolved_upstream_model = (
|
||||||
actual_model_obj.forwarded_model_id or actual_model_obj.id
|
actual_model_obj.forwarded_model_id or actual_model_obj.id
|
||||||
)
|
)
|
||||||
resolved_identity = _normalize_upstream_model_id(
|
resolved_identity = _normalize_upstream_model_id(resolved_upstream_model)
|
||||||
resolved_upstream_model
|
|
||||||
)
|
|
||||||
if resolved_identity != expected_identity:
|
if resolved_identity != expected_identity:
|
||||||
logger.info(
|
logger.info(
|
||||||
"EHBP served model differs from requested, using actual "
|
"EHBP served model differs from requested, using actual "
|
||||||
@@ -500,6 +501,18 @@ class EHBPForwardingTarget:
|
|||||||
profile: ConfidentialInferenceProfile | None = None
|
profile: ConfidentialInferenceProfile | None = None
|
||||||
|
|
||||||
|
|
||||||
|
async def _release_failed_ehbp_charge(
|
||||||
|
reservation: ReservationSnapshot, session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
if await release_reservation(reservation, session, reservation.reserved_msats):
|
||||||
|
return
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
logger.critical(
|
||||||
|
"Failed to release EHBP reservation after rejected charge",
|
||||||
|
extra={"reservation_id": reservation.release_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def finalize_ehbp_actual_cost_payment(
|
async def finalize_ehbp_actual_cost_payment(
|
||||||
key: ApiKey,
|
key: ApiKey,
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
@@ -507,61 +520,29 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
cost_info: dict,
|
cost_info: dict,
|
||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
) -> None:
|
) -> int:
|
||||||
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
|
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
|
||||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||||
await _validate_reservation_snapshot(key, reservation, session)
|
await _validate_reservation_snapshot(key, reservation, session)
|
||||||
if not await _claim_reservation_for_charge(reservation, session):
|
if not await _claim_reservation_for_charge(reservation, session):
|
||||||
return
|
return 0
|
||||||
reserved_cost_for_model = reservation.reserved_msats
|
reserved_cost_for_model = reservation.reserved_msats
|
||||||
billing_key = await get_billing_key(key, session)
|
billing_key = await get_billing_key(key, session)
|
||||||
key_hash = key.hashed_key
|
key_hash = key.hashed_key
|
||||||
billing_key_hash = billing_key.hashed_key
|
billing_key_hash = billing_key.hashed_key
|
||||||
total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
|
total_cost_msats = max(
|
||||||
|
0, int(cost_info.get("total_msats", reserved_cost_for_model))
|
||||||
|
)
|
||||||
now = int(time.time())
|
now = int(time.time())
|
||||||
|
|
||||||
safe_reserved = case(
|
charged = await _charge_reservation_rows(
|
||||||
(
|
session,
|
||||||
col(ApiKey.reserved_balance) >= reserved_cost_for_model,
|
billing_key_hash=billing_key_hash,
|
||||||
col(ApiKey.reserved_balance) - reserved_cost_for_model,
|
key_hash=key_hash,
|
||||||
),
|
reserved_msats=reserved_cost_for_model,
|
||||||
else_=0,
|
charge_msats=total_cost_msats,
|
||||||
)
|
)
|
||||||
cleared_reserved_at = case(
|
if not charged:
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) - reserved_cost_for_model > 0,
|
|
||||||
col(ApiKey.reserved_at),
|
|
||||||
),
|
|
||||||
else_=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
child_result = None
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
|
|
||||||
await session.rollback()
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to finalize EHBP usage-based payment",
|
"Failed to finalize EHBP usage-based payment",
|
||||||
extra={
|
extra={
|
||||||
@@ -570,13 +551,13 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
"model": model_id,
|
"model": model_id,
|
||||||
"reserved_cost_for_model": reserved_cost_for_model,
|
"reserved_cost_for_model": reserved_cost_for_model,
|
||||||
"total_cost_msats": total_cost_msats,
|
"total_cost_msats": total_cost_msats,
|
||||||
"parent_rowcount": result.rowcount,
|
|
||||||
"child_rowcount": getattr(child_result, "rowcount", None),
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return
|
await _release_failed_ehbp_charge(reservation, session)
|
||||||
|
return 0
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -609,6 +590,7 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
"finalized_at": now,
|
"finalized_at": now,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
return total_cost_msats
|
||||||
|
|
||||||
|
|
||||||
async def finalize_ehbp_max_cost_payment(
|
async def finalize_ehbp_max_cost_payment(
|
||||||
@@ -617,7 +599,7 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
max_cost_for_model: int,
|
max_cost_for_model: int,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
) -> None:
|
) -> int:
|
||||||
"""Finalize an EHBP bearer request by charging the reserved max cost.
|
"""Finalize an EHBP bearer request by charging the reserved max cost.
|
||||||
|
|
||||||
EHBP responses are encrypted, so Routstr cannot inspect token usage. Unlike
|
EHBP responses are encrypted, so Routstr cannot inspect token usage. Unlike
|
||||||
@@ -627,7 +609,7 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||||
await _validate_reservation_snapshot(key, reservation, session)
|
await _validate_reservation_snapshot(key, reservation, session)
|
||||||
if not await _claim_reservation_for_charge(reservation, session):
|
if not await _claim_reservation_for_charge(reservation, session):
|
||||||
return
|
return 0
|
||||||
max_cost_for_model = reservation.reserved_msats
|
max_cost_for_model = reservation.reserved_msats
|
||||||
billing_key = await get_billing_key(key, session)
|
billing_key = await get_billing_key(key, session)
|
||||||
key_hash = key.hashed_key
|
key_hash = key.hashed_key
|
||||||
@@ -635,63 +617,14 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
total_cost_msats = max(0, int(max_cost_for_model))
|
total_cost_msats = max(0, int(max_cost_for_model))
|
||||||
now = int(time.time())
|
now = int(time.time())
|
||||||
|
|
||||||
cleared_reserved_at = case(
|
charged = await _charge_reservation_rows(
|
||||||
(
|
session,
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
|
billing_key_hash=billing_key_hash,
|
||||||
col(ApiKey.reserved_at),
|
key_hash=key_hash,
|
||||||
),
|
reserved_msats=max_cost_for_model,
|
||||||
else_=None,
|
charge_msats=total_cost_msats,
|
||||||
)
|
)
|
||||||
safe_reserved = case(
|
if not charged:
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= max_cost_for_model,
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
child_safe_reserved = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= max_cost_for_model,
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
child_cleared_reserved_at = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
|
|
||||||
col(ApiKey.reserved_at),
|
|
||||||
),
|
|
||||||
else_=None,
|
|
||||||
)
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=child_safe_reserved,
|
|
||||||
reserved_at=child_cleared_reserved_at,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
else:
|
|
||||||
child_result = None
|
|
||||||
|
|
||||||
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
|
|
||||||
await session.rollback()
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to finalize EHBP max-cost payment",
|
"Failed to finalize EHBP max-cost payment",
|
||||||
extra={
|
extra={
|
||||||
@@ -699,14 +632,13 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
"billing_key_hash": billing_key_hash[:8] + "...",
|
"billing_key_hash": billing_key_hash[:8] + "...",
|
||||||
"model": model_id,
|
"model": model_id,
|
||||||
"max_cost_for_model": max_cost_for_model,
|
"max_cost_for_model": max_cost_for_model,
|
||||||
"parent_rowcount": result.rowcount,
|
|
||||||
"child_rowcount": getattr(child_result, "rowcount", None),
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return
|
await _release_failed_ehbp_charge(reservation, session)
|
||||||
|
return 0
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
await session.refresh(billing_key)
|
await session.refresh(billing_key)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
if billing_key.hashed_key != key.hashed_key:
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -739,6 +671,7 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
"finalized_at": now,
|
"finalized_at": now,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
return total_cost_msats
|
||||||
|
|
||||||
|
|
||||||
async def send_cashu_refund(
|
async def send_cashu_refund(
|
||||||
@@ -893,10 +826,9 @@ async def forward_ehbp_request(
|
|||||||
cost_info = await _compute_ehbp_actual_cost(
|
cost_info = await _compute_ehbp_actual_cost(
|
||||||
usage_header, model_obj, max_cost_for_model
|
usage_header, model_obj, max_cost_for_model
|
||||||
)
|
)
|
||||||
# Use the actual served model for billing when it differs from
|
|
||||||
# the requested model.
|
|
||||||
billing_model = cost_info.pop("actual_model", None) or model_obj.id
|
billing_model = cost_info.pop("actual_model", None) or model_obj.id
|
||||||
await finalize_ehbp_actual_cost_payment(
|
computed_msats = int(cost_info["total_msats"])
|
||||||
|
charged_msats = await finalize_ehbp_actual_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
@@ -904,7 +836,14 @@ async def forward_ehbp_request(
|
|||||||
cost_info,
|
cost_info,
|
||||||
reservation_snapshot,
|
reservation_snapshot,
|
||||||
)
|
)
|
||||||
cost_data = {**cost_info, "total_usd": 0.0}
|
cost_data = {
|
||||||
|
**cost_info,
|
||||||
|
"total_msats": charged_msats,
|
||||||
|
"charged_msats": charged_msats,
|
||||||
|
"total_usd": 0.0,
|
||||||
|
}
|
||||||
|
if computed_msats != charged_msats:
|
||||||
|
cost_data["computed_msats"] = computed_msats
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"EHBP usage metrics not found in headers or trailers, "
|
"EHBP usage metrics not found in headers or trailers, "
|
||||||
@@ -915,7 +854,7 @@ async def forward_ehbp_request(
|
|||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key.hashed_key[:8] + "...",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
await finalize_ehbp_max_cost_payment(
|
charged_msats = await finalize_ehbp_max_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
@@ -923,11 +862,14 @@ async def forward_ehbp_request(
|
|||||||
reservation_snapshot,
|
reservation_snapshot,
|
||||||
)
|
)
|
||||||
cost_data = {
|
cost_data = {
|
||||||
"total_msats": max_cost_for_model,
|
"total_msats": charged_msats,
|
||||||
|
"charged_msats": charged_msats,
|
||||||
"total_usd": 0.0,
|
"total_usd": 0.0,
|
||||||
"input_tokens": 0,
|
"input_tokens": 0,
|
||||||
"output_tokens": 0,
|
"output_tokens": 0,
|
||||||
}
|
}
|
||||||
|
if charged_msats != max_cost_for_model:
|
||||||
|
cost_data["computed_msats"] = max_cost_for_model
|
||||||
|
|
||||||
# Build the cost_info dict from what adjust_payment_for_tokens returned
|
# Build the cost_info dict from what adjust_payment_for_tokens returned
|
||||||
# or from the max-cost fallback. Fields match CostData/MaxCostData.dict().
|
# or from the max-cost fallback. Fields match CostData/MaxCostData.dict().
|
||||||
@@ -940,6 +882,8 @@ async def forward_ehbp_request(
|
|||||||
"input_msats": cost_data.get("input_msats", 0),
|
"input_msats": cost_data.get("input_msats", 0),
|
||||||
"output_msats": cost_data.get("output_msats", 0),
|
"output_msats": cost_data.get("output_msats", 0),
|
||||||
}
|
}
|
||||||
|
if "computed_msats" in cost_data:
|
||||||
|
cost_info["computed_msats"] = cost_data["computed_msats"]
|
||||||
cost_usd = cost_data.get("total_usd", 0.0)
|
cost_usd = cost_data.get("total_usd", 0.0)
|
||||||
|
|
||||||
# Build response headers, filtering out hop-by-hop headers
|
# Build response headers, filtering out hop-by-hop headers
|
||||||
@@ -1028,7 +972,9 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
target_url = _resolve_ehbp_target_url(
|
target_url = _resolve_ehbp_target_url(
|
||||||
target.url, path, headers, provider_type, profile
|
target.url, path, headers, provider_type, profile
|
||||||
)
|
)
|
||||||
upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile)
|
upstream_headers = _prepare_ehbp_upstream_headers(
|
||||||
|
headers, target.headers, profile
|
||||||
|
)
|
||||||
request_body = await request.body()
|
request_body = await request.body()
|
||||||
|
|
||||||
# Merge query params into the target URL
|
# Merge query params into the target URL
|
||||||
@@ -1076,9 +1022,7 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
usage_source = (
|
usage_source = (
|
||||||
"header"
|
"header"
|
||||||
if usage_header_name
|
if usage_header_name
|
||||||
and any(
|
and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers)
|
||||||
k.lower() == usage_header_name.lower() for k, _ in resp.headers
|
|
||||||
)
|
|
||||||
else ("trailer" if usage_header else "none")
|
else ("trailer" if usage_header else "none")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from unittest.mock import patch
|
|||||||
import pytest
|
import pytest
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from routstr.core.db import ApiKey
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
from routstr.payment.cost_calculation import CostData
|
from routstr.payment.cost_calculation import CostData
|
||||||
|
|
||||||
|
|
||||||
@@ -34,20 +34,30 @@ def _cost_data(total_msats: int) -> CostData:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_overrun_charges_after_reservation_swept(
|
async def test_overrun_with_corrupted_aggregate_releases_without_charging(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Overrun finalize must charge even when the reservation was already released."""
|
"""An overrun whose aggregate reservation was externally zeroed must not
|
||||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
charge: subtracting the reservation would have to clamp, which can erase
|
||||||
|
sibling reservations. The reservation is released and the charge dropped.
|
||||||
|
"""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
deducted_max_cost = 990 # discounted reservation
|
deducted_max_cost = 990 # discounted reservation
|
||||||
actual_token_cost = 1000 # actual cost overruns the reservation
|
actual_token_cost = 1000 # actual cost overruns the reservation
|
||||||
|
|
||||||
# Sweeper has zeroed reserved_balance but left balance untouched.
|
# Something zeroed reserved_balance under an active durable reservation.
|
||||||
key = _make_key(balance=1000, reserved=0)
|
key = _make_key(balance=1000, reserved=0)
|
||||||
|
key_hash = key.hashed_key
|
||||||
integration_session.add(key)
|
integration_session.add(key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
key.reserved_balance = 0
|
key.reserved_balance = 0
|
||||||
integration_session.add(key)
|
integration_session.add(key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
@@ -61,20 +71,24 @@ async def test_overrun_charges_after_reservation_swept(
|
|||||||
"routstr.auth.calculate_cost",
|
"routstr.auth.calculate_cost",
|
||||||
return_value=_cost_data(actual_token_cost),
|
return_value=_cost_data(actual_token_cost),
|
||||||
):
|
):
|
||||||
await adjust_payment_for_tokens(
|
result = await adjust_payment_for_tokens(
|
||||||
key, response_data, integration_session, deducted_max_cost, None, None
|
key, response_data, integration_session, deducted_max_cost, None, None
|
||||||
)
|
)
|
||||||
|
|
||||||
await integration_session.refresh(key)
|
assert result["charged_msats"] == 0
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key_row = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key_row is not None
|
||||||
|
|
||||||
assert key.total_spent == actual_token_cost, (
|
assert key_row.total_spent == 0, "corrupted aggregate must not be charged into"
|
||||||
f"Request was not billed (total_spent={key.total_spent}) — free response bug"
|
assert key_row.balance == 1000
|
||||||
)
|
assert key_row.reserved_balance == 0
|
||||||
assert key.balance == 1000 - actual_token_cost, (
|
|
||||||
f"Balance not charged: {key.balance}"
|
# The corrupt reservation must reach a terminal state — an active leftover
|
||||||
)
|
# would be renewed by its heartbeat forever and poison stale cleanup.
|
||||||
assert key.balance >= 0
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
assert key.reserved_balance == 0
|
assert record is not None and record.status == "released"
|
||||||
|
assert reservation.release_id not in auth._reservation_heartbeats
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -0,0 +1,747 @@
|
|||||||
|
"""Regression tests for the production "negative available balance" 402.
|
||||||
|
|
||||||
|
Invariants protected:
|
||||||
|
* billing and the admin API agree on what "available" means
|
||||||
|
* an in-flight (heartbeaten) reservation is never swept; an abandoned one is
|
||||||
|
released and can never be charged afterwards
|
||||||
|
* a cost overrun spends only its own reservation plus unreserved balance
|
||||||
|
* corrupt reservations are repaired terminally instead of poisoning cleanup
|
||||||
|
* under concurrency, balance / reserved / available never go negative
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from typing import Awaitable, Callable
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient
|
||||||
|
from sqlmodel import col, func, select, update
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
|
from routstr.payment.cost_calculation import CostData
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.integration
|
||||||
|
|
||||||
|
# Realistic sweeper timeout: a renewed (heartbeaten) reservation stays alive,
|
||||||
|
# a reservation backdated past this is released.
|
||||||
|
STALE_TIMEOUT_SECONDS = 300
|
||||||
|
|
||||||
|
# created_at is whole seconds; -1 makes every reservation stale immediately.
|
||||||
|
# Only used in the fuzz test to stress the terminal-release path.
|
||||||
|
SWEEP_EVERYTHING = -1
|
||||||
|
|
||||||
|
|
||||||
|
def _cost_data(total_msats: int) -> CostData:
|
||||||
|
return CostData(
|
||||||
|
base_msats=0,
|
||||||
|
input_msats=total_msats // 2,
|
||||||
|
output_msats=total_msats - total_msats // 2,
|
||||||
|
total_msats=total_msats,
|
||||||
|
total_usd=0.0,
|
||||||
|
input_tokens=100,
|
||||||
|
output_tokens=100,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _response(model: str = "test-model") -> dict:
|
||||||
|
return {
|
||||||
|
"model": model,
|
||||||
|
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _new_key(session: AsyncSession, balance: int) -> str:
|
||||||
|
key_hash = f"test_neg_{uuid.uuid4().hex}"
|
||||||
|
session.add(
|
||||||
|
ApiKey(
|
||||||
|
hashed_key=key_hash,
|
||||||
|
balance=balance,
|
||||||
|
reserved_balance=0,
|
||||||
|
total_spent=0,
|
||||||
|
total_requests=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return key_hash
|
||||||
|
|
||||||
|
|
||||||
|
async def _backdate_reservation(
|
||||||
|
session: AsyncSession, release_id: str, seconds: int
|
||||||
|
) -> None:
|
||||||
|
"""Age a reservation's lease as if it had not been renewed for `seconds`."""
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.id) == release_id)
|
||||||
|
.values(created_at=col(ReservationRelease.created_at) - seconds)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _wait_for(
|
||||||
|
predicate: "Callable[[], Awaitable[bool]]",
|
||||||
|
timeout: float = 10.0,
|
||||||
|
interval: float = 0.1,
|
||||||
|
) -> bool:
|
||||||
|
"""Bounded polling instead of fixed sleeps for background-task effects."""
|
||||||
|
deadline = asyncio.get_event_loop().time() + timeout
|
||||||
|
while True:
|
||||||
|
if await predicate():
|
||||||
|
return True
|
||||||
|
if asyncio.get_event_loop().time() > deadline:
|
||||||
|
return False
|
||||||
|
await asyncio.sleep(interval)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_402_reports_negative_available_while_admin_shows_positive(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
) -> None:
|
||||||
|
"""The production symptom: billing rejects on balance - reserved_balance,
|
||||||
|
so the admin endpoint must expose reserved/available, not just balance."""
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from routstr.auth import _validate_bearer_key_locked
|
||||||
|
from routstr.core.db import set_admin_password
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=263_000)
|
||||||
|
# Leaked reservations slightly exceeding the balance, as in production.
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=267_215)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
with pytest.raises(HTTPException) as exc:
|
||||||
|
await _validate_bearer_key_locked(
|
||||||
|
"sk-" + key_hash, integration_session, min_cost=1
|
||||||
|
)
|
||||||
|
|
||||||
|
assert exc.value.status_code == 402
|
||||||
|
message = exc.value.detail["error"]["message"] # type: ignore[index]
|
||||||
|
assert "-4.215 sats (-4215 msats) available" in message, message
|
||||||
|
|
||||||
|
await set_admin_password(integration_session, "test-admin-pw")
|
||||||
|
login = await integration_client.post(
|
||||||
|
"/admin/api/login", json={"password": "test-admin-pw"}
|
||||||
|
)
|
||||||
|
assert login.status_code == 200
|
||||||
|
token = login.json()["token"]
|
||||||
|
resp = await integration_client.get(
|
||||||
|
"/admin/api/temporary-balances",
|
||||||
|
params={"search": key_hash},
|
||||||
|
headers={"Authorization": f"Bearer {token}"},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 200
|
||||||
|
rows = [row for row in resp.json()["balances"] if row["hashed_key"] == key_hash]
|
||||||
|
assert len(rows) == 1
|
||||||
|
row = rows[0]
|
||||||
|
assert row["balance"] == 263_000
|
||||||
|
assert row["reserved_balance"] == 267_215
|
||||||
|
assert row["available_balance"] == -4_215
|
||||||
|
assert resp.json()["totals"] == {
|
||||||
|
"total_balance": 263_000,
|
||||||
|
"total_reserved_balance": 267_215,
|
||||||
|
"total_available_balance": -4_215,
|
||||||
|
"total_spent": 0,
|
||||||
|
"total_requests": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_abandoned_reservation_is_swept_and_cannot_finalize(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Release is terminal: a swept reservation's late finalizer must not
|
||||||
|
charge, or it could spend funds since reserved by another request."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
cost = 5_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
# Client vanished: the lease is never renewed and ages past the timeout.
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 1
|
||||||
|
)
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
assert released == 1
|
||||||
|
|
||||||
|
# A zombie finalizer shows up afterwards; it must not charge.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.balance == 10_000
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweeper_cannot_release_a_reservation_renewed_after_selection(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A heartbeat landing between the sweeper's select and its transition
|
||||||
|
must win — exercised through the public sweeper entry point."""
|
||||||
|
from routstr.auth import (
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
renew_reservation,
|
||||||
|
)
|
||||||
|
from routstr.core import db as core_db
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=1_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, 1_000, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
|
||||||
|
)
|
||||||
|
|
||||||
|
real_transition = core_db._transition_stale_reservation
|
||||||
|
|
||||||
|
async def renew_then_transition(
|
||||||
|
session: AsyncSession, reservation_id: str, cutoff: int
|
||||||
|
) -> bool:
|
||||||
|
# The sweeper selected this reservation as stale; the heartbeat
|
||||||
|
# renews exactly between that select and the transition.
|
||||||
|
assert await renew_reservation(reservation, session)
|
||||||
|
return await real_transition(session, reservation_id, cutoff)
|
||||||
|
|
||||||
|
with patch.object(core_db, "_transition_stale_reservation", renew_then_transition):
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
|
||||||
|
assert released == 0
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "active"
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.reserved_balance == 1_000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_legacy_cleanup_cannot_erase_a_reservation_committed_after_its_read(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A reservation committing between the legacy sweep's read and its
|
||||||
|
zeroing must survive — exercised through the public sweeper entry point."""
|
||||||
|
from routstr.core import db as core_db
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
stale_reserved_at = int(time.time()) - (STALE_TIMEOUT_SECONDS + 100)
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=500, reserved_at=stale_reserved_at)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
real_release = core_db._release_legacy_aggregate
|
||||||
|
|
||||||
|
async def commit_reservation_then_release(
|
||||||
|
session: AsyncSession,
|
||||||
|
target_key_hash: str,
|
||||||
|
observed_reserved: int,
|
||||||
|
observed_reserved_at: int | None,
|
||||||
|
) -> bool:
|
||||||
|
# The sweeper read the legacy aggregate and found no active durable
|
||||||
|
# owner; a new reservation commits exactly before the zeroing lands.
|
||||||
|
session.add(
|
||||||
|
ReservationRelease(
|
||||||
|
id=uuid.uuid4().hex,
|
||||||
|
key_hash=target_key_hash,
|
||||||
|
billing_key_hash=target_key_hash,
|
||||||
|
reserved_msats=700,
|
||||||
|
status="active",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == target_key_hash)
|
||||||
|
.values(
|
||||||
|
reserved_balance=col(ApiKey.reserved_balance) + 700,
|
||||||
|
reserved_at=int(time.time()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return await real_release(
|
||||||
|
session, target_key_hash, observed_reserved, observed_reserved_at
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
core_db, "_release_legacy_aggregate", commit_reservation_then_release
|
||||||
|
):
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
|
||||||
|
assert released == 0
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.reserved_balance == 1_200, "legacy cleanup erased a live reservation"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweeper_repairs_corrupt_reservation_and_continues_batch(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""One corrupt durable reservation (aggregate no longer holds its msats)
|
||||||
|
must be terminalized without aggregate subtraction, and must not stop the
|
||||||
|
rest of the batch from being released normally."""
|
||||||
|
from routstr.auth import get_reservation_snapshot, pay_for_request
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
corrupt_hash = await _new_key(integration_session, balance=cost)
|
||||||
|
healthy_hash = await _new_key(integration_session, balance=cost)
|
||||||
|
|
||||||
|
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
|
||||||
|
assert corrupt_key is not None
|
||||||
|
await pay_for_request(corrupt_key, cost, integration_session)
|
||||||
|
corrupt_reservation = await get_reservation_snapshot(
|
||||||
|
corrupt_key, integration_session
|
||||||
|
)
|
||||||
|
|
||||||
|
healthy_key = await integration_session.get(ApiKey, healthy_hash)
|
||||||
|
assert healthy_key is not None
|
||||||
|
await pay_for_request(healthy_key, cost, integration_session)
|
||||||
|
healthy_reservation = await get_reservation_snapshot(
|
||||||
|
healthy_key, integration_session
|
||||||
|
)
|
||||||
|
|
||||||
|
# Corrupt the first key: its aggregate no longer holds the reservation.
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == corrupt_hash)
|
||||||
|
.values(reserved_balance=0)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
for reservation in (corrupt_reservation, healthy_reservation):
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
|
||||||
|
)
|
||||||
|
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
assert released == 2, "a corrupt record must not abort the sweep batch"
|
||||||
|
|
||||||
|
for reservation in (corrupt_reservation, healthy_reservation):
|
||||||
|
record = await integration_session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
healthy_key = await integration_session.get(ApiKey, healthy_hash)
|
||||||
|
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
|
||||||
|
assert healthy_key is not None and corrupt_key is not None
|
||||||
|
assert healthy_key.reserved_balance == 0
|
||||||
|
assert corrupt_key.reserved_balance == 0
|
||||||
|
assert corrupt_key.balance == cost, "repair must not touch balances"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_heartbeat_survives_a_rolled_back_charge_attempt(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""Claiming a reservation must not stop its heartbeat: a rollback restores
|
||||||
|
the active reservation, which then still needs lease renewal."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
_claim_reservation_for_charge,
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=2_000)
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# A charge attempt claims the reservation, then its transaction
|
||||||
|
# fails and rolls back.
|
||||||
|
async with create_session() as session:
|
||||||
|
assert await _claim_reservation_for_charge(reservation, session)
|
||||||
|
await session.rollback()
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "active"
|
||||||
|
await _backdate_reservation(
|
||||||
|
session, reservation.release_id, timeout * 10
|
||||||
|
)
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None
|
||||||
|
backdated_lease = record.created_at
|
||||||
|
|
||||||
|
async def lease_renewed() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
return record is not None and record.created_at > backdated_lease
|
||||||
|
|
||||||
|
assert await _wait_for(lease_renewed), (
|
||||||
|
"heartbeat did not survive the rolled-back charge attempt"
|
||||||
|
)
|
||||||
|
|
||||||
|
# The restored reservation still finalizes normally.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "charged"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_heartbeat_dies_with_its_request_so_sweeper_can_recover(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""A request that vanishes without finalizing must not renew forever —
|
||||||
|
its heartbeat stops with the owning task and the sweeper reclaims the
|
||||||
|
funds."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import get_reservation_snapshot, pay_for_request
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=cost)
|
||||||
|
|
||||||
|
holder: dict = {}
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
|
||||||
|
async def doomed_request() -> None:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
holder["reservation"] = await get_reservation_snapshot(key, session)
|
||||||
|
# ...request control dies here, no finalize and no release.
|
||||||
|
|
||||||
|
await asyncio.create_task(doomed_request())
|
||||||
|
release_id = holder["reservation"].release_id
|
||||||
|
|
||||||
|
try:
|
||||||
|
# While the heartbeat is still winding down it may renew once
|
||||||
|
# more; keep backdating until the sweeper wins, which it must as
|
||||||
|
# soon as the dead owner is noticed.
|
||||||
|
async def sweeper_recovered() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
await _backdate_reservation(session, release_id, timeout * 10)
|
||||||
|
return await release_stale_reservations(session, timeout) == 1
|
||||||
|
|
||||||
|
assert await _wait_for(sweeper_recovered), (
|
||||||
|
"sweeper never recovered the abandoned reservation"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(release_id)
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, release_id)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert key.total_spent == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reservation_heartbeat_covers_the_whole_request_lifecycle(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""pay_for_request starts the heartbeat, finalization stops it, and a
|
||||||
|
backdated lease is renewed in the background without any manual call."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=2 * cost)
|
||||||
|
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Only a background renewal can keep this alive now.
|
||||||
|
async with create_session() as session:
|
||||||
|
await _backdate_reservation(
|
||||||
|
session, reservation.release_id, timeout * 10
|
||||||
|
)
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None
|
||||||
|
backdated_lease = record.created_at
|
||||||
|
|
||||||
|
async def lease_renewed() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
return record is not None and record.created_at > backdated_lease
|
||||||
|
|
||||||
|
assert await _wait_for(lease_renewed), "heartbeat never renewed the lease"
|
||||||
|
async with create_session() as session:
|
||||||
|
assert await release_stale_reservations(session, timeout) == 0
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "charged"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_overrun_cannot_spend_a_concurrent_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""An overrun is capped to its own reservation plus unreserved balance;
|
||||||
|
the sibling's reserved funds stay untouched and available stays >= 0."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved_each = 100
|
||||||
|
overrun_cost = 150 # A's real token cost exceeds its reservation
|
||||||
|
|
||||||
|
# Balance covers exactly two reservations; nothing free on top.
|
||||||
|
key_hash = await _new_key(integration_session, balance=2 * reserved_each)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved_each, integration_session)
|
||||||
|
reservation_a = await get_reservation_snapshot(key, integration_session)
|
||||||
|
await pay_for_request(key, reserved_each, integration_session)
|
||||||
|
await get_reservation_snapshot(key, integration_session) # B stays in flight
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.reserved_balance == 2 * reserved_each
|
||||||
|
|
||||||
|
# Only A finalizes; B is still streaming and its funds must stay reserved.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(overrun_cost)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved_each,
|
||||||
|
reservation_snapshot=reservation_a,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == reserved_each
|
||||||
|
assert result["total_msats"] == overrun_cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == reserved_each
|
||||||
|
assert key.reserved_balance == reserved_each
|
||||||
|
assert key.balance == reserved_each
|
||||||
|
assert key.total_balance >= 0, (
|
||||||
|
f"available balance went negative: balance={key.balance} "
|
||||||
|
f"reserved={key.reserved_balance} -> {key.total_balance} msats; "
|
||||||
|
"the overrun charge consumed the still-reserved funds of request B"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_concurrent_requests_with_sweeper_keep_balance_invariants(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""Fuzz: concurrent requests against an everything-is-stale sweeper.
|
||||||
|
Some requests legitimately finish uncharged (release is terminal), but
|
||||||
|
balances must never go negative and no reservation may stay active."""
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
rng = random.Random(1337)
|
||||||
|
starting_balance = 200_000
|
||||||
|
n_requests = 24
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=starting_balance)
|
||||||
|
|
||||||
|
completed_costs: list[int] = []
|
||||||
|
rejected_requests = 0
|
||||||
|
|
||||||
|
async def one_request(index: int) -> None:
|
||||||
|
nonlocal rejected_requests
|
||||||
|
reserved = rng.randrange(1_000, 4_000)
|
||||||
|
actual = max(1, int(reserved * rng.uniform(0.5, 1.1)))
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, reserved, session)
|
||||||
|
except HTTPException as exc:
|
||||||
|
# A depleted balance is the only legitimate rejection.
|
||||||
|
assert exc.status_code == 402, exc.detail
|
||||||
|
rejected_requests += 1
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
except RuntimeError:
|
||||||
|
# The everything-is-stale sweeper can release the reservation
|
||||||
|
# before the stream even starts; the request aborts uncharged.
|
||||||
|
return
|
||||||
|
|
||||||
|
await asyncio.sleep(rng.uniform(0, 0.02)) # the "stream"
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(str(actual)),
|
||||||
|
session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
completed_costs.append(actual)
|
||||||
|
|
||||||
|
sweeping = True
|
||||||
|
|
||||||
|
async def sweeper() -> None:
|
||||||
|
while sweeping:
|
||||||
|
async with create_session() as session:
|
||||||
|
await release_stale_reservations(session, SWEEP_EVERYTHING)
|
||||||
|
await asyncio.sleep(0.002)
|
||||||
|
|
||||||
|
sweep_task = asyncio.create_task(sweeper())
|
||||||
|
try:
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
side_effect=lambda response_data, *a, **k: _cost_data(
|
||||||
|
int(response_data["model"])
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await asyncio.gather(*(one_request(i) for i in range(n_requests)))
|
||||||
|
finally:
|
||||||
|
sweeping = False
|
||||||
|
await sweep_task
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
leftover_active = (
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
select(func.count())
|
||||||
|
.select_from(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.status) == "active")
|
||||||
|
)
|
||||||
|
).one()
|
||||||
|
|
||||||
|
assert completed_costs or rejected_requests, "no request made any progress"
|
||||||
|
assert key.balance >= 0, f"balance went negative: {key.balance}"
|
||||||
|
assert key.reserved_balance >= 0, (
|
||||||
|
f"reserved_balance went negative: {key.reserved_balance}"
|
||||||
|
)
|
||||||
|
assert key.total_balance >= 0, (
|
||||||
|
f"available balance negative: balance={key.balance} "
|
||||||
|
f"reserved={key.reserved_balance} (this is the production symptom)"
|
||||||
|
)
|
||||||
|
assert key.total_spent <= starting_balance, (
|
||||||
|
f"spent {key.total_spent} of a {starting_balance} balance"
|
||||||
|
)
|
||||||
|
# Requests whose reservation was swept mid-flight finish uncharged, so the
|
||||||
|
# charged total can only be at most the sum of completed request costs.
|
||||||
|
max_expected_spend = sum(completed_costs)
|
||||||
|
assert key.total_spent <= max_expected_spend, (
|
||||||
|
f"charged {key.total_spent} msats but completed requests only cost "
|
||||||
|
f"{max_expected_spend}"
|
||||||
|
)
|
||||||
|
assert leftover_active == 0, (
|
||||||
|
f"{leftover_active} reservations still active after all requests finished"
|
||||||
|
)
|
||||||
@@ -0,0 +1,549 @@
|
|||||||
|
"""Invariant coverage for the reserve → charge → release money path.
|
||||||
|
|
||||||
|
Every finalization branch of ``adjust_payment_for_tokens`` must respect the
|
||||||
|
same accounting rules: a completed request is charged exactly once, its
|
||||||
|
reported ``charged_msats`` matches the actual debit, it never spends more than
|
||||||
|
its own reservation leaves available, and child keys spend their parent's
|
||||||
|
balance without raiding sibling reservations.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi import HTTPException
|
||||||
|
from sqlmodel import col, select, update
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
|
from routstr.payment.cost_calculation import (
|
||||||
|
CostData,
|
||||||
|
CostDataError,
|
||||||
|
MaxCostData,
|
||||||
|
)
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.integration
|
||||||
|
|
||||||
|
|
||||||
|
def _cost_data(total_msats: int, cls: type[CostData] = CostData) -> CostData:
|
||||||
|
return cls(
|
||||||
|
base_msats=0,
|
||||||
|
input_msats=total_msats // 2,
|
||||||
|
output_msats=total_msats - total_msats // 2,
|
||||||
|
total_msats=total_msats,
|
||||||
|
total_usd=0.0,
|
||||||
|
input_tokens=100,
|
||||||
|
output_tokens=100,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _response() -> dict:
|
||||||
|
return {
|
||||||
|
"model": "test-model",
|
||||||
|
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _new_key(
|
||||||
|
session: AsyncSession,
|
||||||
|
balance: int,
|
||||||
|
*,
|
||||||
|
parent_key_hash: str | None = None,
|
||||||
|
balance_limit: int | None = None,
|
||||||
|
) -> str:
|
||||||
|
key_hash = f"test_inv_{uuid.uuid4().hex}"
|
||||||
|
session.add(
|
||||||
|
ApiKey(
|
||||||
|
hashed_key=key_hash,
|
||||||
|
balance=balance,
|
||||||
|
reserved_balance=0,
|
||||||
|
total_spent=0,
|
||||||
|
total_requests=0,
|
||||||
|
parent_key_hash=parent_key_hash,
|
||||||
|
balance_limit=balance_limit,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return key_hash
|
||||||
|
|
||||||
|
|
||||||
|
async def _active_reservations(session: AsyncSession) -> int:
|
||||||
|
rows = await session.exec(
|
||||||
|
select(ReservationRelease).where(col(ReservationRelease.status) == "active")
|
||||||
|
)
|
||||||
|
return len(rows.all())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_corrupt_revert_still_decrements_request_count(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
revert_pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, 3_000, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=0)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
assert await revert_pay_for_request(
|
||||||
|
key,
|
||||||
|
integration_session,
|
||||||
|
3_000,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
updated = await integration_session.get(ApiKey, key_hash)
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert updated is not None
|
||||||
|
assert updated.total_requests == 0
|
||||||
|
assert updated.reserved_balance == 0
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_exact_cost_branch_charges_the_reservation_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert key.total_requests == 1
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_max_cost_branch_charges_the_reservation_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""No token pricing configured -> flat MaxCostData charge of the reservation."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 2_500
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
return_value=_cost_data(cost, cls=MaxCostData),
|
||||||
|
):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_underrun_branch_refunds_the_unused_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 5_000
|
||||||
|
actual = 1_200
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == actual
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == actual, "user must pay the real cost, not the reservation"
|
||||||
|
assert key.balance == 10_000 - actual
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_overrun_branch_charges_full_cost_when_balance_is_free(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 1_000
|
||||||
|
actual = 1_400
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == actual
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == actual
|
||||||
|
assert key.balance == 10_000 - actual
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_zero_cost_response_is_free_and_releases_the_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""An empty/unusable upstream response must cost the user nothing."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 4_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(0)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
{"model": "test-model"},
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cost_error_releases_the_reservation_without_charging(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 4_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
return_value=CostDataError(message="no pricing", code="pricing_error"),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc:
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
assert exc.value.status_code == 400
|
||||||
|
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.balance == 10_000, "a pricing failure must not charge the user"
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.reserved_balance == 0, "funds must not stay locked after a 400"
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_repeated_finalization_charges_only_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A retried finalizer (proxy retry, duplicate stream end) must not double-bill."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
for _ in range(3):
|
||||||
|
# A declined re-charge rolls its session back, so re-load the key
|
||||||
|
# the way a fresh request would instead of reusing a stale instance.
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
results.append(
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Only the first finalization debits; duplicates report a zero charge.
|
||||||
|
assert [r["charged_msats"] for r in results] == [cost, 0, 0]
|
||||||
|
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost, f"charged {key.total_spent} for one request"
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_concurrent_duplicate_finalization_charges_only_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=10_000)
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
async def finalize() -> None:
|
||||||
|
async with create_session() as session:
|
||||||
|
fresh = await session.get(ApiKey, key_hash)
|
||||||
|
assert fresh is not None
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
fresh,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await asyncio.gather(finalize(), finalize(), finalize())
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_release_after_charge_does_not_credit_the_user_back(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A late cleanup path must not turn a charged request into a free one."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
release_reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
released = await release_reservation(reservation, integration_session, cost)
|
||||||
|
assert released is False, "a charged reservation must not be releasable"
|
||||||
|
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_child_request_spends_parent_balance_and_records_child_ledger(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
parent_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
child_hash = await _new_key(
|
||||||
|
integration_session, balance=0, parent_key_hash=parent_hash
|
||||||
|
)
|
||||||
|
child = await integration_session.get(ApiKey, child_hash)
|
||||||
|
assert child is not None
|
||||||
|
|
||||||
|
await pay_for_request(child, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(child, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
child,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
parent = await integration_session.get(ApiKey, parent_hash)
|
||||||
|
child = await integration_session.get(ApiKey, child_hash)
|
||||||
|
assert parent is not None and child is not None
|
||||||
|
assert parent.balance == 10_000 - cost
|
||||||
|
assert parent.reserved_balance == 0
|
||||||
|
assert parent.total_balance >= 0
|
||||||
|
assert child.reserved_balance == 0, "child reservation must be released too"
|
||||||
|
assert child.total_balance >= 0
|
||||||
|
assert child.total_spent == cost, "child ledger must record the spend"
|
||||||
|
assert parent.total_spent == cost
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_child_overrun_does_not_raid_a_sibling_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Same overrun defect as the parent case, reached through a child key."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved_each = 100
|
||||||
|
overrun = 150
|
||||||
|
parent_hash = await _new_key(integration_session, balance=2 * reserved_each)
|
||||||
|
child_a = await _new_key(
|
||||||
|
integration_session, balance=0, parent_key_hash=parent_hash
|
||||||
|
)
|
||||||
|
child_b = await _new_key(
|
||||||
|
integration_session, balance=0, parent_key_hash=parent_hash
|
||||||
|
)
|
||||||
|
|
||||||
|
key_a = await integration_session.get(ApiKey, child_a)
|
||||||
|
key_b = await integration_session.get(ApiKey, child_b)
|
||||||
|
assert key_a is not None and key_b is not None
|
||||||
|
|
||||||
|
await pay_for_request(key_a, reserved_each, integration_session)
|
||||||
|
reservation_a = await get_reservation_snapshot(key_a, integration_session)
|
||||||
|
await pay_for_request(key_b, reserved_each, integration_session)
|
||||||
|
await get_reservation_snapshot(key_b, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(overrun)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key_a,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved_each,
|
||||||
|
reservation_snapshot=reservation_a,
|
||||||
|
)
|
||||||
|
|
||||||
|
parent = await integration_session.get(ApiKey, parent_hash)
|
||||||
|
assert parent is not None
|
||||||
|
assert parent.total_balance >= 0, (
|
||||||
|
f"child A's overrun ate child B's reservation: balance={parent.balance} "
|
||||||
|
f"reserved={parent.reserved_balance}"
|
||||||
|
)
|
||||||
@@ -133,14 +133,11 @@ async def test_reserved_balance_with_successful_requests(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_revert_with_zero_reserved_balance_is_noop(
|
async def test_revert_with_zero_reserved_balance_repairs_terminally(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test that revert_pay_for_request is a no-op when reserved_balance is 0.
|
"""Reverting after the aggregate was already zeroed must not drive it
|
||||||
|
negative: the corrupt durable reservation is released without subtraction."""
|
||||||
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 pay_for_request, 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]}"
|
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
|
||||||
@@ -153,21 +150,66 @@ async def test_revert_with_zero_reserved_balance_is_noop(
|
|||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
await pay_for_request(test_key, 100, integration_session)
|
await pay_for_request(test_key, 100, integration_session)
|
||||||
test_key.reserved_balance = 0
|
test_key.reserved_balance = 0
|
||||||
|
test_key.total_requests = 0
|
||||||
integration_session.add(test_key)
|
integration_session.add(test_key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
|
|
||||||
# A stale cleanup already released the aggregate reservation.
|
# A stale cleanup already released the aggregate reservation. The revert
|
||||||
|
# terminalizes the durable row (repair) without driving the aggregate
|
||||||
|
# negative.
|
||||||
result = await revert_pay_for_request(test_key, integration_session, 100)
|
result = await revert_pay_for_request(test_key, integration_session, 100)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
integration_session.expunge_all()
|
||||||
|
test_key = await integration_session.get(ApiKey, unique_key)
|
||||||
|
assert test_key is not None
|
||||||
|
|
||||||
assert result is False, "Revert should return False when reservation already released"
|
assert result is True, "Revert must terminalize the corrupt reservation"
|
||||||
assert test_key.reserved_balance == 0, (
|
assert test_key.reserved_balance == 0, (
|
||||||
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
|
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
|
||||||
)
|
)
|
||||||
assert test_key.total_requests == 1, (
|
assert test_key.total_requests == 0
|
||||||
f"Total requests should remain 1, got: {test_key.total_requests}"
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_child_corrupt_revert_clamps_zero_request_counts(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
revert_pay_for_request,
|
||||||
)
|
)
|
||||||
|
from routstr.core.db import ReservationRelease
|
||||||
|
|
||||||
|
suffix = uuid.uuid4().hex[:8]
|
||||||
|
parent = ApiKey(hashed_key=f"repair-parent-{suffix}", balance=5_000)
|
||||||
|
child = ApiKey(
|
||||||
|
hashed_key=f"repair-child-{suffix}",
|
||||||
|
parent_key_hash=parent.hashed_key,
|
||||||
|
)
|
||||||
|
integration_session.add(parent)
|
||||||
|
integration_session.add(child)
|
||||||
|
await integration_session.commit()
|
||||||
|
await pay_for_request(child, 500, integration_session)
|
||||||
|
snapshot = await get_reservation_snapshot(child, integration_session)
|
||||||
|
|
||||||
|
parent.total_requests = 0
|
||||||
|
child.total_requests = 0
|
||||||
|
child.reserved_balance = 0
|
||||||
|
integration_session.add(parent)
|
||||||
|
integration_session.add(child)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
assert await revert_pay_for_request(child, integration_session, 500, snapshot)
|
||||||
|
|
||||||
|
integration_session.expunge_all()
|
||||||
|
parent = await integration_session.get(ApiKey, snapshot.billing_key_hash)
|
||||||
|
child = await integration_session.get(ApiKey, snapshot.key_hash)
|
||||||
|
release = await integration_session.get(ReservationRelease, snapshot.release_id)
|
||||||
|
assert parent is not None and child is not None
|
||||||
|
assert (parent.total_requests, child.total_requests) == (0, 0)
|
||||||
|
assert (parent.reserved_balance, child.reserved_balance) == (500, 0)
|
||||||
|
assert release is not None and release.status == "released"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -203,10 +245,11 @@ async def test_revert_with_sufficient_reserved_balance_succeeds(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_revert_partial_reserved_balance_is_noop(
|
async def test_revert_partial_reserved_balance_repairs_terminally(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test that reverting more than the current reserved_balance is a no-op."""
|
"""Reverting more than the aggregate holds must not clamp or go negative:
|
||||||
|
the corrupt durable reservation is released without subtraction."""
|
||||||
from routstr.auth import pay_for_request, 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]}"
|
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
|
||||||
@@ -223,18 +266,19 @@ async def test_revert_partial_reserved_balance_is_noop(
|
|||||||
integration_session.add(test_key)
|
integration_session.add(test_key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
|
|
||||||
# Try to revert 500 when only 50 is reserved — should be no-op
|
# Reverting 500 when only 50 is reserved cannot subtract; the corrupt
|
||||||
|
# reservation is terminalized and the aggregate left untouched.
|
||||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
integration_session.expunge_all()
|
||||||
|
test_key = await integration_session.get(ApiKey, unique_key)
|
||||||
|
assert test_key is not None
|
||||||
|
|
||||||
assert result is False, "Revert should fail when cost > reserved_balance"
|
assert result is True, "Revert must terminalize the corrupt reservation"
|
||||||
assert test_key.reserved_balance == 50, (
|
assert test_key.reserved_balance == 50, (
|
||||||
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
|
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
|
||||||
)
|
)
|
||||||
assert test_key.total_requests == 1, (
|
assert test_key.total_requests == 0
|
||||||
f"Total requests should stay at 1, got: {test_key.total_requests}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -265,9 +309,7 @@ async def test_double_revert_prevented(
|
|||||||
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
||||||
|
|
||||||
# First revert — should succeed
|
# First revert — should succeed
|
||||||
result1 = await revert_pay_for_request(
|
result1 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
|
|
||||||
assert result1 is True
|
assert result1 is True
|
||||||
@@ -275,9 +317,7 @@ async def test_double_revert_prevented(
|
|||||||
assert test_key.total_requests == 4
|
assert test_key.total_requests == 4
|
||||||
|
|
||||||
# Second revert of the same amount — should be no-op
|
# Second revert of the same amount — should be no-op
|
||||||
result2 = await revert_pay_for_request(
|
result2 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
|
|
||||||
assert result2 is False, "Second revert should be a no-op"
|
assert result2 is False, "Second revert should be a no-op"
|
||||||
@@ -319,9 +359,7 @@ async def test_sequential_reverts_never_go_negative(
|
|||||||
# Run 5 sequential reverts for the same 500 reservation
|
# Run 5 sequential reverts for the same 500 reservation
|
||||||
results = []
|
results = []
|
||||||
for _ in range(5):
|
for _ in range(5):
|
||||||
r = await revert_pay_for_request(
|
r = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
results.append(r)
|
results.append(r)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ async def _add_key(
|
|||||||
hashed_key: str,
|
hashed_key: str,
|
||||||
*,
|
*,
|
||||||
balance: int = 0,
|
balance: int = 0,
|
||||||
|
reserved_balance: int = 0,
|
||||||
total_spent: int = 0,
|
total_spent: int = 0,
|
||||||
total_requests: int = 0,
|
total_requests: int = 0,
|
||||||
created_at: int | None = None,
|
created_at: int | None = None,
|
||||||
@@ -31,6 +32,7 @@ async def _add_key(
|
|||||||
key = ApiKey(
|
key = ApiKey(
|
||||||
hashed_key=hashed_key,
|
hashed_key=hashed_key,
|
||||||
balance=balance,
|
balance=balance,
|
||||||
|
reserved_balance=reserved_balance,
|
||||||
total_spent=total_spent,
|
total_spent=total_spent,
|
||||||
total_requests=total_requests,
|
total_requests=total_requests,
|
||||||
parent_key_hash=parent_key_hash,
|
parent_key_hash=parent_key_hash,
|
||||||
@@ -157,6 +159,41 @@ async def test_temporary_balances_totals_exclude_child_balance(
|
|||||||
assert totals["total_requests"] == 10
|
assert totals["total_requests"] == 10
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_child_reservations_are_not_double_counted(
|
||||||
|
integration_client: httpx.AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
await _add_key(
|
||||||
|
integration_session,
|
||||||
|
"parent",
|
||||||
|
balance=5_000,
|
||||||
|
reserved_balance=700,
|
||||||
|
created_at=1_000,
|
||||||
|
)
|
||||||
|
await _add_key(
|
||||||
|
integration_session,
|
||||||
|
"child",
|
||||||
|
reserved_balance=700,
|
||||||
|
parent_key_hash="parent",
|
||||||
|
created_at=1_001,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await integration_client.get(
|
||||||
|
"/admin/api/temporary-balances", headers=_admin_headers()
|
||||||
|
)
|
||||||
|
|
||||||
|
body = response.json()
|
||||||
|
rows = {row["hashed_key"]: row for row in body["balances"]}
|
||||||
|
assert rows["parent"]["available_balance"] == 4_300
|
||||||
|
assert rows["parent"]["reserved_balance"] == 700
|
||||||
|
assert rows["child"]["available_balance"] is None
|
||||||
|
assert rows["child"]["reserved_balance"] == 700
|
||||||
|
assert body["totals"]["total_reserved_balance"] == 700
|
||||||
|
assert body["totals"]["total_available_balance"] == 4_300
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_temporary_balances_search_filters_total_and_totals(
|
async def test_temporary_balances_search_filters_total_and_totals(
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ def _assert_cost_contract(response: Any) -> None:
|
|||||||
"input_msats": 1_200,
|
"input_msats": 1_200,
|
||||||
"output_msats": 300,
|
"output_msats": 300,
|
||||||
"total_msats": 1_500,
|
"total_msats": 1_500,
|
||||||
|
"charged_msats": 1_500,
|
||||||
"total_usd": 0.0001,
|
"total_usd": 0.0001,
|
||||||
"cache_read_input_tokens": 8,
|
"cache_read_input_tokens": 8,
|
||||||
"cache_creation_input_tokens": 2,
|
"cache_creation_input_tokens": 2,
|
||||||
@@ -69,6 +70,36 @@ def _assert_cost_contract(response: Any) -> None:
|
|||||||
assert response.headers["X-Routstr-Output-Cost-Msats"] == "300"
|
assert response.headers["X-Routstr-Output-Cost-Msats"] == "300"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_duplicate_finalization_publishes_zero_debit_and_computed_usage() -> None:
|
||||||
|
provider = _provider()
|
||||||
|
duplicate_cost = {**COST_DATA, "charged_msats": 0}
|
||||||
|
with patch(
|
||||||
|
"routstr.upstream.base.adjust_payment_for_tokens",
|
||||||
|
new=AsyncMock(return_value=duplicate_cost),
|
||||||
|
):
|
||||||
|
response = await provider.handle_non_streaming_chat_completion(
|
||||||
|
_upstream_response(
|
||||||
|
{
|
||||||
|
"model": "test-model",
|
||||||
|
"usage": {"prompt_tokens": 10, "completion_tokens": 3},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
_key(),
|
||||||
|
_session(),
|
||||||
|
deducted_max_cost=10_000,
|
||||||
|
)
|
||||||
|
|
||||||
|
body = json.loads(response.body)
|
||||||
|
assert response.headers["X-Routstr-Cost-Msats"] == "0"
|
||||||
|
assert response.headers["X-Routstr-Computed-Cost-Msats"] == "1500"
|
||||||
|
assert body["usage"]["cost"]["total_msats"] == 0
|
||||||
|
assert body["usage"]["cost"]["charged_msats"] == 0
|
||||||
|
assert body["usage"]["cost"]["computed_msats"] == 1_500
|
||||||
|
assert body["cost"]["total_msats"] == 0
|
||||||
|
assert body["cost"]["computed_msats"] == 1_500
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_balance_chat_completion_uses_shared_cost_contract() -> None:
|
async def test_balance_chat_completion_uses_shared_cost_contract() -> None:
|
||||||
provider = _provider()
|
provider = _provider()
|
||||||
|
|||||||
@@ -6,12 +6,14 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||||
from sqlalchemy.pool import StaticPool
|
from sqlalchemy.pool import StaticPool
|
||||||
from sqlmodel import SQLModel, select
|
from sqlmodel import SQLModel, col, select, update
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
import routstr.auth as auth_module
|
||||||
from routstr.auth import get_reservation_snapshot, pay_for_request
|
from routstr.auth import get_reservation_snapshot, pay_for_request
|
||||||
from routstr.core.db import ApiKey, ReservationRelease
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
from routstr.upstream.ehbp import (
|
from routstr.upstream.ehbp import (
|
||||||
|
_inject_cost_response_headers,
|
||||||
finalize_ehbp_actual_cost_payment,
|
finalize_ehbp_actual_cost_payment,
|
||||||
finalize_ehbp_max_cost_payment,
|
finalize_ehbp_max_cost_payment,
|
||||||
)
|
)
|
||||||
@@ -26,7 +28,9 @@ def _make_engine() -> AsyncEngine:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSession, None]:
|
async def session(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> AsyncGenerator[AsyncSession, None]:
|
||||||
monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0)
|
monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0)
|
||||||
engine = _make_engine()
|
engine = _make_engine()
|
||||||
async with engine.begin() as conn:
|
async with engine.begin() as conn:
|
||||||
@@ -35,6 +39,8 @@ async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSessio
|
|||||||
try:
|
try:
|
||||||
yield db_session
|
yield db_session
|
||||||
finally:
|
finally:
|
||||||
|
for release_id in list(auth_module._reservation_heartbeats):
|
||||||
|
await auth_module._stop_reservation_heartbeat(release_id)
|
||||||
await db_session.close()
|
await db_session.close()
|
||||||
await engine.dispose()
|
await engine.dispose()
|
||||||
|
|
||||||
@@ -54,9 +60,7 @@ def _fail_nth_api_key_update(
|
|||||||
original_exec = session.exec
|
original_exec = session.exec
|
||||||
api_key_updates = 0
|
api_key_updates = 0
|
||||||
|
|
||||||
async def exec_with_failure(
|
async def exec_with_failure(statement: Any, *args: Any, **kwargs: Any) -> Any:
|
||||||
statement: Any, *args: Any, **kwargs: Any
|
|
||||||
) -> Any:
|
|
||||||
nonlocal api_key_updates
|
nonlocal api_key_updates
|
||||||
table = getattr(statement, "table", None)
|
table = getattr(statement, "table", None)
|
||||||
if getattr(table, "name", None) == "api_keys":
|
if getattr(table, "name", None) == "api_keys":
|
||||||
@@ -78,7 +82,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve
|
|||||||
await pay_for_request(key, 3_000, session)
|
await pay_for_request(key, 3_000, session)
|
||||||
reservation = await get_reservation_snapshot(key, session)
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
await finalize_ehbp_actual_cost_payment(
|
charged = await finalize_ehbp_actual_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
reserved_cost_for_model=3_000,
|
reserved_cost_for_model=3_000,
|
||||||
@@ -93,6 +97,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve
|
|||||||
reservation_snapshot=reservation,
|
reservation_snapshot=reservation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
assert charged == 1_200
|
||||||
updated = await _api_key(session, "ehbp-actual")
|
updated = await _api_key(session, "ehbp-actual")
|
||||||
assert updated is not None
|
assert updated is not None
|
||||||
assert updated.balance == 8_800
|
assert updated.balance == 8_800
|
||||||
@@ -106,16 +111,14 @@ async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
|
|||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
parent = ApiKey(hashed_key="ehbp-parent", balance=10_000)
|
parent = ApiKey(hashed_key="ehbp-parent", balance=10_000)
|
||||||
child = ApiKey(
|
child = ApiKey(hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent")
|
||||||
hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent"
|
|
||||||
)
|
|
||||||
session.add(parent)
|
session.add(parent)
|
||||||
session.add(child)
|
session.add(child)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await pay_for_request(child, 3_000, session)
|
await pay_for_request(child, 3_000, session)
|
||||||
reservation = await get_reservation_snapshot(child, session)
|
reservation = await get_reservation_snapshot(child, session)
|
||||||
|
|
||||||
await finalize_ehbp_max_cost_payment(
|
charged = await finalize_ehbp_max_cost_payment(
|
||||||
child,
|
child,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model=3_000,
|
max_cost_for_model=3_000,
|
||||||
@@ -123,6 +126,7 @@ async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
|
|||||||
reservation_snapshot=reservation,
|
reservation_snapshot=reservation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
assert charged == 3_000
|
||||||
updated_parent = await _api_key(session, "ehbp-parent")
|
updated_parent = await _api_key(session, "ehbp-parent")
|
||||||
updated_child = await _api_key(session, "ehbp-child")
|
updated_child = await _api_key(session, "ehbp-child")
|
||||||
assert updated_parent is not None
|
assert updated_parent is not None
|
||||||
@@ -151,7 +155,7 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche
|
|||||||
rollback_spy = AsyncMock(wraps=session.rollback)
|
rollback_spy = AsyncMock(wraps=session.rollback)
|
||||||
monkeypatch.setattr(session, "rollback", rollback_spy)
|
monkeypatch.setattr(session, "rollback", rollback_spy)
|
||||||
|
|
||||||
await finalize_ehbp_actual_cost_payment(
|
charged = await finalize_ehbp_actual_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
reserved_cost_for_model=3_000,
|
reserved_cost_for_model=3_000,
|
||||||
@@ -160,15 +164,17 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche
|
|||||||
reservation_snapshot=reservation,
|
reservation_snapshot=reservation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
assert charged == 0
|
||||||
rollback_spy.assert_awaited_once()
|
rollback_spy.assert_awaited_once()
|
||||||
updated = await _api_key(session, "ehbp-missing-parent")
|
updated = await _api_key(session, "ehbp-missing-parent")
|
||||||
assert updated is not None
|
assert updated is not None
|
||||||
assert updated.balance == 10_000
|
assert updated.balance == 10_000
|
||||||
assert updated.reserved_balance == 3_000
|
assert updated.reserved_balance == 0
|
||||||
assert updated.total_spent == 0
|
assert updated.total_spent == 0
|
||||||
release = await session.get(ReservationRelease, reservation.release_id)
|
release = await session.get(ReservationRelease, reservation.release_id)
|
||||||
assert release is not None
|
assert release is not None
|
||||||
assert release.status == "active"
|
assert release.status == "released"
|
||||||
|
assert reservation.release_id not in auth_module._reservation_heartbeats
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -189,7 +195,7 @@ async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_mat
|
|||||||
reservation = await get_reservation_snapshot(child, session)
|
reservation = await get_reservation_snapshot(child, session)
|
||||||
_fail_nth_api_key_update(session, monkeypatch, target_update=2)
|
_fail_nth_api_key_update(session, monkeypatch, target_update=2)
|
||||||
|
|
||||||
await finalize_ehbp_max_cost_payment(
|
charged = await finalize_ehbp_max_cost_payment(
|
||||||
child,
|
child,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model=3_000,
|
max_cost_for_model=3_000,
|
||||||
@@ -197,12 +203,84 @@ async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_mat
|
|||||||
reservation_snapshot=reservation,
|
reservation_snapshot=reservation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
assert charged == 0
|
||||||
updated_parent = await _api_key(session, "ehbp-rollback-parent")
|
updated_parent = await _api_key(session, "ehbp-rollback-parent")
|
||||||
assert updated_parent is not None
|
assert updated_parent is not None
|
||||||
assert updated_parent.balance == 10_000
|
assert updated_parent.balance == 10_000
|
||||||
assert updated_parent.reserved_balance == 3_000
|
assert updated_parent.reserved_balance == 0
|
||||||
assert updated_parent.total_spent == 0
|
assert updated_parent.total_spent == 0
|
||||||
updated_child = await _api_key(session, "ehbp-missing-child")
|
updated_child = await _api_key(session, "ehbp-missing-child")
|
||||||
assert updated_child is not None
|
assert updated_child is not None
|
||||||
assert updated_child.reserved_balance == 3_000
|
assert updated_child.reserved_balance == 0
|
||||||
assert updated_child.total_spent == 0
|
assert updated_child.total_spent == 0
|
||||||
|
release = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert release is not None and release.status == "released"
|
||||||
|
assert reservation.release_id not in auth_module._reservation_heartbeats
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_corrupt_child_aggregate_does_not_erase_parent_sibling_reserve(
|
||||||
|
session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
parent = ApiKey(hashed_key="ehbp-corrupt-parent", balance=10_000)
|
||||||
|
child = ApiKey(
|
||||||
|
hashed_key="ehbp-corrupt-child",
|
||||||
|
balance=0,
|
||||||
|
parent_key_hash=parent.hashed_key,
|
||||||
|
)
|
||||||
|
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 session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == "ehbp-corrupt-parent")
|
||||||
|
.values(reserved_balance=5_000)
|
||||||
|
)
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == "ehbp-corrupt-child")
|
||||||
|
.values(reserved_balance=1_000)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
charged = await finalize_ehbp_actual_cost_payment(
|
||||||
|
child,
|
||||||
|
session,
|
||||||
|
reserved_cost_for_model=3_000,
|
||||||
|
model_id="tinfoil/model",
|
||||||
|
cost_info={"total_msats": 1_200},
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert charged == 0
|
||||||
|
updated_parent = await _api_key(session, "ehbp-corrupt-parent")
|
||||||
|
updated_child = await _api_key(session, "ehbp-corrupt-child")
|
||||||
|
assert updated_parent is not None and updated_child is not None
|
||||||
|
assert updated_parent.balance == 10_000
|
||||||
|
assert updated_parent.reserved_balance == 5_000
|
||||||
|
assert updated_parent.total_spent == 0
|
||||||
|
assert updated_child.reserved_balance == 1_000
|
||||||
|
assert updated_child.total_spent == 0
|
||||||
|
release = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert release is not None and release.status == "released"
|
||||||
|
assert reservation.release_id not in auth_module._reservation_heartbeats
|
||||||
|
|
||||||
|
|
||||||
|
def test_zero_debit_ehbp_headers_preserve_computed_cost() -> None:
|
||||||
|
headers: dict[str, str] = {}
|
||||||
|
|
||||||
|
_inject_cost_response_headers(
|
||||||
|
headers,
|
||||||
|
{
|
||||||
|
"total_msats": 0,
|
||||||
|
"computed_msats": 1_500,
|
||||||
|
"input_msats": 1_200,
|
||||||
|
"output_msats": 300,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert headers["X-Routstr-Cost-Msats"] == "0"
|
||||||
|
assert headers["X-Routstr-Computed-Cost-Msats"] == "1500"
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
from collections.abc import AsyncGenerator
|
from collections.abc import AsyncGenerator
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
@@ -96,7 +97,10 @@ async def test_release_updates_parent_and_child_atomically() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_release_rolls_back_partial_parent_child_update() -> None:
|
async def test_release_repairs_partial_parent_child_corruption() -> None:
|
||||||
|
"""A child aggregate that no longer holds the reservation must not leave
|
||||||
|
the durable row active forever: the release rolls the subtraction back and
|
||||||
|
terminalizes the reservation without touching aggregates."""
|
||||||
engine = await _engine()
|
engine = await _engine()
|
||||||
parent = ApiKey(hashed_key="parent", balance=1_000)
|
parent = ApiKey(hashed_key="parent", balance=1_000)
|
||||||
child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0)
|
child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0)
|
||||||
@@ -109,12 +113,13 @@ async def test_release_rolls_back_partial_parent_child_update() -> None:
|
|||||||
session.add(child)
|
session.add(child)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
||||||
assert await release_reservation(snapshot, session, 500) is False
|
assert await release_reservation(snapshot, session, 500) is True
|
||||||
await session.refresh(parent)
|
await session.refresh(parent)
|
||||||
await session.refresh(child)
|
await session.refresh(child)
|
||||||
record = await session.get(ReservationRelease, snapshot.release_id)
|
record = await session.get(ReservationRelease, snapshot.release_id)
|
||||||
|
# Aggregates untouched — legacy cleanup reconciles them when stale.
|
||||||
assert (parent.reserved_balance, child.reserved_balance) == (500, 100)
|
assert (parent.reserved_balance, child.reserved_balance) == (500, 100)
|
||||||
assert record is not None and record.status == "active"
|
assert record is not None and record.status == "released"
|
||||||
await engine.dispose()
|
await engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
@@ -333,6 +338,73 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure(
|
|||||||
release.assert_awaited_once_with(snapshot, session, 500)
|
release.assert_awaited_once_with(snapshot, session, 500)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_responses_streaming_duplicate_publishes_zero_settled_cost() -> 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":2,"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-duplicate"
|
||||||
|
key.balance = 10_000
|
||||||
|
session = MagicMock()
|
||||||
|
session.get = AsyncMock(return_value=key)
|
||||||
|
session_context = MagicMock()
|
||||||
|
session_context.__aenter__ = AsyncMock(return_value=session)
|
||||||
|
session_context.__aexit__ = AsyncMock(return_value=None)
|
||||||
|
cost_data = {
|
||||||
|
"input_tokens": 2,
|
||||||
|
"output_tokens": 1,
|
||||||
|
"input_msats": 1_000,
|
||||||
|
"output_msats": 500,
|
||||||
|
"total_msats": 1_500,
|
||||||
|
"charged_msats": 0,
|
||||||
|
"total_usd": 0.0001,
|
||||||
|
}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"routstr.upstream.base.adjust_payment_for_tokens",
|
||||||
|
AsyncMock(return_value=cost_data),
|
||||||
|
),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
chunks = [
|
||||||
|
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||||
|
async for chunk in response.body_iterator
|
||||||
|
]
|
||||||
|
|
||||||
|
completed = next(
|
||||||
|
json.loads(line[6:])
|
||||||
|
for line in "".join(chunks).splitlines()
|
||||||
|
if line.startswith("data: {")
|
||||||
|
)
|
||||||
|
nested_usage = completed["response"]["usage"]
|
||||||
|
assert nested_usage["cost_sats"] == 0
|
||||||
|
assert nested_usage["cost"]["total_msats"] == 0
|
||||||
|
assert nested_usage["cost"]["charged_msats"] == 0
|
||||||
|
assert nested_usage["cost"]["computed_msats"] == 1_500
|
||||||
|
assert completed["cost"]["total_msats"] == 0
|
||||||
|
assert completed["cost"]["computed_msats"] == 1_500
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("via_litellm", [False, True])
|
@pytest.mark.parametrize("via_litellm", [False, True])
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
|
|||||||
@@ -96,6 +96,8 @@ export function TemporaryBalances({
|
|||||||
const total = data?.total ?? 0;
|
const total = data?.total ?? 0;
|
||||||
const totals = data?.totals ?? {
|
const totals = data?.totals ?? {
|
||||||
total_balance: 0,
|
total_balance: 0,
|
||||||
|
total_reserved_balance: 0,
|
||||||
|
total_available_balance: 0,
|
||||||
total_spent: 0,
|
total_spent: 0,
|
||||||
total_requests: 0,
|
total_requests: 0,
|
||||||
};
|
};
|
||||||
@@ -180,7 +182,7 @@ export function TemporaryBalances({
|
|||||||
<Card>
|
<Card>
|
||||||
<CardHeader className='flex flex-row items-center justify-between space-y-0 pb-2'>
|
<CardHeader className='flex flex-row items-center justify-between space-y-0 pb-2'>
|
||||||
<CardTitle className='text-muted-foreground text-sm font-medium'>
|
<CardTitle className='text-muted-foreground text-sm font-medium'>
|
||||||
Total Balance
|
Total Available
|
||||||
</CardTitle>
|
</CardTitle>
|
||||||
<span className='inline-flex size-8 items-center justify-center'>
|
<span className='inline-flex size-8 items-center justify-center'>
|
||||||
<DollarSign className='size-4 text-green-600 dark:text-green-300' />
|
<DollarSign className='size-4 text-green-600 dark:text-green-300' />
|
||||||
@@ -188,7 +190,11 @@ export function TemporaryBalances({
|
|||||||
</CardHeader>
|
</CardHeader>
|
||||||
<CardContent className='pt-0'>
|
<CardContent className='pt-0'>
|
||||||
<p className='text-2xl font-semibold tracking-tight tabular-nums'>
|
<p className='text-2xl font-semibold tracking-tight tabular-nums'>
|
||||||
{formatBalance(totals.total_balance)}
|
{formatBalance(totals.total_available_balance)}
|
||||||
|
</p>
|
||||||
|
<p className='text-muted-foreground mt-1 text-xs'>
|
||||||
|
{formatBalance(totals.total_balance)} raw ·{' '}
|
||||||
|
{formatBalance(totals.total_reserved_balance)} reserved
|
||||||
</p>
|
</p>
|
||||||
</CardContent>
|
</CardContent>
|
||||||
</Card>
|
</Card>
|
||||||
@@ -263,7 +269,7 @@ export function TemporaryBalances({
|
|||||||
<TableHeader>
|
<TableHeader>
|
||||||
<TableRow>
|
<TableRow>
|
||||||
<TableHead>Hashed Key</TableHead>
|
<TableHead>Hashed Key</TableHead>
|
||||||
<TableHead className='text-right'>Balance</TableHead>
|
<TableHead className='text-right'>Available</TableHead>
|
||||||
<TableHead className='text-right'>
|
<TableHead className='text-right'>
|
||||||
Total Spent
|
Total Spent
|
||||||
</TableHead>
|
</TableHead>
|
||||||
@@ -284,7 +290,9 @@ export function TemporaryBalances({
|
|||||||
<TableRow
|
<TableRow
|
||||||
key={`${balance.hashed_key}-${balance.parent_key_hash ?? 'root'}-${index}`}
|
key={`${balance.hashed_key}-${balance.parent_key_hash ?? 'root'}-${index}`}
|
||||||
className={cn(
|
className={cn(
|
||||||
balance.balance === 0 && !isChild && 'opacity-60',
|
balance.available_balance === 0 &&
|
||||||
|
!isChild &&
|
||||||
|
'opacity-60',
|
||||||
isChild && 'bg-muted/30'
|
isChild && 'bg-muted/30'
|
||||||
)}
|
)}
|
||||||
>
|
>
|
||||||
@@ -307,7 +315,19 @@ export function TemporaryBalances({
|
|||||||
(Parent)
|
(Parent)
|
||||||
</span>
|
</span>
|
||||||
) : (
|
) : (
|
||||||
formatBalance(balance.balance)
|
<div>
|
||||||
|
<div>
|
||||||
|
{formatBalance(
|
||||||
|
balance.available_balance ??
|
||||||
|
balance.balance
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
<div className='text-muted-foreground text-xs'>
|
||||||
|
{formatBalance(balance.balance)} raw ·{' '}
|
||||||
|
{formatBalance(balance.reserved_balance)}{' '}
|
||||||
|
reserved
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
)}
|
)}
|
||||||
</TableCell>
|
</TableCell>
|
||||||
<TableCell className='text-right font-mono'>
|
<TableCell className='text-right font-mono'>
|
||||||
@@ -350,7 +370,9 @@ export function TemporaryBalances({
|
|||||||
<Card
|
<Card
|
||||||
key={`${balance.hashed_key}-${balance.parent_key_hash ?? 'root'}-mobile-${index}`}
|
key={`${balance.hashed_key}-${balance.parent_key_hash ?? 'root'}-mobile-${index}`}
|
||||||
className={cn(
|
className={cn(
|
||||||
balance.balance === 0 && !isChild && 'opacity-80',
|
balance.available_balance === 0 &&
|
||||||
|
!isChild &&
|
||||||
|
'opacity-80',
|
||||||
isChild && 'bg-muted/30'
|
isChild && 'bg-muted/30'
|
||||||
)}
|
)}
|
||||||
>
|
>
|
||||||
@@ -372,13 +394,22 @@ export function TemporaryBalances({
|
|||||||
<CardContent className='grid grid-cols-2 gap-3 p-4 pt-0'>
|
<CardContent className='grid grid-cols-2 gap-3 p-4 pt-0'>
|
||||||
<div>
|
<div>
|
||||||
<p className='text-muted-foreground text-xs'>
|
<p className='text-muted-foreground text-xs'>
|
||||||
Balance
|
Available
|
||||||
</p>
|
</p>
|
||||||
<p className='font-mono text-sm'>
|
<p className='font-mono text-sm'>
|
||||||
{isChild
|
{isChild
|
||||||
? '(Uses Parent)'
|
? '(Uses Parent)'
|
||||||
: formatBalance(balance.balance)}
|
: formatBalance(
|
||||||
|
balance.available_balance ?? balance.balance
|
||||||
|
)}
|
||||||
</p>
|
</p>
|
||||||
|
{!isChild && (
|
||||||
|
<p className='text-muted-foreground text-xs'>
|
||||||
|
{formatBalance(balance.balance)} raw ·{' '}
|
||||||
|
{formatBalance(balance.reserved_balance)}{' '}
|
||||||
|
reserved
|
||||||
|
</p>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
<div>
|
<div>
|
||||||
<p className='text-muted-foreground text-xs'>
|
<p className='text-muted-foreground text-xs'>
|
||||||
|
|||||||
@@ -1082,6 +1082,8 @@ export interface CliTokenCreated {
|
|||||||
export const TemporaryBalanceSchema = z.object({
|
export const TemporaryBalanceSchema = z.object({
|
||||||
hashed_key: z.string(),
|
hashed_key: z.string(),
|
||||||
balance: z.number(),
|
balance: z.number(),
|
||||||
|
reserved_balance: z.number(),
|
||||||
|
available_balance: z.number().nullable(),
|
||||||
total_spent: z.number(),
|
total_spent: z.number(),
|
||||||
total_requests: z.number(),
|
total_requests: z.number(),
|
||||||
refund_address: z.string().nullable(),
|
refund_address: z.string().nullable(),
|
||||||
@@ -1097,6 +1099,8 @@ export interface TemporaryBalancesResponse {
|
|||||||
total: number;
|
total: number;
|
||||||
totals: {
|
totals: {
|
||||||
total_balance: number;
|
total_balance: number;
|
||||||
|
total_reserved_balance: number;
|
||||||
|
total_available_balance: number;
|
||||||
total_spent: number;
|
total_spent: number;
|
||||||
total_requests: number;
|
total_requests: number;
|
||||||
};
|
};
|
||||||
|
|||||||
Reference in New Issue
Block a user