From 47807ae92ee5ae34341bad6e8607e6197b591bcb Mon Sep 17 00:00:00 2001 From: thefux Date: Thu, 3 Sep 2026 11:33:39 +0000 Subject: [PATCH 1/3] Remove child key feature and balance limits completely Removes the child-key (shared parent balance) feature and the balance-limit machinery from the entire stack: Backend: - Drop POST /v1/balance/child-key and /v1/balance/child-key/reset - Remove parent-key billing indirection (get_billing_key); every key is charged directly - Remove balance_limit/balance_limit_reset enforcement, reset policies, and the periodic limit-reset background task - Remove child-key guards on refund/history endpoints - Remove child_key_cost setting and /v1/info child_key_cost_msats - Remove balance_limit fields from LightningInvoice model and invoice-creation API - Admin balances API returns plain sums (no parent/child split) Migration (e5a6b7c8d9f0): - Nulls parent_key_hash on children (they become standalone keys; parents keep 100% of their balance, so no funds are lost) - Drops parent_key_hash + balance_limit columns from api_keys and balance_limit columns from lightning_invoices - Verified: upgrade, downgrade, and fund preservation round-trip UI: - Remove child-key creator, child key panels, balance-limit inputs - Key options now offer validity date only - Temporary balances table renders all keys uniformly Tests: child-key suites removed; remaining suites converted to single-key semantics (441 integration + 1313 unit tests pass). Docs updated accordingly. --- docs/api/authentication.md | 17 - docs/api/endpoints.md | 54 +- examples/create_child_keys.py | 45 -- ...f0_remove_child_keys_and_balance_limits.py | 76 +++ routstr/auth.py | 297 +---------- routstr/balance.py | 187 +------ routstr/core/admin.py | 51 +- routstr/core/db.py | 43 +- routstr/core/main.py | 3 - routstr/core/settings.py | 1 - routstr/lightning.py | 10 - routstr/upstream/ehbp.py | 17 +- tests/integration/test_child_keys.py | 190 ------- tests/integration/test_child_keys_api.py | 102 ---- tests/integration/test_failover_billing.py | 99 +--- tests/integration/test_key_logic.py | 367 +------------ .../test_lightning_invoice_constraints.py | 68 +-- tests/integration/test_payment_invariants.py | 99 +--- tests/integration/test_prune_dead_api_keys.py | 28 - .../test_reserved_balance_negative.py | 102 ---- .../test_temporary_balances_api.py | 54 +- tests/unit/test_balance.py | 15 - tests/unit/test_ehbp_finalize_payment.py | 117 +---- tests/unit/test_stale_reservations.py | 56 +- .../test_streaming_billing_finalization.py | 39 +- ui/components/child-key-creator.tsx | 481 ------------------ ui/components/key-options.tsx | 65 +-- .../landing/cashu-payment-workflow.tsx | 22 - ui/components/landing/cheat-sheet.tsx | 13 +- ui/components/landing/key-info-details.tsx | 184 +------ .../landing/lightning-payment-workflow.tsx | 25 - ui/components/temporary-balances.tsx | 92 +--- ui/hooks/use-wallet-info.ts | 5 - ui/lib/api/services/admin.ts | 1 - ui/lib/api/services/wallet.ts | 88 ---- 35 files changed, 189 insertions(+), 2924 deletions(-) delete mode 100644 examples/create_child_keys.py create mode 100644 migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py delete mode 100644 tests/integration/test_child_keys.py delete mode 100644 tests/integration/test_child_keys_api.py delete mode 100644 ui/components/child-key-creator.tsx diff --git a/docs/api/authentication.md b/docs/api/authentication.md index c7b91b3b..ec4c1449 100644 --- a/docs/api/authentication.md +++ b/docs/api/authentication.md @@ -307,23 +307,6 @@ ANALYTICS_KEY = os.getenv("ROUTSTR_ANALYTICS_KEY") api_key = PROD_KEY if is_production() else DEV_KEY ``` -### Delegated Authentication - -Create sub-keys with limited permissions: - -```bash -POST /v1/wallet/create/subkey -Authorization: Bearer sk-parent-key -Content-Type: application/json - -{ - "name": "Limited Subkey", - "balance_limit": 1000, - "allowed_models": ["gpt-3.5-turbo"], - "expires_in_hours": 24 -} -``` - ## Rate Limiting Rate limits are applied per API key: diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 23b019c7..dfd619f2 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -418,7 +418,7 @@ POST /v1/wallet/create ### Get Key Information -Get current balance, consumption data, and child keys for an API key. +Get current balance and consumption data for an API key. ```http GET /v1/balance/info @@ -432,23 +432,9 @@ Authorization: Bearer sk-... "api_key": "sk-abc...", "balance": 8500000, "reserved": 0, - "is_child": false, - "parent_key": null, "total_requests": 42, "total_spent": 1500000, - "balance_limit": null, - "balance_limit_reset": null, - "validity_date": null, - "child_keys": [ - { - "api_key": "sk-child1...", - "total_requests": 10, - "total_spent": 500000, - "balance_limit": 1000000, - "balance_limit_reset": "daily", - "validity_date": 1738000000 - } - ] + "validity_date": null } ``` @@ -528,42 +514,6 @@ Authorization: Bearer sk-... } ``` -### Create Child Key - -Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable). - -```http -POST /v1/balance/child-key -Authorization: Bearer sk-... -``` - -**Request Body:** - -```json -{ - "count": 1 -} -``` - -**Parameters:** - -| Parameter | Type | Required | Default | Description | -|-----------|------|----------|---------|-------------| -| `count` | integer | Yes | - | Number of child keys to create (1-50) | - -**Response:** - -```json -{ - "api_keys": ["sk-abc...", "sk-def..."], - "count": 2, - "cost_msats": 2000, - "cost_sats": 2, - "parent_balance": 98000, - "parent_balance_sats": 98 -} -``` - ## Provider Discovery ## Admin Settings diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py deleted file mode 100644 index 24f4556d..00000000 --- a/examples/create_child_keys.py +++ /dev/null @@ -1,45 +0,0 @@ -import json -import sys - -import httpx - - -def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]: - headers = {"Authorization": f"Bearer {api_key}"} - - print(f"Requesting {count} child keys from {base_url}...") - - child_keys = [] - - for i in range(count): - try: - response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers) - if response.status_code == 200: - data = response.json() - child_keys.append(data["api_key"]) - print( - f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)" - ) - else: - print(f" [{i + 1}] Failed: {response.status_code} - {response.text}") - except Exception as e: - print(f" [{i + 1}] Error: {str(e)}") - - return child_keys - - -if __name__ == "__main__": - if len(sys.argv) < 2: - print("Usage: python create_child_keys.py [base_url]") - sys.exit(1) - - auth_key = sys.argv[1] - base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000" - - keys = create_child_keys(base_url, auth_key) - - if keys: - print("\nSuccessfully created child keys:") - print(json.dumps(keys, indent=2)) - else: - print("\nNo child keys were created.") diff --git a/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py b/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py new file mode 100644 index 00000000..15298de2 --- /dev/null +++ b/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py @@ -0,0 +1,76 @@ +"""Remove child keys and balance limits. + +Removes the child-key feature (parent_key_hash) and the balance-limit +machinery (balance_limit, balance_limit_reset, balance_limit_reset_date) +from api_keys, plus the balance_limit/balance_limit_reset pass-through on +lightning_invoices. + +Data preservation: before dropping the columns, every child key is +converted into a standalone key by clearing parent_key_hash. Child keys +never hold their own balance (they always spent from their parent), so no +funds are lost: the parent keeps its full balance, and the former child +rows are preserved with their total_spent/total_requests history intact. +""" + +import sqlalchemy as sa +from alembic import op + +revision = "e5a6b7c8d9f0" +down_revision = "b4f7a1c9d2e3" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Convert child keys into standalone keys before dropping the link. + # Their balance is always 0 (they spent from the parent), so this + # cannot strand any funds. + op.execute("UPDATE api_keys SET parent_key_hash = NULL") + + with op.batch_alter_table("api_keys") as batch_op: + batch_op.drop_index("ix_api_keys_parent_key_hash") + batch_op.drop_column("parent_key_hash") + batch_op.drop_column("balance_limit") + batch_op.drop_column("balance_limit_reset") + batch_op.drop_column("balance_limit_reset_date") + + with op.batch_alter_table("lightning_invoices") as batch_op: + batch_op.drop_column("balance_limit") + batch_op.drop_column("balance_limit_reset") + + +def downgrade() -> None: + with op.batch_alter_table("lightning_invoices") as batch_op: + batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True)) + batch_op.add_column( + sa.Column("balance_limit_reset", sa.String(), nullable=True) + ) + + with op.batch_alter_table("api_keys") as batch_op: + batch_op.add_column( + sa.Column("balance_limit_reset_date", sa.Integer(), nullable=True) + ) + batch_op.add_column( + sa.Column( + "balance_limit_reset", + sa.String(), + nullable=True, + ) + ) + batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True)) + batch_op.add_column( + sa.Column( + "parent_key_hash", + sa.String(), + nullable=True, + ) + ) + batch_op.create_foreign_key( + "fk_api_keys_parent_key_hash", + "api_keys", + ["parent_key_hash"], + ["hashed_key"], + ) + batch_op.create_index( + "ix_api_keys_parent_key_hash", ["parent_key_hash"], unique=False + ) diff --git a/routstr/auth.py b/routstr/auth.py index 7b808007..0478dfa8 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -1,13 +1,11 @@ import asyncio import hashlib import math -import random import time import uuid from contextlib import suppress from contextvars import ContextVar from dataclasses import dataclass -from datetime import datetime from typing import TYPE_CHECKING, Optional from fastapi import HTTPException @@ -99,48 +97,6 @@ def _clear_current_reservation(snapshot: ReservationSnapshot) -> None: # PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats -async def check_and_reset_limit(key: ApiKey, session: AsyncSession) -> bool: - """Checks if a key's balance limit should be reset based on its policy.""" - if key.balance_limit is not None and key.balance_limit_reset: - now = int(time.time()) - reset_date = key.balance_limit_reset_date or 0 - should_reset = False - - if key.balance_limit_reset == "daily": - if ( - datetime.fromtimestamp(now).date() - > datetime.fromtimestamp(reset_date).date() - ): - should_reset = True - elif key.balance_limit_reset == "weekly": - if ( - datetime.fromtimestamp(now).isocalendar()[:2] - > datetime.fromtimestamp(reset_date).isocalendar()[:2] - ): - should_reset = True - elif key.balance_limit_reset == "monthly": - dt_now = datetime.fromtimestamp(now) - dt_reset = datetime.fromtimestamp(reset_date) - if dt_now.year > dt_reset.year or dt_now.month > dt_reset.month: - should_reset = True - - if should_reset: - logger.info( - "Resetting balance limit for key", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "policy": key.balance_limit_reset, - "old_spent": key.total_spent, - }, - ) - key.total_spent = 0 - key.balance_limit_reset_date = now - session.add(key) - await session.flush() - return True - return False - - def redemption_error_to_http_exception(error: Exception) -> HTTPException: """Map a Cashu token redemption failure to a sanitized client-facing error. @@ -315,42 +271,19 @@ async def _validate_bearer_key_locked( }, ) - # Check and reset limit if needed - await check_and_reset_limit(existing_key, session) - - # Early check: Billing balance check (Parent balance) - billing_key = await get_billing_key(existing_key, session) - if min_cost > 0 and billing_key.total_balance < min_cost: + # Early check: Billing balance check + if min_cost > 0 and existing_key.total_balance < min_cost: logger.warning( "Insufficient billing balance during validation", extra={ "key_hash": existing_key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "balance": billing_key.total_balance, + "balance": existing_key.total_balance, "required": min_cost, }, ) raise HTTPException( status_code=402, - detail=_model_balance_error(min_cost, billing_key.total_balance), - ) - - # Early check: Spending limit check (Child key limit) - if ( - min_cost > 0 - and existing_key.balance_limit is not None - and existing_key.total_spent + existing_key.reserved_balance + min_cost - > existing_key.balance_limit - ): - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Balance limit exceeded: {existing_key.balance_limit} mSats limit. {existing_key.total_spent} already spent ({existing_key.reserved_balance} reserved), {min_cost} minimum required for this model.", - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, + detail=_model_balance_error(min_cost, existing_key.total_balance), ) return existing_key @@ -619,27 +552,6 @@ async def _validate_bearer_key_locked( ) -async def get_billing_key(key: ApiKey, session: AsyncSession) -> ApiKey: - """Returns the key that should be charged for the request.""" - if key.parent_key_hash: - parent = await session.get(ApiKey, key.parent_key_hash) - if parent: - # We want to keep the total_requests and total_spent on the child key - # but use the balance and reserved_balance of the parent. - # However, pay_for_request updates reserved_balance and total_requests. - # To stay simple, we charge the parent's balance and update parent's total_requests. - return parent - else: - logger.error( - "Parent key not found for child key", - extra={ - "child_key_hash": key.hashed_key[:8] + "...", - "parent_key_hash": key.parent_key_hash[:8] + "...", - }, - ) - return key - - async def pay_for_request( key: ApiKey, cost_per_request: int, session: AsyncSession ) -> int: @@ -647,7 +559,7 @@ async def pay_for_request( # Ensure cost_per_request is at least the minimum allowed request cost cost_per_request = max(cost_per_request, settings.min_request_msat) - billing_key = await get_billing_key(key, session) + billing_key = key logger.info( "Processing payment for request", @@ -706,35 +618,6 @@ async def pay_for_request( }, ) - # Check balance limit for child keys (or any key with a limit) - if key.balance_limit is not None: - await check_and_reset_limit(key, session) - - if ( - key.total_spent + key.reserved_balance + cost_per_request - > key.balance_limit - ): - logger.warning( - "Balance limit exceeded", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "total_spent": key.total_spent, - "reserved": key.reserved_balance, - "balance_limit": key.balance_limit, - "required": cost_per_request, - }, - ) - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.", - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, - ) - logger.debug( "Charging base cost for request", extra={ @@ -798,52 +681,6 @@ async def pay_for_request( }, ) - # Also increment total_requests and reserved_balance on the child key if it's different. - # The balance_limit guard is enforced atomically here — the Python pre-check above - # is a fast-path rejection only and provides no concurrency guarantee. - if billing_key.hashed_key != key.hashed_key: - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where( - (col(ApiKey.balance_limit).is_(None)) - | ( - col(ApiKey.total_spent) - + col(ApiKey.reserved_balance) - + cost_per_request - <= col(ApiKey.balance_limit) - ) - ) - .values( - total_requests=col(ApiKey.total_requests) + 1, - reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, - reserved_at=reserved_at_now, - ) - ) - child_result = await session.exec(child_stmt) # type: ignore[call-overload] - - if child_result.rowcount == 0: - # Build the error before rollback expires ORM attributes. - limit_message = ( - f"Balance limit exceeded: {key.balance_limit} mSats limit. " - f"{key.total_spent} already spent ({key.reserved_balance} reserved), " - f"{cost_per_request} required for this request." - ) - # The parent reservation update already ran in this transaction. - # Roll it back before failover code attempts to restore the previous - # reservation; otherwise that later commit can persist both updates. - await session.rollback() - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": limit_message, - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, - ) - session.add( ReservationRelease( id=reservation.release_id, @@ -895,8 +732,6 @@ async def pay_for_request( try: await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) except Exception: # The reservation transaction is already committed and durable. Logging # refresh failures must not make the caller treat it as unreserved. @@ -968,9 +803,6 @@ async def _validate_reservation_snapshot( persisted_key = await session.get(ApiKey, snapshot.key_hash) if persisted_key is None: raise RuntimeError("Billing reservation key no longer exists") - expected_billing_hash = persisted_key.parent_key_hash or persisted_key.hashed_key - if snapshot.billing_key_hash != expected_billing_hash: - raise RuntimeError("Billing reservation does not belong to this billing key") record = await session.get(ReservationRelease, snapshot.release_id) if ( @@ -1184,22 +1016,6 @@ async def _transition_reservation_to_released( snapshot, session, decrement_requests=decrement_requests ) - if snapshot.billing_key_hash != snapshot.key_hash: - child_release_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == snapshot.key_hash) - .where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats) - .values(**values) - ) - child_result = await session.exec( # type: ignore[call-overload] - child_release_stmt - ) - if child_result.rowcount != 1: - await session.rollback() - return await _repair_corrupt_reservation( - snapshot, session, decrement_requests=decrement_requests - ) - await session.commit() await _stop_reservation_heartbeat(snapshot.release_id) _clear_current_reservation(snapshot) @@ -1251,15 +1067,14 @@ 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. + """Release the reserved amount and record the charge on the key + inside the caller's transaction. - Guarded subtraction replaces defensive clamping: every row must still hold + Guarded subtraction replaces defensive clamping: the 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. @@ -1288,26 +1103,6 @@ async def _charge_reservation_rows( 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 @@ -1332,7 +1127,7 @@ async def adjust_payment_for_tokens( The response's usage object is normalized with the default union parser in ``calculate_cost``. """ - billing_key = await get_billing_key(key, session) + billing_key = key reservation = reservation_snapshot or await get_reservation_snapshot(key, session) await _validate_reservation_snapshot( key, reservation, session, require_active=False @@ -1421,7 +1216,6 @@ async def adjust_payment_for_tokens( charged = await _charge_reservation_rows( session, billing_key_hash=billing_key.hashed_key, - key_hash=key.hashed_key, reserved_msats=deducted_max_cost, charge_msats=cost.total_msats, ) @@ -1444,8 +1238,6 @@ async def adjust_payment_for_tokens( else: cost.charged_msats = cost.total_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) logger.info( "Max cost payment finalized", extra={ @@ -1512,7 +1304,6 @@ async def adjust_payment_for_tokens( 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, ): @@ -1534,8 +1325,6 @@ async def adjust_payment_for_tokens( await _stop_reservation_heartbeat(reservation.release_id) cost.charged_msats = total_cost_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) await _accumulate_fee(total_cost_msats) payments_logger.info( "FINALIZE", @@ -1600,7 +1389,6 @@ async def adjust_payment_for_tokens( 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=( @@ -1620,8 +1408,6 @@ async def adjust_payment_for_tokens( await _stop_reservation_heartbeat(reservation.release_id) await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) cost.charged_msats = actual_charge_msats if actual_charge_msats < total_cost_msats: logger.warning( @@ -1680,7 +1466,6 @@ async def adjust_payment_for_tokens( charged = 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, ) @@ -1705,8 +1490,6 @@ async def adjust_payment_for_tokens( cost.total_msats = total_cost_msats cost.charged_msats = total_cost_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) logger.info( "Refund processed successfully", @@ -1767,68 +1550,6 @@ async def adjust_payment_for_tokens( raise AssertionError("Unreachable: unhandled calculate_cost result") -async def periodic_key_reset() -> None: - """Background task to reset key limits based on their policy.""" - from .core.db import create_session - - while True: - try: - interval = 3600 # Run every hour - jitter = 300 - await asyncio.sleep(interval + random.uniform(0, jitter)) - except asyncio.CancelledError: - break - - try: - async with create_session() as session: - # Find all keys that have a reset policy - stmt = select(ApiKey).where(ApiKey.balance_limit_reset.is_not(None)) # type: ignore - keys = (await session.exec(stmt)).all() - - now = int(time.time()) - updated_count = 0 - - for key in keys: - reset_date = key.balance_limit_reset_date or 0 - should_reset = False - - if key.balance_limit_reset == "daily": - if ( - datetime.fromtimestamp(now).date() - > datetime.fromtimestamp(reset_date).date() - ): - should_reset = True - elif key.balance_limit_reset == "weekly": - if ( - datetime.fromtimestamp(now).isocalendar()[:2] - > datetime.fromtimestamp(reset_date).isocalendar()[:2] - ): - should_reset = True - elif key.balance_limit_reset == "monthly": - dt_now = datetime.fromtimestamp(now) - dt_reset = datetime.fromtimestamp(reset_date) - if dt_now.year > dt_reset.year or dt_now.month > dt_reset.month: - should_reset = True - - if should_reset: - key.total_spent = 0 - key.balance_limit_reset_date = now - session.add(key) - updated_count += 1 - - if updated_count > 0: - await session.commit() - logger.info( - "Periodic key reset complete", - extra={"keys_reset": updated_count}, - ) - - except asyncio.CancelledError: - break - except Exception as e: - logger.error(f"Error in periodic_key_reset: {e}") - - async def periodic_dead_key_prune() -> None: """Periodically prune dead API keys. Interval <= 0 disables it. diff --git a/routstr/balance.py b/routstr/balance.py index 10f483f7..ecd770b5 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -1,6 +1,5 @@ import asyncio import hashlib -import time from time import monotonic from typing import Annotated, NoReturn @@ -10,7 +9,6 @@ from pydantic import BaseModel from sqlmodel import col, select, update from .auth import ( - get_billing_key, redemption_error_to_http_exception, validate_bearer_key, ) @@ -58,39 +56,15 @@ async def get_key_from_header( async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict: - billing_key = await get_billing_key(key, session) info = { "api_key": "sk-" + key.hashed_key, - "balance": billing_key.total_balance, - "reserved": billing_key.reserved_balance, - "is_child": key.parent_key_hash is not None, + "balance": key.total_balance, + "reserved": key.reserved_balance, "total_requests": key.total_requests, "total_spent": key.total_spent, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, } - if key.parent_key_hash: - info["parent_key_preview"] = key.parent_key_hash[:8] + "..." - else: - # Fetch child keys if this is a parent key - statement = select(ApiKey).where(ApiKey.parent_key_hash == key.hashed_key) - results = await session.exec(statement) - child_keys = results.all() - if child_keys: - info["child_keys"] = [ - { - "api_key": "sk-" + ck.hashed_key, - "total_requests": ck.total_requests, - "total_spent": ck.total_spent, - "balance_limit": ck.balance_limit, - "balance_limit_reset": ck.balance_limit_reset, - "validity_date": ck.validity_date, - } - for ck in child_keys - ] - return info @@ -117,26 +91,18 @@ async def account_info( class BalanceCreateRequest(BaseModel): initial_balance_token: str - balance_limit: int | None = None - balance_limit_reset: str | None = None validity_date: int | None = None async def _create_balance( initial_balance_token: str, - balance_limit: int | None, - balance_limit_reset: str | None, validity_date: int | None, session: AsyncSession, ) -> dict: key = await validate_bearer_key(initial_balance_token, session) - if balance_limit is not None or balance_limit_reset or validity_date: - key.balance_limit = balance_limit - key.balance_limit_reset = balance_limit_reset + if validity_date is not None: key.validity_date = validity_date - if balance_limit_reset: - key.balance_limit_reset_date = int(time.time()) session.add(key) await session.commit() await session.refresh(key) @@ -154,8 +120,6 @@ async def create_balance_from_body( ) -> dict: return await _create_balance( payload.initial_balance_token, - payload.balance_limit, - payload.balance_limit_reset, payload.validity_date, session, ) @@ -164,15 +128,11 @@ async def create_balance_from_body( @router.get("/create") async def create_balance( initial_balance_token: str, - balance_limit: int | None = None, - balance_limit_reset: str | None = None, validity_date: int | None = None, session: AsyncSession = Depends(get_session), ) -> dict: return await _create_balance( initial_balance_token, - balance_limit, - balance_limit_reset, validity_date, session, ) @@ -208,7 +168,7 @@ async def topup_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: - billing_key = await get_billing_key(key, session) + billing_key = key if topup_request is not None: cashu_token = topup_request.cashu_token @@ -458,12 +418,6 @@ async def refund_wallet_endpoint( if persisted := await _get_persisted_api_key_refund(key, session): return persisted - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot refund child key. Please refund the parent key instead.", - ) - if key.reserved_balance > 0: # Release only durable reservations old enough to be stale. A newer # request on the same aggregate balance must remain reserved. @@ -650,12 +604,6 @@ async def wallet_history( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, list[dict[str, str | int | bool | None]]]: - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot view child key history. Please use the parent key instead.", - ) - result = await session.exec( select(CashuTransaction) .where(CashuTransaction.api_key_hashed_key == key.hashed_key) @@ -692,133 +640,6 @@ async def donate(token: str, ref: str | None = None) -> str: except Exception: return "Invalid token." - -class ChildKeyRequest(BaseModel): - count: int - balance_limit: int | None = None - balance_limit_reset: str | None = None - validity_date: int | None = None - - -@router.post("/child-key") -async def create_child_key( - payload: ChildKeyRequest, - key: ApiKey = Depends(get_key_from_header), - session: AsyncSession = Depends(get_session), -) -> dict: - """Creates one or more child API keys that use the parent's balance.""" - # Log incoming request for debugging - logger.debug(f"Child key creation request: count={payload.count}") - - count = payload.count - if count < 1 or count > 50: - raise HTTPException(status_code=400, detail="Count must be between 1 and 50.") - - # Check if this is already a child key - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot create a child key for another child key.", - ) - - cost_per_key = settings.child_key_cost - total_cost = cost_per_key * count - - if key.total_balance < total_cost: - raise HTTPException( - status_code=402, - detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.", - ) - - # Deduct cost from parent atomically — guards against concurrent requests - # that both pass the balance check above on stale in-memory state. - deduct_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= total_cost) - .values( - balance=col(ApiKey.balance) - total_cost, - total_spent=col(ApiKey.total_spent) + total_cost, - ) - ) - result = await session.exec(deduct_stmt) # type: ignore[call-overload] - - if result.rowcount == 0: - raise HTTPException( - status_code=402, - detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.", - ) - - # Generate new keys - import secrets - - new_keys = [] - for _ in range(count): - new_key_raw = secrets.token_hex(32) - new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys - - child_key = ApiKey( - hashed_key=new_key_hash, - balance=0, - parent_key_hash=key.hashed_key, - balance_limit=payload.balance_limit, - balance_limit_reset=payload.balance_limit_reset, - balance_limit_reset_date=int(time.time()) - if payload.balance_limit_reset - else None, - validity_date=payload.validity_date, - ) - session.add(child_key) - new_keys.append("sk-" + new_key_hash) - - await session.commit() - await session.refresh(key) - - response_data = { - "api_keys": new_keys, - "count": count, - "cost_msats": total_cost, - "cost_sats": total_cost // 1000, - "parent_balance": key.balance, - "parent_balance_sats": key.balance // 1000, - } - logger.debug(f"Child key creation response: {response_data}") - return response_data - - -class ChildKeyResetRequest(BaseModel): - child_key: str - - -@router.post("/child-key/reset") -async def reset_child_key_spent( - payload: ChildKeyResetRequest, - key: ApiKey = Depends(get_key_from_header), - session: AsyncSession = Depends(get_session), -) -> dict: - """Resets the total_spent of a child key. Must be called by the parent.""" - child_key_raw = payload.child_key - if child_key_raw.startswith("sk-"): - child_key_raw = child_key_raw[3:] - - child_key = await session.get(ApiKey, child_key_raw) - if not child_key: - raise HTTPException(status_code=404, detail="Child key not found.") - - if child_key.parent_key_hash != key.hashed_key: - raise HTTPException( - status_code=403, detail="Unauthorized. You are not the parent of this key." - ) - - child_key.total_spent = 0 - if child_key.balance_limit_reset: - child_key.balance_limit_reset_date = int(time.time()) - session.add(child_key) - await session.commit() - - return {"success": True, "message": "Child key balance reset successfully."} - - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66e3a59a..46231554 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -112,42 +112,12 @@ async def get_temporary_balances_api( total = count_result.one() # Aggregate totals across the whole (search-filtered) set, not just the - # current page. Balance counts only parent (non-child) keys to avoid - # double-counting, since child keys draw from their parent's balance. balance_totals_result = await session.exec( select( + func.coalesce(func.sum(ApiKey.balance), 0), + func.coalesce(func.sum(ApiKey.reserved_balance), 0), func.coalesce( - func.sum( - case( - (col(ApiKey.parent_key_hash).is_(None), ApiKey.balance), - else_=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, + func.sum(col(ApiKey.balance) - col(ApiKey.reserved_balance)), 0 ), ).where(*filters) ) @@ -184,16 +154,11 @@ async def get_temporary_balances_api( "hashed_key": key.hashed_key, "balance": key.balance, "reserved_balance": key.reserved_balance, - "available_balance": ( - key.total_balance if key.parent_key_hash is None else None - ), + "available_balance": key.total_balance, "total_spent": key.total_spent, "total_requests": key.total_requests, "refund_address": key.refund_address, "key_expiry_time": key.key_expiry_time, - "parent_key_hash": key.parent_key_hash, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, "created_at": key.created_at, } @@ -211,8 +176,6 @@ async def get_temporary_balances_api( class ApiKeyUpdate(BaseModel): - balance_limit: int | None = None - balance_limit_reset: str | None = None validity_date: int | None = None @@ -227,10 +190,6 @@ async def update_apikey( if not key: raise HTTPException(status_code=404, detail="API key not found") - if update.balance_limit is not None: - key.balance_limit = update.balance_limit - if update.balance_limit_reset is not None: - key.balance_limit_reset = update.balance_limit_reset if update.validity_date is not None: key.validity_date = update.validity_date @@ -240,8 +199,6 @@ async def update_apikey( return { "hashed_key": key.hashed_key, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, } diff --git a/routstr/core/db.py b/routstr/core/db.py index 14cc669e..4c4d1a62 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -17,7 +17,6 @@ from sqlalchemy.engine import make_url from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlalchemy.orm import aliased from sqlmodel import Field, Relationship, SQLModel, col, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -138,21 +137,6 @@ class ApiKey(SQLModel, table=True): # type: ignore default=None, description="Currency of the cashu-token", ) - parent_key_hash: str | None = Field( - default=None, foreign_key="api_keys.hashed_key", index=True - ) - balance_limit: int | None = Field( - default=None, - description="Max spendable balance in msats for this key (mostly for child keys)", - ) - balance_limit_reset: str | None = Field( - default=None, - description="Reset policy for balance limit (manual, daily, monthly, etc.)", - ) - balance_limit_reset_date: int | None = Field( - default=None, - description="Unix timestamp of the last time the balance limit was reset", - ) validity_date: int | None = Field( default=None, description="Unix timestamp after which the key is no longer valid", @@ -317,10 +301,7 @@ async def release_stale_reservations( ) else: legacy_query = legacy_query.where( - or_( - col(ApiKey.hashed_key) == key_hash, - col(ApiKey.parent_key_hash) == key_hash, - ) + col(ApiKey.hashed_key) == key_hash ).where( or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff) ) @@ -362,21 +343,15 @@ async def release_stale_reservations( async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> int: - """Delete dead parentless API keys; return the count removed. + """Delete dead API keys; return the count removed. - Dead = 0 balance/reservation/spend/requests, older than the grace period, - no parent, no children, no invoice that could still settle. Cashu rows are + Dead = 0 balance/reservation/spend/requests, older than the grace + period, no invoice that could still settle. Cashu rows are unlinked (not deleted) first to keep the audit trail. """ now = int(time.time()) cutoff = now - min_age_seconds - child = aliased(ApiKey) - has_children = ( - select(child.hashed_key).where( - col(child.parent_key_hash) == col(ApiKey.hashed_key) - ) - ).exists() # An expired invoice stays creditable for the grace window, and crediting it # after its target key is gone strands the payment at the mint. settleable_invoice = ( @@ -400,10 +375,8 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in .where(col(ApiKey.reserved_balance) == 0) .where(col(ApiKey.total_spent) == 0) .where(col(ApiKey.total_requests) == 0) - .where(col(ApiKey.parent_key_hash).is_(None)) .where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)) .where(~settleable_invoice) - .where(~has_children) ) # Unlink transactions rather than cascade-deleting them, so the financial @@ -542,14 +515,6 @@ class LightningInvoice(SQLModel, table=True): # type: ignore ) expires_at: int = Field(description="Unix timestamp when invoice expires") paid_at: int | None = Field(default=None, description="Unix timestamp when paid") - balance_limit: int | None = Field( - default=None, - description="Max spendable msats for the created key", - ) - balance_limit_reset: str | None = Field( - default=None, - description="Reset policy for balance limit (daily, weekly, monthly)", - ) validity_date: int | None = Field( default=None, description="Unix timestamp after which the created key expires", diff --git a/routstr/core/main.py b/routstr/core/main.py index b89fa64d..07f11aa2 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -14,7 +14,6 @@ from starlette.types import Scope from ..auth import ( periodic_dead_key_prune, - periodic_key_reset, periodic_stale_reservation_sweep, ) from ..balance import balance_router, deprecated_wallet_router @@ -149,7 +148,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: analytics_task = asyncio.create_task(publish_usage_analytics()) if global_settings.providers_refresh_interval_seconds > 0: providers_task = asyncio.create_task(providers_cache_refresher()) - key_reset_task = asyncio.create_task(periodic_key_reset()) stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) @@ -308,7 +306,6 @@ async def info() -> dict: "mints": global_settings.cashu_mints, "http_url": global_settings.http_url, "onion_url": global_settings.onion_url, - "child_key_cost_msats": global_settings.child_key_cost, } diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 7190f058..014a5795 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -77,7 +77,6 @@ class Settings(BaseSettings): exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") - child_key_cost: int = Field(default=0, env="CHILD_KEY_COST") # Minimum per-request charge in millisatoshis when model pricing is free/zero min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") reset_reserved_balance_on_startup: bool = Field( diff --git a/routstr/lightning.py b/routstr/lightning.py index 13c6b321..e45da084 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -76,8 +76,6 @@ class _InvoiceSettlement: purpose: str api_key_hash: str | None mint_url: str | None - balance_limit: int | None - balance_limit_reset: str | None validity_date: int | None @classmethod @@ -89,8 +87,6 @@ class _InvoiceSettlement: purpose=invoice.purpose, api_key_hash=invoice.api_key_hash, mint_url=invoice.mint_url, - balance_limit=invoice.balance_limit, - balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, ) @@ -114,8 +110,6 @@ class InvoiceCreateRequest(BaseModel): default=None, description="Deprecated: legacy field for topup. Prefer Authorization header.", ) - balance_limit: int | None = Field(default=None) - balance_limit_reset: str | None = Field(default=None) validity_date: int | None = Field(default=None) @@ -312,8 +306,6 @@ async def create_invoice( api_key_hash=api_key_token[3:] if api_key_token else None, purpose=request.purpose, mint_url=mint_url, - balance_limit=request.balance_limit, - balance_limit_reset=request.balance_limit_reset, validity_date=request.validity_date, expires_at=expires_at, ) @@ -697,8 +689,6 @@ async def _create_api_key_record( balance=invoice.amount_sats * 1000, refund_currency="sat", refund_mint_url=mint_url, - balance_limit=invoice.balance_limit, - balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, ) session.add(api_key) diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index d6b7e7d2..839cb079 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -18,7 +18,6 @@ from ..auth import ( _claim_reservation_for_charge, _stop_reservation_heartbeat, _validate_reservation_snapshot, - get_billing_key, get_reservation_snapshot, payments_logger, release_reservation, @@ -527,9 +526,8 @@ async def finalize_ehbp_actual_cost_payment( if not await _claim_reservation_for_charge(reservation, session): return 0 reserved_cost_for_model = reservation.reserved_msats - billing_key = await get_billing_key(key, session) key_hash = key.hashed_key - billing_key_hash = billing_key.hashed_key + billing_key_hash = key_hash total_cost_msats = max( 0, int(cost_info.get("total_msats", reserved_cost_for_model)) ) @@ -538,7 +536,6 @@ async def finalize_ehbp_actual_cost_payment( charged = await _charge_reservation_rows( session, billing_key_hash=billing_key_hash, - key_hash=key_hash, reserved_msats=reserved_cost_for_model, charge_msats=total_cost_msats, ) @@ -558,9 +555,7 @@ async def finalize_ehbp_actual_cost_payment( await session.commit() await _stop_reservation_heartbeat(reservation.release_id) - await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) + await session.refresh(key) if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0: fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100) @@ -577,15 +572,15 @@ async def finalize_ehbp_actual_cost_payment( extra={ "event": "finalize", "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "billing_key_hash": key.hashed_key[:8] + "...", "model": model_id, "cost_reserved": reserved_cost_for_model, "cost_charged": total_cost_msats, "input_tokens": cost_info.get("input_tokens", 0), "output_tokens": cost_info.get("output_tokens", 0), - "balance": billing_key.balance, - "reserved_balance": billing_key.reserved_balance, - "total_spent": billing_key.total_spent, + "balance": key.balance, + "reserved_balance": key.reserved_balance, + "total_spent": key.total_spent, "finalize_type": "ehbp_usage", "finalized_at": now, }, diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py deleted file mode 100644 index 5226d357..00000000 --- a/tests/integration/test_child_keys.py +++ /dev/null @@ -1,190 +0,0 @@ -import asyncio -import secrets -from typing import Any - -import pytest -from fastapi import HTTPException -from sqlmodel.ext.asyncio.session import AsyncSession - -from routstr.auth import adjust_payment_for_tokens, pay_for_request -from routstr.balance import ChildKeyRequest, create_child_key -from routstr.core.db import ApiKey, create_session -from routstr.core.settings import settings - - -@pytest.mark.asyncio -async def test_child_key_flow(integration_session: AsyncSession) -> None: - # 1. Create a parent key with balance - parent_raw = "parent_test_key_" + secrets.token_hex(4) - parent_key = ApiKey( - hashed_key=parent_raw, - balance=10000, # 10 sats - ) - integration_session.add(parent_key) - await integration_session.commit() - await integration_session.refresh(parent_key) - - # Mock settings - settings.child_key_cost = 1000 # 1 sat - - # 2. Call create_child_key - result = await create_child_key( - ChildKeyRequest(count=1), parent_key, integration_session - ) - - assert "api_keys" in result - assert result["cost_msats"] == 1000 - assert result["parent_balance"] == 9000 - - child_key_raw = result["api_keys"][0][3:] # remove sk- - - # 3. Verify child key exists in DB - child_key_db = await integration_session.get(ApiKey, child_key_raw) - assert child_key_db is not None - assert child_key_db.parent_key_hash == parent_key.hashed_key - assert child_key_db.balance == 0 - - # 4. Test payment with child key - cost = 500 - await pay_for_request(child_key_db, cost, integration_session) - - # Refresh keys - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key_db) - - # Parent should be charged - assert parent_key.reserved_balance == 500 - assert parent_key.total_requests == 1 - - # Child should have total_requests incremented - assert child_key_db.total_requests == 1 - - # 5. Test adjustment - response_data = {"model": "test-model", "usage": {"total_tokens": 10}} - - # Mock calculate_cost - import routstr.auth - from routstr.payment.cost_calculation import CostData - - async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData: - return CostData( - base_msats=0, input_msats=200, output_msats=200, total_msats=400 - ) - - # Patch calculate_cost - original_calculate_cost = routstr.auth.calculate_cost - routstr.auth.calculate_cost = mock_calculate_cost - - try: - adjustment = await adjust_payment_for_tokens( - child_key_db, response_data, integration_session, 500, None, None - ) - assert adjustment["total_msats"] == 400 - - # Refresh keys - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key_db) - - # Parent should have updated balance and total_spent - assert parent_key.reserved_balance == 0 - assert parent_key.balance == 9000 - 400 - assert ( - parent_key.total_spent == 1400 - ) # 1000 for child key creation + 400 for request - - # Child should also have total_spent updated - assert child_key_db.total_spent == 400 - - finally: - routstr.auth.calculate_cost = original_calculate_cost - - -@pytest.mark.asyncio -async def test_child_key_insufficient_balance( - integration_session: AsyncSession, -) -> None: - parent_key = ApiKey( - hashed_key="poor_parent_" + secrets.token_hex(4), - balance=500, - ) - integration_session.add(parent_key) - await integration_session.commit() - await integration_session.refresh(parent_key) - - settings.child_key_cost = 1000 - - with pytest.raises(HTTPException) as exc: - await create_child_key( - ChildKeyRequest(count=1), parent_key, integration_session - ) - assert exc.value.status_code == 402 - - -@pytest.mark.asyncio -async def test_concurrent_child_key_creation_is_atomic( - patched_db_engine: None, -) -> None: - """Two concurrent create_child_key() calls with balance for exactly one must - result in exactly one success and one 402, with the parent balance deducted - only once.""" - child_key_cost = 1000 - settings.child_key_cost = child_key_cost - - parent_hash = f"parent_concurrent_{secrets.token_hex(8)}" - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=child_key_cost) - session.add(parent) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_parent = await session.get(ApiKey, parent_hash) - assert fresh_parent is not None - try: - await create_child_key(ChildKeyRequest(count=1), fresh_parent, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final = await session.get(ApiKey, parent_hash) - assert final is not None - - assert final.balance == 0, ( - f"Balance should be fully deducted once: expected 0, got {final.balance}" - ) - assert final.total_spent == child_key_cost, ( - f"total_spent should equal one deduction: expected {child_key_cost}, " - f"got {final.total_spent}" - ) - - -@pytest.mark.asyncio -async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None: - parent_key = ApiKey( - hashed_key="parent_" + secrets.token_hex(4), - balance=10000, - ) - child_key = ApiKey( - hashed_key="child_" + secrets.token_hex(4), - balance=0, - parent_key_hash=parent_key.hashed_key, - ) - integration_session.add(parent_key) - integration_session.add(child_key) - await integration_session.commit() - await integration_session.refresh(child_key) - - with pytest.raises(HTTPException) as exc: - await create_child_key(ChildKeyRequest(count=1), child_key, integration_session) - assert exc.value.status_code == 400 - assert "Cannot create a child key for another child key" in str(exc.value.detail) diff --git a/tests/integration/test_child_keys_api.py b/tests/integration/test_child_keys_api.py deleted file mode 100644 index 1adb9496..00000000 --- a/tests/integration/test_child_keys_api.py +++ /dev/null @@ -1,102 +0,0 @@ -from typing import Any - -import pytest -from httpx import AsyncClient - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_wallet_info_returns_child_keys( - integration_client: AsyncClient, - authenticated_client: AsyncClient, - integration_session: Any, -) -> None: - """Test that GET /v1/wallet/info returns child keys for a parent key""" - - # 1. Get parent info to find its hashed_key - response = await authenticated_client.get("/v1/wallet/info") - assert response.status_code == 200 - parent_data = response.json() - parent_data["api_key"] - - # 2. Create child keys for this parent - # We need to use the parent's authentication for this - child_payload = {"count": 2, "balance_limit": 1000, "balance_limit_reset": "daily"} - create_response = await authenticated_client.post( - "/v1/wallet/child-key", json=child_payload - ) - assert create_response.status_code == 200 - create_data = create_response.json() - child_keys = create_data["api_keys"] - assert len(child_keys) == 2 - - # 3. Call /info again and check for child_keys - info_response = await authenticated_client.get("/v1/wallet/info") - assert info_response.status_code == 200 - info_data = info_response.json() - - assert "child_keys" in info_data - assert len(info_data["child_keys"]) == 2 - - # Verify child key details - for ck in info_data["child_keys"]: - assert ck["api_key"] in child_keys - assert ck["balance_limit"] == 1000 - assert ck["balance_limit_reset"] == "daily" - assert "total_spent" in ck - assert "total_requests" in ck - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_wallet_info_child_key_no_child_keys( - integration_client: AsyncClient, - authenticated_client: AsyncClient, - integration_session: Any, -) -> None: - """Test that GET /v1/wallet/info for a child key does NOT return child_keys""" - - # 1. Create a child key - child_payload = {"count": 1} - create_response = await authenticated_client.post( - "/v1/wallet/child-key", json=child_payload - ) - assert create_response.status_code == 200 - child_key = create_response.json()["api_keys"][0] - - # 2. Use the child key to get its info - integration_client.headers["Authorization"] = f"Bearer {child_key}" - info_response = await integration_client.get("/v1/wallet/info") - assert info_response.status_code == 200 - info_data = info_response.json() - parent_key = authenticated_client._test_api_key # type: ignore[attr-defined] - parent_key_hash = parent_key.removeprefix("sk-") - - assert info_data["is_child"] is True - assert "child_keys" not in info_data - assert "parent_key" not in info_data - assert info_data["parent_key_preview"] == parent_key_hash[:8] + "..." - assert info_data["parent_key_preview"] not in {parent_key, parent_key_hash} - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_account_info_root_returns_child_keys( - authenticated_client: AsyncClient, -) -> None: - """Test that GET / returns child keys for a parent key (root endpoint)""" - - # 1. Create a child key - child_payload = {"count": 1} - await authenticated_client.post("/v1/wallet/child-key", json=child_payload) - - # 2. Call root endpoint /v1/balance/ - # Note: routstr/balance.py defines router = APIRouter() - # and it is included in balance_router with prefix /v1/balance - # The endpoint is @router.get("/") - response = await authenticated_client.get("/v1/balance/") - assert response.status_code == 200 - data = response.json() - - assert "child_keys" in data - assert len(data["child_keys"]) >= 1 diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index a0c31401..d67dc32b 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -17,7 +17,7 @@ from httpx import AsyncClient from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.db import ApiKey, ReservationRelease +from routstr.core.db import ReservationRelease from routstr.payment.models import Architecture, Model, Pricing from routstr.proxy import refresh_model_maps from routstr.upstream.base import BaseUpstreamProvider @@ -531,103 +531,6 @@ async def test_failover_beyond_balance_envelope_is_rejected( # fallback must be rejected before its upstream is ever contacted. assert response.status_code == 402 assert [r.url.host for r in sent_requests] == ["cheap.example.com"] - - -@pytest.fixture -async def three_candidate_child_maps( - patched_db_engine: None, -) -> AsyncGenerator[None, None]: - """Second candidate cannot fit the child limit; third restores and serves.""" - first = _StaticProvider( - CHEAP_BASE_URL, - "key-first", - 1.0, - _make_model("dual-model", 0.001, 0.002, max_cost=50.0), - ) - too_large = _StaticProvider( - EXPENSIVE_BASE_URL, - "key-too-large", - 1.0, - _make_model("dual-model", 0.002, 0.003, max_cost=100.0), - ) - third = _StaticProvider( - THIRD_BASE_URL, - "key-third", - 1.0, - _make_model("dual-model", 0.003, 0.004, max_cost=50.0), - ) - async for _ in _install_providers([first, too_large, third]): - yield - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring( - authenticated_client: AsyncClient, - three_candidate_child_maps: None, - integration_session: AsyncSession, -) -> None: - """A failed child guard cannot leak its parent update into restoration.""" - key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined] - child = await integration_session.get(ApiKey, key_hash) - assert child is not None - parent = ApiKey(hashed_key="failover-parent", balance=10_000_000) - child.parent_key_hash = parent.hashed_key - child.balance_limit = 75_000 - integration_session.add(parent) - integration_session.add(child) - await integration_session.commit() - - sent_requests: list[httpx.Request] = [] - - async def fake_transport( - request: httpx.Request, *args: Any, **kwargs: Any - ) -> httpx.Response: - sent_requests.append(request) - return _upstream_response(request) - - with ( - patch( - "httpx.AsyncHTTPTransport.handle_async_request", - side_effect=fake_transport, - ), - patch( - "routstr.payment.cost_calculation.sats_usd_price", - return_value=0.0005, - ), - ): - response = await authenticated_client.post( - "/v1/chat/completions", - json={ - "model": "dual-model", - "messages": [{"role": "user", "content": "hello"}], - }, - ) - - assert response.status_code == 200 - # The 100-sat candidate is rejected before forwarding; the third serves. - assert [request.url.host for request in sent_requests] == [ - "cheap.example.com", - "third.example.com", - ] - - await integration_session.refresh(parent) - await integration_session.refresh(child) - assert parent.reserved_balance == 0 - assert child.reserved_balance == 0 - assert parent.total_spent == response.json()["cost"]["total_msats"] - - records = ( - await integration_session.exec( - select(ReservationRelease).where(ReservationRelease.key_hash == key_hash) - ) - ).all() - assert len(records) == 2 - assert sorted(record.status for record in records) == ["charged", "released"] - assert len({record.reserved_msats for record in records}) == 1 - assert all(record.status != "active" for record in records) - - @pytest.fixture async def raised_envelope_provider_maps( patched_db_engine: None, diff --git a/tests/integration/test_key_logic.py b/tests/integration/test_key_logic.py index 79d2c9fe..0a0fd573 100644 --- a/tests/integration/test_key_logic.py +++ b/tests/integration/test_key_logic.py @@ -1,14 +1,10 @@ -import asyncio import time -from datetime import datetime, timedelta import pytest -from fastapi import HTTPException -from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession from routstr.auth import pay_for_request -from routstr.core.db import ApiKey, create_session +from routstr.core.db import ApiKey @pytest.mark.asyncio @@ -25,367 +21,6 @@ async def test_key_validity_date(integration_session: AsyncSession) -> None: assert "expired" in str(excinfo.value).lower() -@pytest.mark.asyncio -async def test_key_balance_limit(integration_session: AsyncSession) -> None: - # 1. Create a key with a balance limit - key = ApiKey( - hashed_key="limited_key", balance=10000, balance_limit=500, total_spent=450 - ) - integration_session.add(key) - await integration_session.commit() - - # 2. Try to pay for a request that exceeds the limit - with pytest.raises(Exception) as excinfo: - await pay_for_request(key, 100, integration_session) - assert "limit exceeded" in str(excinfo.value).lower() - - # 3. Try to pay for a request that fits - await pay_for_request(key, 50, integration_session) - await integration_session.refresh(key) - # Note: total_spent is updated in adjust_payment_for_tokens, - # but pay_for_request checks it. - # In our current logic, pay_for_request checks (total_spent + cost) > balance_limit. - - -@pytest.mark.asyncio -async def test_key_daily_reset_policy(integration_session: AsyncSession) -> None: - # 1. Create a key with a daily reset policy and old reset date - yesterday = int((datetime.now() - timedelta(days=1)).timestamp()) - key = ApiKey( - hashed_key="daily_reset_key", - balance=10000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=900, - ) - integration_session.add(key) - await integration_session.commit() - - # 2. Pay for a request - should trigger reset first because it's a new day - # Request is 200, total_spent is 900. 900+200 > 1000, - # but reset should happen making total_spent 0, then 0+200 < 1000. - await pay_for_request(key, 200, integration_session) - - await integration_session.refresh(key) - assert key.total_spent == 0 # Reset in pay_for_request happens before charging - # Wait, the charging logic in pay_for_request increments parent/billing_key's total_requests, - # but total_spent is updated in adjust_payment_for_tokens. - # However, the reset logic sets total_spent to 0. - assert key.balance_limit_reset_date is not None - assert key.balance_limit_reset_date > yesterday - - -@pytest.mark.asyncio -async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None: - # 1. Create multiple keys needing reset - yesterday = int((datetime.now() - timedelta(days=1)).timestamp()) - key1 = ApiKey( - hashed_key="job_reset_key_1", - balance=1000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=500, - ) - key2 = ApiKey( - hashed_key="job_reset_key_2", - balance=1000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=800, - ) - integration_session.add(key1) - integration_session.add(key2) - await integration_session.commit() - - # 2. Run the periodic reset logic manually (mocking the background task loop) - # We can't easily run the actual loop because it has a sleep, - # but we can test the logic inside. - - # Implementation of periodic_key_reset logic for testing: - stmt = select(ApiKey).where(ApiKey.balance_limit_reset != None) # noqa: E711 - keys = (await integration_session.exec(stmt)).all() - now = int(time.time()) - for k in keys: - if k.hashed_key in ["job_reset_key_1", "job_reset_key_2"]: - k.total_spent = 0 - k.balance_limit_reset_date = now - integration_session.add(k) - await integration_session.commit() - - # 3. Verify resets - await integration_session.refresh(key1) - await integration_session.refresh(key2) - assert key1.total_spent == 0 - assert key2.total_spent == 0 - - -@pytest.mark.asyncio -async def test_balance_limit_enforced_atomically_under_concurrency( - patched_db_engine: None, -) -> None: - parent_hash = "parent_limit_atomic" - child_hash = "child_limit_atomic" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=cost, - total_spent=0, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_child = await session.get(ApiKey, child_hash) - assert final_child is not None - - assert final_child.reserved_balance == cost, ( - f"Child reserved_balance should equal one reservation, " - f"got {final_child.reserved_balance}" - ) - - -@pytest.mark.asyncio -async def test_parallel_payments_with_parent_and_child_key( - patched_db_engine: None, -) -> None: - parent_hash = "parent_parallel_mixed" - child_hash = "child_parallel_mixed" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=2 * cost, - ) - session.add(parent) - session.add(child) - await session.commit() - - async def attempt(key_hash: str) -> str: - async with create_session() as session: - fresh_key = await session.get(ApiKey, key_hash) - assert fresh_key is not None - try: - await pay_for_request(fresh_key, cost, session) - return "success" - except HTTPException as exc: - assert exc.status_code == 402 - return "blocked" - - results = await asyncio.gather(attempt(parent_hash), attempt(child_hash)) - - assert results == ["success", "success"], ( - f"Both parent and child payments should succeed, got: {results}" - ) - - async with create_session() as session: - final_parent = await session.get(ApiKey, parent_hash) - final_child = await session.get(ApiKey, child_hash) - assert final_parent is not None - assert final_child is not None - - # Both requests bill the parent; only the child request reserves on the child. - assert final_parent.reserved_balance == 2 * cost - assert final_parent.total_requests == 2 - assert final_child.reserved_balance == cost - assert final_child.total_requests == 1 - - -@pytest.mark.asyncio -async def test_balance_limit_with_existing_total_spent_under_concurrency( - patched_db_engine: None, -) -> None: - parent_hash = "parent_total_spent" - child_hash = "child_total_spent" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - # 700 already spent against a 1000 limit: only one more 300 request fits. - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - total_spent=700, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_child = await session.get(ApiKey, child_hash) - assert final_child is not None - - assert final_child.reserved_balance == cost - assert final_child.total_spent == 700 - - -@pytest.mark.asyncio -async def test_balance_limit_with_existing_reserved_balance( - patched_db_engine: None, -) -> None: - parent_hash = "parent_reserved_set" - blocked_hash = "child_reserved_blocked" - allowed_hash = "child_reserved_allowed" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - # 800 already reserved against a 1000 limit: another 300 must be rejected. - blocked_child = ApiKey( - hashed_key=blocked_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - reserved_balance=800, - ) - # 500 reserved against a 1000 limit: another 300 still fits. - allowed_child = ApiKey( - hashed_key=allowed_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - reserved_balance=500, - ) - session.add(parent) - session.add(blocked_child) - session.add(allowed_child) - await session.commit() - - async with create_session() as session: - fresh_blocked = await session.get(ApiKey, blocked_hash) - assert fresh_blocked is not None - with pytest.raises(HTTPException) as exc_info: - await pay_for_request(fresh_blocked, cost, session) - assert exc_info.value.status_code == 402 - - async with create_session() as session: - fresh_allowed = await session.get(ApiKey, allowed_hash) - assert fresh_allowed is not None - await pay_for_request(fresh_allowed, cost, session) - - async with create_session() as session: - final_blocked = await session.get(ApiKey, blocked_hash) - final_allowed = await session.get(ApiKey, allowed_hash) - final_parent = await session.get(ApiKey, parent_hash) - assert final_blocked is not None - assert final_allowed is not None - assert final_parent is not None - - assert final_blocked.reserved_balance == 800, "Rejected request must not reserve" - assert final_blocked.total_requests == 0 - assert final_allowed.reserved_balance == 500 + cost - assert final_allowed.total_requests == 1 - # Only the allowed request should have billed the parent. - assert final_parent.reserved_balance == cost - assert final_parent.total_requests == 1 - - -@pytest.mark.asyncio -async def test_child_reservation_discarded_when_parent_balance_depleted( - patched_db_engine: None, -) -> None: - parent_hash = "parent_depleted" - child_hash = "child_depleted" - cost = 300 - - async with create_session() as session: - # Parent can only afford one request; child has no balance_limit. - parent = ApiKey(hashed_key=parent_hash, balance=cost) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_parent = await session.get(ApiKey, parent_hash) - final_child = await session.get(ApiKey, child_hash) - assert final_parent is not None - assert final_child is not None - - assert final_parent.reserved_balance == cost - # The failed request must not leave a committed reservation on the child. - assert final_child.reserved_balance == cost, ( - f"Child reserved_balance should reflect only the successful request, " - f"got {final_child.reserved_balance}" - ) - assert final_child.total_requests == 1 - - @pytest.mark.asyncio async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None: # This requires mocking the router call or testing the logic in balance.py diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index 1a6d94b9..7370ebe4 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -1,10 +1,10 @@ """Integration tests for Lightning invoice key constraint fields. Covers two things: -- The three constraint fields (balance_limit, balance_limit_reset, validity_date) - are persisted on LightningInvoice and survive a DB round-trip. -- The production-path API-key record helper propagates those fields to the - created ApiKey, so the constraints are actually enforced when the key is used. +- The validity_date constraint field is persisted on LightningInvoice and + survives a DB round-trip. +- The production-path API-key record helper propagates it to the created + ApiKey, so the constraint is actually enforced when the key is used. """ from __future__ import annotations @@ -67,32 +67,6 @@ def mock_wallet_mint() -> object: # --------------------------------------------------------------------------- -@pytest.mark.asyncio -async def test_invoice_persists_balance_limit( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=5000) - integration_session.add(invoice) - await integration_session.commit() - - stored = await integration_session.get(LightningInvoice, invoice.id) - assert stored is not None - assert stored.balance_limit == 5000 - - -@pytest.mark.asyncio -async def test_invoice_persists_balance_limit_reset( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=5000, balance_limit_reset="daily") - integration_session.add(invoice) - await integration_session.commit() - - stored = await integration_session.get(LightningInvoice, invoice.id) - assert stored is not None - assert stored.balance_limit_reset == "daily" - - @pytest.mark.asyncio async def test_invoice_persists_validity_date( integration_session: AsyncSession, @@ -112,38 +86,6 @@ async def test_invoice_persists_validity_date( # --------------------------------------------------------------------------- -@pytest.mark.asyncio -async def test_created_key_receives_balance_limit( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=8000) - integration_session.add(invoice) - await integration_session.flush() - - api_key = await _create_api_key_record(invoice, integration_session) - await integration_session.commit() - - stored_key = await integration_session.get(ApiKey, api_key.hashed_key) - assert stored_key is not None - assert stored_key.balance_limit == 8000 - - -@pytest.mark.asyncio -async def test_created_key_receives_balance_limit_reset( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=8000, balance_limit_reset="monthly") - integration_session.add(invoice) - await integration_session.flush() - - api_key = await _create_api_key_record(invoice, integration_session) - await integration_session.commit() - - stored_key = await integration_session.get(ApiKey, api_key.hashed_key) - assert stored_key is not None - assert stored_key.balance_limit_reset == "monthly" - - @pytest.mark.asyncio async def test_created_key_receives_validity_date( integration_session: AsyncSession, @@ -422,8 +364,6 @@ async def test_created_key_without_constraints_has_none_fields( stored_key = await integration_session.get(ApiKey, api_key.hashed_key) assert stored_key is not None - assert stored_key.balance_limit is None - assert stored_key.balance_limit_reset is None assert stored_key.validity_date is None diff --git a/tests/integration/test_payment_invariants.py b/tests/integration/test_payment_invariants.py index ca7bf892..cb10b80a 100644 --- a/tests/integration/test_payment_invariants.py +++ b/tests/integration/test_payment_invariants.py @@ -45,13 +45,7 @@ def _response() -> dict: } -async def _new_key( - session: AsyncSession, - balance: int, - *, - parent_key_hash: str | None = None, - balance_limit: int | None = None, -) -> str: +async def _new_key(session: AsyncSession, balance: int) -> str: key_hash = f"test_inv_{uuid.uuid4().hex}" session.add( ApiKey( @@ -60,8 +54,6 @@ async def _new_key( reserved_balance=0, total_spent=0, total_requests=0, - parent_key_hash=parent_key_hash, - balance_limit=balance_limit, ) ) await session.commit() @@ -458,92 +450,3 @@ async def test_release_after_charge_does_not_credit_the_user_back( 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}" - ) diff --git a/tests/integration/test_prune_dead_api_keys.py b/tests/integration/test_prune_dead_api_keys.py index de2b9d44..d0987af9 100644 --- a/tests/integration/test_prune_dead_api_keys.py +++ b/tests/integration/test_prune_dead_api_keys.py @@ -98,34 +98,6 @@ async def test_used_key_never_pruned(patched_db_engine: None) -> None: assert await _exists(k.hashed_key) -@pytest.mark.asyncio -async def test_parent_and_child_keys_are_not_pruned( - patched_db_engine: None, -) -> None: - """Pruning must not orphan child keys or delete valid children.""" - parent = _dead_key(LONG_AGO) - child = ApiKey( - hashed_key=f"child_{uuid.uuid4().hex}", - balance=0, - reserved_balance=0, - total_spent=0, - total_requests=0, - created_at=LONG_AGO, - parent_key_hash=parent.hashed_key, - ) - async with create_session() as session: - session.add(parent) - session.add(child) - await session.commit() - - async with create_session() as session: - pruned = await prune_dead_api_keys(session, OLD) - - assert pruned == 0 - assert await _exists(parent.hashed_key) - assert await _exists(child.hashed_key) - - @pytest.mark.asyncio @pytest.mark.parametrize( ("status", "expires_at"), diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py index bf23b0be..49e290cb 100644 --- a/tests/integration/test_reserved_balance_negative.py +++ b/tests/integration/test_reserved_balance_negative.py @@ -170,48 +170,6 @@ async def test_revert_with_zero_reserved_balance_repairs_terminally( assert updated.total_requests == 0 -@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_row = await integration_session.get(ApiKey, snapshot.billing_key_hash) - child_row = await integration_session.get(ApiKey, snapshot.key_hash) - release = await integration_session.get(ReservationRelease, snapshot.release_id) - assert parent_row is not None and child_row is not None - assert (parent_row.total_requests, child_row.total_requests) == (0, 0) - assert (parent_row.reserved_balance, child_row.reserved_balance) == (500, 0) - assert release is not None and release.status == "released" - - @pytest.mark.asyncio async def test_revert_with_sufficient_reserved_balance_succeeds( integration_session: AsyncSession, @@ -375,63 +333,3 @@ async def test_sequential_reverts_never_go_negative( assert test_key.reserved_balance >= 0, ( f"Reserved balance went negative: {test_key.reserved_balance}" ) - - -@pytest.mark.asyncio -async def test_child_key_revert_floor_guard( - integration_session: AsyncSession, -) -> None: - """Test that child key reserved_balance also has floor guard on revert.""" - from routstr.auth import ( - get_reservation_snapshot, - pay_for_request, - revert_pay_for_request, - ) - - parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}" - child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}" - - parent_key = ApiKey( - hashed_key=parent_key_hash, - balance=10000, - reserved_balance=0, - total_requests=2, - ) - child_key = ApiKey( - hashed_key=child_key_hash, - balance=0, - reserved_balance=0, - total_requests=2, - parent_key_hash=parent_key_hash, - ) - integration_session.add(parent_key) - integration_session.add(child_key) - await integration_session.commit() - await pay_for_request(child_key, 500, integration_session) - snapshot = await get_reservation_snapshot(child_key, integration_session) - - # First revert succeeds - result1 = await revert_pay_for_request( - child_key, integration_session, 500, snapshot - ) - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key) - - assert result1 is True - assert parent_key.reserved_balance == 0 - assert child_key.reserved_balance == 0 - - # Second revert is a no-op for both parent and child - result2 = await revert_pay_for_request( - child_key, integration_session, 500, snapshot - ) - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key) - - assert result2 is False - assert parent_key.reserved_balance == 0, ( - f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}" - ) - assert child_key.reserved_balance == 0, ( - f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}" - ) diff --git a/tests/integration/test_temporary_balances_api.py b/tests/integration/test_temporary_balances_api.py index 5439dcba..a6a2d5e5 100644 --- a/tests/integration/test_temporary_balances_api.py +++ b/tests/integration/test_temporary_balances_api.py @@ -26,7 +26,6 @@ async def _add_key( total_spent: int = 0, total_requests: int = 0, created_at: int | None = None, - parent_key_hash: str | None = None, refund_address: str | None = None, ) -> ApiKey: key = ApiKey( @@ -35,7 +34,6 @@ async def _add_key( reserved_balance=reserved_balance, total_spent=total_spent, total_requests=total_requests, - parent_key_hash=parent_key_hash, refund_address=refund_address, ) key.created_at = created_at @@ -125,7 +123,7 @@ async def test_temporary_balances_pagination( @pytest.mark.integration @pytest.mark.asyncio -async def test_temporary_balances_totals_exclude_child_balance( +async def test_temporary_balances_totals( integration_client: httpx.AsyncClient, integration_session: AsyncSession, ) -> None: @@ -137,17 +135,6 @@ async def test_temporary_balances_totals_exclude_child_balance( total_requests=3, created_at=1000, ) - # Child draws from parent's balance, so its balance must NOT be summed, - # but its spent/requests still count. - await _add_key( - integration_session, - "child", - balance=0, - total_spent=200, - total_requests=7, - created_at=1001, - parent_key_hash="parent", - ) response = await integration_client.get( "/admin/api/temporary-balances", headers=_admin_headers() @@ -155,43 +142,8 @@ async def test_temporary_balances_totals_exclude_child_balance( totals = response.json()["totals"] assert totals["total_balance"] == 5000 - assert totals["total_spent"] == 300 - 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 + assert totals["total_spent"] == 100 + assert totals["total_requests"] == 3 @pytest.mark.integration diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 93ac02b8..3812b722 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -208,7 +208,6 @@ def _make_api_key( refund_currency: str | None = "sat", refund_mint_url: str | None = "https://mint.example.com", refund_address: str | None = None, - parent_key_hash: str | None = None, ) -> ApiKey: key = ApiKey(hashed_key="testhash") key.balance = balance @@ -216,7 +215,6 @@ def _make_api_key( key.refund_currency = refund_currency key.refund_mint_url = refund_mint_url key.refund_address = refund_address - key.parent_key_hash = parent_key_hash key.total_spent = 0 key.total_requests = 0 return key @@ -306,7 +304,6 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No session.commit = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store, patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), @@ -341,7 +338,6 @@ async def test_apikey_refund_logs_token() -> None: session.commit = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), @@ -370,7 +366,6 @@ async def test_apikey_refund_log_includes_path() -> None: session.commit = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), @@ -409,7 +404,6 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None: mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted") with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", mock_send_token), patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), @@ -470,7 +464,6 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: session.commit = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.send_token", AsyncMock(side_effect=MintConnectionError("raw mint outage detail")), @@ -508,7 +501,6 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None: session.commit = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))), patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), @@ -609,7 +601,6 @@ async def test_topup_mint_unreachable_returns_503(error: Exception) -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -635,7 +626,6 @@ async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible error = SourceMintConnectionError("Issuing Cashu mint is unreachable") with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -660,7 +650,6 @@ async def test_topup_already_spent_still_returns_400() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=ValueError("Token already spent")), @@ -688,7 +677,6 @@ async def test_topup_zero_value_returns_400_zero_value_message() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock( @@ -720,7 +708,6 @@ async def test_topup_token_consumed_returns_500() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=TokenConsumedError("credit failed")), @@ -786,7 +773,6 @@ async def test_topup_fee_and_swap_failures_return_422( session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -811,7 +797,6 @@ async def test_topup_unexpected_non_valueerror_returns_500() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=RuntimeError("db exploded")), diff --git a/tests/unit/test_ehbp_finalize_payment.py b/tests/unit/test_ehbp_finalize_payment.py index 9ed9fb5c..1ed2dda4 100644 --- a/tests/unit/test_ehbp_finalize_payment.py +++ b/tests/unit/test_ehbp_finalize_payment.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.pool import StaticPool -from sqlmodel import SQLModel, col, select, update +from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession import routstr.auth as auth_module @@ -107,19 +107,17 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve @pytest.mark.asyncio -async def test_unmeasured_ehbp_releases_parent_and_child_reservation( +async def test_unmeasured_ehbp_releases_reservation( session: AsyncSession, ) -> None: - parent = ApiKey(hashed_key="ehbp-parent", balance=10_000) - child = ApiKey(hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent") - session.add(parent) - session.add(child) + key = ApiKey(hashed_key="ehbp-key", balance=10_000) + session.add(key) await session.commit() - await pay_for_request(child, 3_000, session) - reservation = await get_reservation_snapshot(child, session) + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) charged = await finalize_ehbp_max_cost_payment( - child, + key, session, max_cost_for_model=3_000, model_id="tinfoil/model", @@ -127,18 +125,12 @@ async def test_unmeasured_ehbp_releases_parent_and_child_reservation( ) assert charged == 0 - updated_parent = await _api_key(session, "ehbp-parent") - updated_child = await _api_key(session, "ehbp-child") - assert updated_parent is not None - assert updated_child is not None - assert updated_parent.balance == 10_000 - assert updated_parent.reserved_balance == 0 - assert updated_parent.reserved_at is None - assert updated_parent.total_spent == 0 - assert updated_child.balance == 0 - assert updated_child.reserved_balance == 0 - assert updated_child.reserved_at is None - assert updated_child.total_spent == 0 + updated = await _api_key(session, "ehbp-key") + assert updated is not None + assert updated.balance == 10_000 + assert updated.reserved_balance == 0 + assert updated.reserved_at is None + assert updated.total_spent == 0 @pytest.mark.asyncio @@ -182,21 +174,15 @@ async def test_unmeasured_ehbp_release_is_safe_when_charge_update_would_fail( session: AsyncSession, monkeypatch: pytest.MonkeyPatch, ) -> None: - parent = ApiKey(hashed_key="ehbp-rollback-parent", balance=10_000) - child = ApiKey( - hashed_key="ehbp-missing-child", - balance=0, - parent_key_hash="ehbp-rollback-parent", - ) - session.add(parent) - session.add(child) + key = ApiKey(hashed_key="ehbp-rollback-key", balance=10_000) + session.add(key) await session.commit() - await pay_for_request(child, 3_000, session) - reservation = await get_reservation_snapshot(child, session) - _fail_nth_api_key_update(session, monkeypatch, target_update=2) + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + _fail_nth_api_key_update(session, monkeypatch, target_update=1) charged = await finalize_ehbp_max_cost_payment( - child, + key, session, max_cost_for_model=3_000, model_id="tinfoil/model", @@ -204,68 +190,13 @@ async def test_unmeasured_ehbp_release_is_safe_when_charge_update_would_fail( ) assert charged == 0 - updated_parent = await _api_key(session, "ehbp-rollback-parent") - assert updated_parent is not None - assert updated_parent.balance == 10_000 + updated = await _api_key(session, "ehbp-rollback-key") + assert updated is not None + assert updated.balance == 10_000 # The injected partial-update failure rolls aggregate subtraction back; # terminal fencing prevents a charge or retry from consuming those funds. - assert updated_parent.reserved_balance == 3_000 - assert updated_parent.total_spent == 0 - updated_child = await _api_key(session, "ehbp-missing-child") - assert updated_child is not None - assert updated_child.reserved_balance == 3_000 - assert updated_child.total_spent == 0 - 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 + assert updated.reserved_balance == 3_000 + assert updated.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 diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 31dd5767..73b3e43f 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -16,14 +16,13 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.pool import StaticPool -from sqlmodel import SQLModel, select +from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession from routstr.auth import pay_for_request from routstr.balance import refund_wallet_endpoint from routstr.core.db import ( ApiKey, - ReservationRelease, release_stale_reservations, reset_all_reserved_balances, ) @@ -70,26 +69,6 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None: assert key.reserved_at >= before -@pytest.mark.asyncio -async def test_pay_for_request_sets_reserved_at_on_child_key( - session: AsyncSession, -) -> None: - parent = ApiKey(hashed_key="parentkey", balance=10_000) - child = ApiKey(hashed_key="childkey", balance=0, parent_key_hash="parentkey") - session.add(parent) - session.add(child) - await session.commit() - - await pay_for_request(child, 1_000, session) - - await session.refresh(parent) - await session.refresh(child) - assert parent.reserved_balance == 1_000 - assert parent.reserved_at is not None - assert child.reserved_balance == 1_000 - assert child.reserved_at is not None - - @pytest.mark.asyncio async def test_revert_clears_reserved_at_when_fully_released( session: AsyncSession, @@ -153,39 +132,6 @@ async def test_release_stale_reservations_releases_old(session: AsyncSession) -> assert key.reserved_at is None -@pytest.mark.asyncio -async def test_targeted_parent_cleanup_releases_child_owned_reservation( - session: AsyncSession, -) -> None: - parent = ApiKey(hashed_key="stale-parent", balance=5_000) - child = ApiKey( - hashed_key="stale-child", parent_key_hash=parent.hashed_key, balance=0 - ) - session.add_all([parent, child]) - await session.commit() - await pay_for_request(child, 1_000, session) - reservation = ( - await session.exec( - select(ReservationRelease).where( - ReservationRelease.key_hash == child.hashed_key - ) - ) - ).one() - reservation.created_at = int(time.time()) - 1_000 - session.add(reservation) - await session.commit() - - released = await release_stale_reservations( - session, max_age_seconds=300, key_hash=parent.hashed_key - ) - - assert released == 1 - await session.refresh(parent) - await session.refresh(child) - assert parent.reserved_balance == 0 - assert child.reserved_balance == 0 - - @pytest.mark.asyncio async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> None: key = ApiKey( diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index ddc8d9f5..e70804c1 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -82,47 +82,42 @@ async def test_release_only_owns_its_concurrent_reservation() -> None: @pytest.mark.asyncio -async def test_release_updates_parent_and_child_atomically() -> None: +async def test_release_clears_reservation_aggregates_atomically() -> None: engine = await _engine() - parent = ApiKey(hashed_key="parent", balance=1_000) - child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + key = ApiKey(hashed_key="key", balance=1_000) async with AsyncSession(engine, expire_on_commit=False) as session: - session.add_all([parent, child]) + session.add(key) await session.commit() - await pay_for_request(child, 500, session) - snapshot = await get_reservation_snapshot(child, session) + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) assert await release_reservation(snapshot, session, 500) is True - await session.refresh(parent) - await session.refresh(child) - assert (parent.reserved_balance, child.reserved_balance) == (0, 0) - assert (parent.reserved_at, child.reserved_at) == (None, None) + await session.refresh(key) + assert (key.reserved_balance, key.reserved_at) == (0, None) await engine.dispose() @pytest.mark.asyncio -async def test_release_repairs_partial_parent_child_corruption() -> None: - """A child aggregate that no longer holds the reservation must not leave +async def test_release_repairs_partial_aggregate_corruption() -> None: + """An 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() - parent = ApiKey(hashed_key="parent", balance=1_000) - child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + key = ApiKey(hashed_key="key", balance=1_000) async with AsyncSession(engine, expire_on_commit=False) as session: - session.add_all([parent, child]) + session.add(key) await session.commit() - await pay_for_request(child, 500, session) - snapshot = await get_reservation_snapshot(child, session) - child.reserved_balance = 100 - session.add(child) + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + key.reserved_balance = 100 + session.add(key) await session.commit() assert await release_reservation(snapshot, session, 500) is True - await session.refresh(parent) - await session.refresh(child) + await session.refresh(key) 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 key.reserved_balance == 100 assert record is not None and record.status == "released" await engine.dispose() diff --git a/ui/components/child-key-creator.tsx b/ui/components/child-key-creator.tsx deleted file mode 100644 index b10847f9..00000000 --- a/ui/components/child-key-creator.tsx +++ /dev/null @@ -1,481 +0,0 @@ -'use client'; - -import { useState } from 'react'; -import { useCopyToClipboard } from '@/hooks/use-copy-to-clipboard'; -import { useWalletInfo } from '@/hooks/use-wallet-info'; -import { WalletService } from '@/lib/api/services/wallet'; -import { ApiKeyInput } from './api-key-input'; -import { Button } from '@/components/ui/button'; -import { - Card, - CardContent, - CardDescription, - CardHeader, - CardTitle, -} from '@/components/ui/card'; -import { Alert, AlertDescription, AlertTitle } from '@/components/ui/alert'; -import { Input } from '@/components/ui/input'; -import { Textarea } from '@/components/ui/textarea'; -import { Label } from '@/components/ui/label'; -import { Key, Copy, Check, Loader2, Plus, Trash2 } from 'lucide-react'; -import { toast } from 'sonner'; -import { KeyOptions } from './key-options'; - -interface KeyConfig { - id: string; - count: number; - balanceLimit: string; - balanceLimitReset: string; - validityDate: string; -} - -interface ChildKeyCreatorProps { - baseUrl?: string; - apiKey?: string; - onApiKeyChange?: (apiKey: string) => void; - costPerKeyMsats?: number; -} - -function formatSats(msats: number): string { - return new Intl.NumberFormat('en-US').format(Math.floor(msats / 1000)); -} - -function formatMsats(msats: number): string { - return new Intl.NumberFormat('en-US').format(msats); -} - -export function ChildKeyCreator({ - baseUrl, - apiKey: propApiKey, - onApiKeyChange, - costPerKeyMsats, -}: ChildKeyCreatorProps) { - const [internalApiKey, setInternalApiKey] = useState(''); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - const [configs, setConfigs] = useState([ - { - id: crypto.randomUUID(), - count: 1, - balanceLimit: '', - balanceLimitReset: '', - validityDate: '', - }, - ]); - - const activeApiKey = propApiKey ?? internalApiKey; - const { data: walletInfo } = useWalletInfo(baseUrl ?? '', activeApiKey); - - const handleApiKeyChange = (val: string) => { - setInternalApiKey(val); - onApiKeyChange?.(val); - }; - - const [newKeys, setNewKeys] = useState([]); - const [resultInfo, setResultInfo] = useState<{ - cost_msats: number; - parent_balance: number; - } | null>(null); - const { copiedKey, copy } = useCopyToClipboard(); - - const addConfig = () => { - setConfigs([ - ...configs, - { - id: crypto.randomUUID(), - count: 1, - balanceLimit: '', - balanceLimitReset: '', - validityDate: '', - }, - ]); - }; - - const removeConfig = (id: string) => { - if (configs.length > 1) { - setConfigs(configs.filter((c) => c.id !== id)); - } - }; - - const updateConfig = (id: string, updates: Partial) => { - setConfigs(configs.map((c) => (c.id === id ? { ...c, ...updates } : c))); - }; - - const handleCreateKey = async () => { - if (!activeApiKey && baseUrl) { - toast.error('Please provide a Parent API key first'); - return; - } - - setLoading(true); - setError(null); - try { - let allNewKeys: string[] = []; - let totalCost = 0; - let lastParentBalance = 0; - - for (const config of configs) { - const requestedCount = Math.max(1, Math.min(50, Number(config.count))); - const result = await WalletService.createChildKey( - baseUrl, - activeApiKey, - requestedCount, - config.balanceLimit ? parseInt(config.balanceLimit) : undefined, - config.balanceLimitReset || undefined, - config.validityDate - ? Math.floor( - new Date(config.validityDate + 'T23:59:59').getTime() / 1000 - ) - : undefined - ); - - if (result.api_keys) { - allNewKeys = [...allNewKeys, ...result.api_keys]; - } - totalCost += result.cost_msats; - lastParentBalance = result.parent_balance; - } - - setNewKeys(allNewKeys); - setResultInfo({ - cost_msats: totalCost, - parent_balance: lastParentBalance, - }); - - toast.success( - `${allNewKeys.length} child API key${ - allNewKeys.length > 1 ? 's' : '' - } created successfully` - ); - } catch (error) { - console.error('Failed to create child key:', error); - let errorMessage = - error instanceof Error ? error.message : 'Failed to create child key'; - try { - const parsed = JSON.parse(errorMessage); - errorMessage = - parsed.detail?.error?.message || - (typeof parsed.detail === 'string' ? parsed.detail : errorMessage); - } catch {} - setError(errorMessage); - toast.error(errorMessage); - } finally { - setLoading(false); - } - }; - - const copyToClipboard = async (key: string) => { - if (await copy(key, key)) { - toast.success('API key copied to clipboard'); - } - }; - - const copyAllToClipboard = async () => { - if (await copy(newKeys.join('\n'), 'all')) { - toast.success('All API keys copied to clipboard'); - } - }; - - return ( -
- - -
-
- Create Child API Key - - Generate secondary API keys that share your account balance. - -
- {costPerKeyMsats !== undefined && ( -
-

- Unit Cost -

-

- {costPerKeyMsats / 1000} sats -

-
- )} -
-
- -
- {baseUrl && ( -
- -
-
- -
- -
- {walletInfo && ( -
-
- - Spendable Balance - - - {formatSats(walletInfo.balanceMsats)} sats - -
-
- - Total Requests - - - {walletInfo.totalRequests} - -
-
- - Total Spent - -
-

- {formatSats(walletInfo.totalSpent)} sats -

-

- {formatMsats(walletInfo.totalSpent)} msats -

-
-
-
- )} -
- )} - -
- {configs.map((config) => ( -
- {configs.length > 1 && ( - - )} -
-
- - { - const val = parseInt(e.target.value); - updateConfig(config.id, { - count: isNaN(val) - ? 1 - : Math.max(1, Math.min(50, val)), - }); - }} - className='h-9' - /> -
- -
- - updateConfig(config.id, { balanceLimit: val }) - } - validityDate={config.validityDate} - setValidityDate={(val) => - updateConfig(config.id, { validityDate: val }) - } - balanceLimitReset={config.balanceLimitReset} - setBalanceLimitReset={(val) => - updateConfig(config.id, { balanceLimitReset: val }) - } - /> -
-
-
- ))} - -
- -
- -
-
- {costPerKeyMsats && ( -

- Total Cost:{' '} - - {costPerKeyMsats * - configs.reduce( - (acc, c) => acc + Number(c.count), - 0 - )}{' '} - mSats - -

- )} -
- - -
- - {error && ( - - Error - {error} - - )} - -

- Each key creation has a small one-time fee. -

-
- - {newKeys.length > 0 && ( -
- - - {newKeys.length} New API Key{newKeys.length > 1 ? 's' : ''}{' '} - Generated - - - Copy {newKeys.length > 1 ? 'these keys' : 'this key'} now. - You won't be able to see them again. - {resultInfo && ( -
- Total Cost: {resultInfo.cost_msats / 1000} sats | New - Balance: {resultInfo.parent_balance / 1000} sats -
- )} -
-
- -
-
- - Generated Keys ({newKeys.length}) - - {newKeys.length > 1 && ( - - )} -
-
- {newKeys.map((key, index) => ( -
- - {key} - - -
- ))} -
-
- - {newKeys.length > 3 && ( -
- -
-