From fc042c768c9eec3a2e64c1f30af52594eedb4380 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 19:01:52 +0100 Subject: [PATCH 1/6] support child keys mapped to parent balance --- migrations/versions/a86e5348850b_.py | 42 ++++++ routstr/auth.py | 206 +++++++++++++++++++++------ routstr/balance.py | 94 ++++++++++-- routstr/core/admin.py | 3 +- routstr/core/db.py | 3 + routstr/core/settings.py | 1 + 6 files changed, 296 insertions(+), 53 deletions(-) create mode 100644 migrations/versions/a86e5348850b_.py diff --git a/migrations/versions/a86e5348850b_.py b/migrations/versions/a86e5348850b_.py new file mode 100644 index 00000000..12c35e41 --- /dev/null +++ b/migrations/versions/a86e5348850b_.py @@ -0,0 +1,42 @@ +""" + +Revision ID: a86e5348850b +Revises: b9667ffc5701 +Create Date: 2026-01-10 18:57:48.475781 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a86e5348850b" +down_revision = "b9667ffc5701" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Use batch_alter_table for SQLite compatibility + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.add_column( + sa.Column( + "parent_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True + ) + ) + batch_op.create_index( + batch_op.f("ix_api_keys_parent_key_hash"), ["parent_key_hash"], unique=False + ) + batch_op.create_foreign_key( + "fk_api_keys_parent_key_hash", + "api_keys", + ["parent_key_hash"], + ["hashed_key"], + ) + + +def downgrade() -> None: + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.drop_constraint("fk_api_keys_parent_key_hash", type_="foreignkey") + batch_op.drop_index(batch_op.f("ix_api_keys_parent_key_hash")) + batch_op.drop_column("parent_key_hash") diff --git a/routstr/auth.py b/routstr/auth.py index b3be04b2..1f869886 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -286,30 +286,55 @@ async def validate_bearer_key( ) +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: """Process payment for a request.""" + billing_key = await get_billing_key(key, session) + logger.info( "Processing payment for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "current_balance": key.balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "current_balance": billing_key.balance, "required_cost": cost_per_request, - "sufficient_balance": key.balance >= cost_per_request, + "sufficient_balance": billing_key.balance >= cost_per_request, }, ) - if key.total_balance < cost_per_request: + if billing_key.total_balance < cost_per_request: logger.warning( "Insufficient balance for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "balance": key.balance, - "reserved_balance": key.reserved_balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, "required": cost_per_request, - "shortfall": cost_per_request - key.total_balance, + "shortfall": cost_per_request - billing_key.total_balance, }, ) @@ -317,7 +342,7 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.total_balance} available. (reserved: {billing_key.reserved_balance})", "type": "insufficient_quota", "code": "insufficient_balance", } @@ -328,15 +353,16 @@ async def pay_for_request( "Charging base cost for request", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost": cost_per_request, - "balance_before": key.balance, + "balance_before": billing_key.balance, }, ) # Charge the base cost for the request atomically to avoid race conditions stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request) .values( reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, @@ -344,6 +370,16 @@ async def pay_for_request( ) ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also increment total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) + 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -351,8 +387,9 @@ async def pay_for_request( "Concurrent request depleted balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "required_cost": cost_per_request, - "current_balance": key.balance, + "current_balance": billing_key.balance, }, ) @@ -361,23 +398,26 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Payment processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost_per_request, - "new_balance": key.balance, - "total_spent": key.total_spent, - "total_requests": key.total_requests, + "new_balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "total_requests": billing_key.total_requests, }, ) @@ -387,9 +427,11 @@ async def pay_for_request( async def revert_pay_for_request( key: ApiKey, session: AsyncSession, cost_per_request: int ) -> None: + billing_key = await get_billing_key(key, session) + stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, @@ -397,27 +439,40 @@ async def revert_pay_for_request( ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also decrement total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) - 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to revert payment - insufficient reserved balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost_to_revert": cost_per_request, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, }, ) raise HTTPException( status_code=402, detail={ "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.", + "message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "payment_error", "code": "payment_error", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) async def adjust_payment_for_tokens( @@ -428,15 +483,17 @@ async def adjust_payment_for_tokens( This is called after the initial payment and the upstream request is complete. Returns cost data to be included in the response. """ + billing_key = await get_billing_key(key, session) model = response_data.get("model", "unknown") logger.debug( "Starting payment adjustment for tokens", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "deducted_max_cost": deducted_max_cost, - "current_balance": key.balance, + "current_balance": billing_key.balance, "has_usage": "usage" in response_data, }, ) @@ -446,8 +503,10 @@ async def adjust_payment_for_tokens( try: release_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost + ) ) await session.exec(release_stmt) # type: ignore[call-overload] await session.commit() @@ -455,13 +514,18 @@ async def adjust_payment_for_tokens( "Released reservation without charging (fallback)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, }, ) except Exception as e: logger.error( "Failed to release reservation in fallback", - extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."}, + extra={ + "error": str(e), + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + }, ) match await calculate_cost(response_data, deducted_max_cost, session): @@ -470,6 +534,7 @@ async def adjust_payment_for_tokens( "Using max cost data (no token adjustment)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "max_cost": cost.total_msats, }, @@ -477,7 +542,7 @@ async def adjust_payment_for_tokens( # Finalize by releasing reservation and charging max cost finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, balance=col(ApiKey.balance) - cost.total_msats, @@ -485,27 +550,41 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + cost.total_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to finalize max-cost payment - retrying reservation release", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": cost.total_msats, "model": model, }, ) await release_reservation_only() else: - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Max cost payment finalized", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost.total_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -521,6 +600,7 @@ async def adjust_payment_for_tokens( "Calculated token-based cost", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "token_cost": cost.total_msats, "deducted_max_cost": deducted_max_cost, @@ -533,11 +613,15 @@ async def adjust_payment_for_tokens( if cost_difference == 0: logger.debug( "Finalizing with exact reserved cost", - extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -546,8 +630,20 @@ async def adjust_payment_for_tokens( ) ) await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) return cost.dict() # this should never happen why do we handle this??? @@ -557,16 +653,17 @@ async def adjust_payment_for_tokens( "Additional charge required for token usage", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "additional_charge": cost_difference, - "current_balance": key.balance, - "sufficient_balance": key.balance >= cost_difference, + "current_balance": billing_key.balance, + "sufficient_balance": billing_key.balance >= cost_difference, "model": model, }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -575,18 +672,31 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Finalized payment with additional charge", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": total_cost_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -595,6 +705,7 @@ async def adjust_payment_for_tokens( "Failed to finalize additional charge - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "attempted_charge": total_cost_msats, "model": model, }, @@ -607,15 +718,16 @@ async def adjust_payment_for_tokens( "Refunding excess payment", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refund_amount": refund, - "current_balance": key.balance, + "current_balance": billing_key.balance, "model": model, }, ) refund_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -624,6 +736,16 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(refund_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -631,8 +753,9 @@ async def adjust_payment_for_tokens( "Failed to finalize payment - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": total_cost_msats, "model": model, }, @@ -640,14 +763,17 @@ async def adjust_payment_for_tokens( await release_reservation_only() else: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Refund processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refunded_amount": refund, - "new_balance": key.balance, + "new_balance": billing_key.balance, "final_cost": cost.total_msats, "model": model, }, diff --git a/routstr/balance.py b/routstr/balance.py index 6697a5db..472f8b05 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -32,16 +32,30 @@ async def get_key_from_header( ) -# TODO: remove this endpoint when frontend is updated -@router.get("/", include_in_schema=False) -async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: +async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) return { "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, + "balance": billing_key.balance, + "reserved": billing_key.reserved_balance, + "is_child": key.parent_key_hash is not None, + "parent_key": "sk-" + key.parent_key_hash if key.parent_key_hash else None, + "total_requests": key.total_requests, + "total_spent": key.total_spent, } +# TODO: remove this endpoint when frontend is updated +@router.get("/", include_in_schema=False) +async def account_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) + + # TODO: Implement POST /v1/wallet/create endpoint # This endpoint should accept: # - cashu_token (required): The eCash token to deposit @@ -66,12 +80,11 @@ async def create_balance( @router.get("/info") -async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: - return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, - } +async def wallet_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) class TopupRequest(BaseModel): @@ -85,6 +98,10 @@ async def topup_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) + if topup_request is not None: cashu_token = topup_request.cashu_token if cashu_token is None: @@ -94,7 +111,7 @@ async def topup_wallet_endpoint( if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") try: - amount_msats = await credit_balance(cashu_token, key, session) + amount_msats = await credit_balance(cashu_token, billing_key, session) except ValueError as e: error_msg = str(e) if "already spent" in error_msg.lower(): @@ -155,6 +172,12 @@ async def refund_wallet_endpoint( key: ApiKey = await validate_bearer_key(bearer_value, session) + if key.parent_key_hash: + raise HTTPException( + status_code=400, + detail="Cannot refund child key. Please refund the parent key instead.", + ) + remaining_balance_msats: int = key.total_balance if key.refund_currency == "sat": @@ -240,6 +263,53 @@ async def wallet_catch_all(path: str) -> NoReturn: ) +@router.post("/child-key") +async def create_child_key( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + """Creates a child API key that uses the parent's balance.""" + # 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 = settings.child_key_cost + + if key.total_balance < cost: + raise HTTPException( + status_code=402, + detail=f"Insufficient balance to create child key. {cost} mSats required.", + ) + + # Deduct cost from parent + key.balance -= cost + key.total_spent += cost + session.add(key) + + # Generate new key + import secrets + + 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, + ) + session.add(child_key) + await session.commit() + + return { + "api_key": "sk-" + new_key_hash, + "cost_msats": cost, + "parent_balance": key.balance, + } + + balance_router.include_router(lightning_router) balance_router.include_router(router) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index d4f7ec59..626df30c 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -124,7 +124,7 @@ async def partial_apikeys(request: Request) -> str: rows = "".join( [ - f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" + f"{key.hashed_key}{'
(Child of ' + key.parent_key_hash[:8] + '...)' if key.parent_key_hash else ''}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" for key in api_keys ] ) @@ -158,6 +158,7 @@ async def get_temporary_balances_api(request: Request) -> list[dict[str, object] "total_requests": key.total_requests, "refund_address": key.refund_address, "key_expiry_time": key.key_expiry_time, + "parent_key_hash": key.parent_key_hash, } for key in api_keys ] diff --git a/routstr/core/db.py b/routstr/core/db.py index 4c236d3a..46c56582 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -47,6 +47,9 @@ 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 + ) @property def total_balance(self) -> int: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 685f8970..3b659615 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -59,6 +59,7 @@ 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=1000, 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( From daf17f51ab79d60fda29dfb33f18a67a8571e357 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 20:00:52 +0100 Subject: [PATCH 2/6] fix: move child-key route before catch-all and fix indentation --- examples/create_child_keys.py | 44 +++++++++ routstr/balance.py | 24 ++--- routstr/core/main.py | 8 +- tests/integration/test_child_keys.py | 131 +++++++++++++++++++++++++++ 4 files changed, 190 insertions(+), 17 deletions(-) create mode 100644 examples/create_child_keys.py create mode 100644 tests/integration/test_child_keys.py diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py new file mode 100644 index 00000000..eefa3710 --- /dev/null +++ b/examples/create_child_keys.py @@ -0,0 +1,44 @@ +import httpx +import sys +import json + + +def create_child_keys(base_url, api_key, count=3): + 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/routstr/balance.py b/routstr/balance.py index 472f8b05..3ee8ca11 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -251,18 +251,6 @@ async def donate(token: str, ref: str | None = None) -> str: return "Invalid token." -@router.api_route( - "/{path:path}", - methods=["GET", "POST", "PUT", "DELETE"], - include_in_schema=False, - response_model=None, -) -async def wallet_catch_all(path: str) -> NoReturn: - raise HTTPException( - status_code=404, detail="Not found check /docs for available endpoints" - ) - - @router.post("/child-key") async def create_child_key( key: ApiKey = Depends(get_key_from_header), @@ -310,6 +298,18 @@ async def create_child_key( } +@router.api_route( + "/{path:path}", + methods=["GET", "POST", "PUT", "DELETE"], + include_in_schema=False, + response_model=None, +) +async def wallet_catch_all(path: str) -> NoReturn: + raise HTTPException( + status_code=404, detail="Not found check /docs for available endpoints" + ) + + balance_router.include_router(lightning_router) balance_router.include_router(router) diff --git a/routstr/core/main.py b/routstr/core/main.py index 5bce5d32..3083097d 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -13,12 +13,10 @@ from starlette.exceptions import HTTPException from ..balance import balance_router, deprecated_wallet_router from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider -from ..payment.models import ( - models_router, - update_sats_pricing, -) +from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically -from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically +from ..proxy import (initialize_upstreams, proxy_router, + refresh_model_maps_periodically) from ..wallet import periodic_payout from .admin import admin_router from .db import create_session, init_db, run_migrations diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py new file mode 100644 index 00000000..77da7675 --- /dev/null +++ b/tests/integration/test_child_keys.py @@ -0,0 +1,131 @@ +import pytest +import secrets +from fastapi import HTTPException +from routstr.core.db import ApiKey +from routstr.balance import create_child_key +from routstr.auth import pay_for_request, adjust_payment_for_tokens +from routstr.core.settings import settings + + +@pytest.mark.asyncio +async def test_child_key_flow(integration_session): + # 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(parent_key, integration_session) + + assert "api_key" in result + assert result["cost_msats"] == 1000 + assert result["parent_balance"] == 9000 + + child_key_raw = result["api_key"][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, **kwargs): + 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 + ) + 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): + 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(parent_key, integration_session) + assert exc.value.status_code == 402 + + +@pytest.mark.asyncio +async def test_child_key_cannot_create_child(integration_session): + 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(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) From 917a4d32b1630440a6f4ed00fb547ba55a908c07 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 20:10:59 +0100 Subject: [PATCH 3/6] ui: visualize parent-child relationship in balances page --- ui/components/temporary-balances.tsx | 79 +++++++++++++++++++++++----- ui/lib/api/services/admin.ts | 1 + 2 files changed, 67 insertions(+), 13 deletions(-) diff --git a/ui/components/temporary-balances.tsx b/ui/components/temporary-balances.tsx index c33d0d43..004ba77f 100644 --- a/ui/components/temporary-balances.tsx +++ b/ui/components/temporary-balances.tsx @@ -60,7 +60,11 @@ export function TemporaryBalances({ let totalRequests = 0; balances.forEach((balance) => { - totalBalance += balance.balance || 0; + // Only count parents for total balance to avoid double counting + // since child keys use parent balance + if (!balance.parent_key_hash) { + totalBalance += balance.balance || 0; + } totalSpent += balance.total_spent || 0; totalRequests += balance.total_requests || 0; }); @@ -72,6 +76,34 @@ export function TemporaryBalances({ ? calculateTotals(data) : { totalBalance: 0, totalSpent: 0, totalRequests: 0 }; + // Group parents and children + const hierarchicalData = (() => { + if (!data) return []; + + const parents = filteredData.filter((item) => !item.parent_key_hash); + const result: (TemporaryBalance & { isChild?: boolean })[] = []; + + parents.forEach((parent) => { + result.push(parent); + const children = data.filter( + (item) => item.parent_key_hash === parent.hashed_key + ); + children.forEach((child) => { + result.push({ ...child, isChild: true }); + }); + }); + + // Add children whose parents didn't match the search or aren't in the list + const orphans = filteredData.filter( + (item) => + item.parent_key_hash && + !result.some((r) => r.hashed_key === item.hashed_key) + ); + result.push(...orphans.map((o) => ({ ...o, isChild: true }))); + + return result; + })(); + return ( <> @@ -182,22 +214,32 @@ export function TemporaryBalances({
Expiry Time
- {filteredData.length > 0 ? ( - filteredData.map((balance, index) => ( + {hierarchicalData.length > 0 ? ( + hierarchicalData.map((balance, index) => (
{/* Desktop Layout */}
-
+
+ {balance.isChild && ( + + Child + + )} {balance.hashed_key}
- {formatBalance(balance.balance)} + {balance.isChild ? ( + (Parent) + ) : ( + formatBalance(balance.balance) + )}
{formatBalance(balance.total_spent)} @@ -226,13 +268,20 @@ export function TemporaryBalances({ {/* Mobile Layout */}
-
- - Key - -
- {balance.hashed_key} +
+
+ + {balance.isChild ? 'Child Key' : 'Key'} + +
+ {balance.hashed_key} +
+ {balance.isChild && ( + + Child + + )}
@@ -241,7 +290,11 @@ export function TemporaryBalances({ Balance
- {formatBalance(balance.balance)} + {balance.isChild ? ( + (Uses Parent) + ) : ( + formatBalance(balance.balance) + )}
diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index f94729ba..2ea27e7e 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -926,6 +926,7 @@ export const TemporaryBalanceSchema = z.object({ total_requests: z.number(), refund_address: z.string().nullable(), key_expiry_time: z.number().nullable(), + parent_key_hash: z.string().nullable().optional(), }); export type TemporaryBalance = z.infer; From 88bcc0edcb9cebc3532861ddfe816f33718c0d1c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:34:48 +0100 Subject: [PATCH 4/6] fmt --- examples/create_child_keys.py | 6 ++++-- routstr/core/main.py | 3 +-- tests/integration/test_child_keys.py | 8 +++++--- 3 files changed, 10 insertions(+), 7 deletions(-) diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py index eefa3710..5e5d23e7 100644 --- a/examples/create_child_keys.py +++ b/examples/create_child_keys.py @@ -1,6 +1,7 @@ -import httpx -import sys import json +import sys + +import httpx def create_child_keys(base_url, api_key, count=3): @@ -42,3 +43,4 @@ if __name__ == "__main__": print(json.dumps(keys, indent=2)) else: print("\nNo child keys were created.") + diff --git a/routstr/core/main.py b/routstr/core/main.py index 3083097d..00a51ec3 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -15,8 +15,7 @@ from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically -from ..proxy import (initialize_upstreams, proxy_router, - refresh_model_maps_periodically) +from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..wallet import periodic_payout from .admin import admin_router from .db import create_session, init_db, run_migrations diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py index 77da7675..4265430a 100644 --- a/tests/integration/test_child_keys.py +++ b/tests/integration/test_child_keys.py @@ -1,9 +1,11 @@ -import pytest import secrets + +import pytest from fastapi import HTTPException -from routstr.core.db import ApiKey + +from routstr.auth import adjust_payment_for_tokens, pay_for_request from routstr.balance import create_child_key -from routstr.auth import pay_for_request, adjust_payment_for_tokens +from routstr.core.db import ApiKey from routstr.core.settings import settings From d4339287beb8173bfdf80f5202083fad297eb585 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:35:16 +0100 Subject: [PATCH 5/6] chore: add type annotations to example and test files --- examples/create_child_keys.py | 3 +-- tests/integration/test_child_keys.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py index 5e5d23e7..24f4556d 100644 --- a/examples/create_child_keys.py +++ b/examples/create_child_keys.py @@ -4,7 +4,7 @@ import sys import httpx -def create_child_keys(base_url, api_key, count=3): +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}...") @@ -43,4 +43,3 @@ if __name__ == "__main__": print(json.dumps(keys, indent=2)) else: print("\nNo child keys were created.") - diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py index 4265430a..e868d1cb 100644 --- a/tests/integration/test_child_keys.py +++ b/tests/integration/test_child_keys.py @@ -1,7 +1,9 @@ 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 create_child_key @@ -10,7 +12,7 @@ from routstr.core.settings import settings @pytest.mark.asyncio -async def test_child_key_flow(integration_session): +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( @@ -61,7 +63,7 @@ async def test_child_key_flow(integration_session): import routstr.auth from routstr.payment.cost_calculation import CostData - async def mock_calculate_cost(*args, **kwargs): + async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData: return CostData( base_msats=0, input_msats=200, output_msats=200, total_msats=400 ) @@ -95,7 +97,9 @@ async def test_child_key_flow(integration_session): @pytest.mark.asyncio -async def test_child_key_insufficient_balance(integration_session): +async def test_child_key_insufficient_balance( + integration_session: AsyncSession, +) -> None: parent_key = ApiKey( hashed_key="poor_parent_" + secrets.token_hex(4), balance=500, @@ -112,7 +116,7 @@ async def test_child_key_insufficient_balance(integration_session): @pytest.mark.asyncio -async def test_child_key_cannot_create_child(integration_session): +async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None: parent_key = ApiKey( hashed_key="parent_" + secrets.token_hex(4), balance=10000, From db021866d87e4f7e464195b85d4d62d88696146a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:42:21 +0100 Subject: [PATCH 6/6] fmt u --- ui/components/temporary-balances.tsx | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/ui/components/temporary-balances.tsx b/ui/components/temporary-balances.tsx index 004ba77f..7e7927d9 100644 --- a/ui/components/temporary-balances.tsx +++ b/ui/components/temporary-balances.tsx @@ -220,15 +220,18 @@ export function TemporaryBalances({ key={index} className={cn( 'hover:bg-muted/50 border-t p-3 text-sm transition-colors', - balance.balance === 0 && !balance.isChild && 'opacity-60', - balance.isChild && 'bg-blue-50/30 ml-4 border-l-2 border-l-blue-200' + balance.balance === 0 && + !balance.isChild && + 'opacity-60', + balance.isChild && + 'ml-4 border-l-2 border-l-blue-200 bg-blue-50/30' )} > {/* Desktop Layout */}
{balance.isChild && ( - + Child )} @@ -236,7 +239,9 @@ export function TemporaryBalances({
{balance.isChild ? ( - (Parent) + + (Parent) + ) : ( formatBalance(balance.balance) )} @@ -278,7 +283,7 @@ export function TemporaryBalances({
{balance.isChild && ( - + Child )} @@ -291,7 +296,9 @@ export function TemporaryBalances({
{balance.isChild ? ( - (Uses Parent) + + (Uses Parent) + ) : ( formatBalance(balance.balance) )}