diff --git a/README.md b/README.md index 61db87d8..da66b107 100644 --- a/README.md +++ b/README.md @@ -51,14 +51,26 @@ curl https://api.routstr.com/v1/chat/completions \ ## Quick Start (Docker) -If you are a node runner, start a Routstr Core instance and configure upstream access in the dashboard. +If you are a node runner, start a Routstr Core instance using Docker Compose: -```bash -docker run -d \ - --name routstr-proxy \ - -p 8000:8000 \ - ghcr.io/routstr/proxy:latest -``` +1. **Prepare your `.env`**: + ```bash + ADMIN_PASSWORD=mysecretpassword + NAME="My AI Node" + DESCRIPTION="Fast access to models" + NSEC=yournsec + RECEIVE_LN_ADDRESS=yourname@wallet.com + ``` + +2. **Start the services**: + ```bash + docker compose up -d + ``` + +3. **Configure**: + Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing. + +For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**. ## Development diff --git a/docs/provider/deployment.md b/docs/provider/deployment.md index 99f0166a..e1df11e8 100644 --- a/docs/provider/deployment.md +++ b/docs/provider/deployment.md @@ -6,16 +6,7 @@ Production deployment guide for Routstr Provider nodes. For production, use Docker Compose with persistent storage and optional Tor support. -### Unified Setup (All-in-one) -To build and run the node with the UI integrated in a single container using the multi-stage build: - -```bash -docker build -f Dockerfile.full -t routstr-full . -docker run -d -p 8000:8000 --env-file .env routstr-full -``` - -### Advanced Setup (Separated UI & Node) -Use the included `compose.yml` for a more flexible setup that separates the UI build process from the node execution. This is useful for development or when you want to manage Tor as a separate service. +Use the included `compose.yml` for a flexible setup that handles both the UI and the node execution. This is useful for development or when you want to manage Tor as a separate service. ```bash docker compose up -d @@ -184,20 +175,16 @@ docker compose up -d ## Building from Source -### Unified Image (UI + Node) -The easiest way to build everything from source into a single production-ready image: +### Using Docker Compose +The easiest way to build everything from source: ```bash -docker build -f Dockerfile.full -t routstr-full . +docker compose build ``` ### Individual Components -If you prefer building them separately or using Docker Compose: +If you prefer building the node only (requires manual UI build first): ```bash -# Build using compose -docker compose build - -# Or build the node only (requires manual UI build first) docker build -t routstr-node . ``` diff --git a/docs/provider/quickstart.md b/docs/provider/quickstart.md index 28e6c593..5b33c359 100644 --- a/docs/provider/quickstart.md +++ b/docs/provider/quickstart.md @@ -35,6 +35,7 @@ ADMIN_PASSWORD=mysecretpassword # Node Identity NAME="My AI Node" DESCRIPTION="Fast access to models" +NSEC=yournsec # Lightning Payouts RECEIVE_LN_ADDRESS=yourname@wallet.com @@ -43,32 +44,10 @@ RECEIVE_LN_ADDRESS=yourname@wallet.com ## 2. Start the Node -You can run the pre-built image directly: +The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor. ```bash -docker run -d \ - --name routstr \ - -p 8000:8000 \ - --env-file .env \ - -v routstr-data:/app/data \ - ghcr.io/routstr/proxy:latest -``` - -*Note: The pre-built image does not contain the UI. For the all-in-one experience with the Admin Dashboard, use the Build from Source instructions below.* - -### Build from Source (Recommended) - -If you want to build the node and UI yourself from source, use the unified Dockerfile: - -```bash -git clone https://github.com/routstr/routstr-core.git -cd routstr-core -# Edit your .env with ADMIN_PASSWORD and API keys -cp .env.example .env -nano .env - -docker build -f Dockerfile.full -t routstr-local . -docker run -d -p 8000:8000 --env-file .env --name routstr routstr-local +docker compose up -d ``` Verify it's running: @@ -77,6 +56,15 @@ Verify it's running: curl http://localhost:8000/v1/info ``` +### Build from Source (Optional) + +If you've cloned the repository and want to build the images yourself: + +```bash +docker compose build +docker compose up -d +``` + --- ## 3. Configure via Dashboard diff --git a/migrations/versions/b1c2d3e4f5a6_add_upstream_model_id_to_models.py b/migrations/versions/b1c2d3e4f5a6_add_upstream_model_id_to_models.py new file mode 100644 index 00000000..a299b32f --- /dev/null +++ b/migrations/versions/b1c2d3e4f5a6_add_upstream_model_id_to_models.py @@ -0,0 +1,33 @@ +"""add forwarded_model_id to models + +Revision ID: b1c2d3e4f5a6 +Revises: a776ca70e5fe +Create Date: 2026-04-05 00:00:00.000000 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "b1c2d3e4f5a6" +down_revision = "a776ca70e5fe" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "models", + sa.Column( + "forwarded_model_id", + sqlmodel.sql.sqltypes.AutoString(), + nullable=True, + ), + ) + # Backfill: set forwarded_model_id = id for all existing rows + op.execute("UPDATE models SET forwarded_model_id = id WHERE forwarded_model_id IS NULL") + + +def downgrade() -> None: + op.drop_column("models", "forwarded_model_id") diff --git a/migrations/versions/c3d4e5f6a7b8_add_source_to_cashu_transactions.py b/migrations/versions/c3d4e5f6a7b8_add_source_to_cashu_transactions.py new file mode 100644 index 00000000..31446019 --- /dev/null +++ b/migrations/versions/c3d4e5f6a7b8_add_source_to_cashu_transactions.py @@ -0,0 +1,36 @@ +"""add source to cashu_transactions + +Revision ID: c3d4e5f6a7b8 +Revises: b1c2d3e4f5a6 +Create Date: 2026-04-10 00:00:00.000000 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "c3d4e5f6a7b8" +down_revision = "b1c2d3e4f5a6" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + columns = [col["name"] for col in inspector.get_columns("cashu_transactions")] + if "source" not in columns: + op.add_column( + "cashu_transactions", + sa.Column( + "source", + sqlmodel.sql.sqltypes.AutoString(), + nullable=False, + server_default="x-cashu", + ), + ) + + +def downgrade() -> None: + op.drop_column("cashu_transactions", "source") diff --git a/routstr/algorithm.py b/routstr/algorithm.py index e3efa261..73ff75d1 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -217,6 +217,10 @@ def create_model_mappings( if prefixed_id not in aliases: aliases.append(prefixed_id) + # Register forwarded_model_id as a routable alias + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + aliases.append(model_to_use.forwarded_model_id) + # Try to set each alias for alias in aliases: _add_candidate(alias, model_to_use, upstream) @@ -305,6 +309,10 @@ def create_model_mappings( if prefixed_id not in aliases: aliases.append(prefixed_id) + # Register forwarded_model_id as a routable alias + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + aliases.append(model_to_use.forwarded_model_id) + for alias in aliases: _add_candidate(alias, model_to_use, upstream_for_override) seen_model_provider.add(dedupe_key) diff --git a/routstr/auth.py b/routstr/auth.py index 808bd7c4..df761938 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -7,6 +7,7 @@ from datetime import datetime from typing import Optional from fastapi import HTTPException +from sqlalchemy import case from sqlalchemy.exc import IntegrityError from sqlmodel import col, select, update @@ -22,6 +23,7 @@ from .payment.cost_calculation import ( from .wallet import credit_balance, deserialize_token_from_string logger = get_logger(__name__) +payments_logger = get_logger("routstr.payments") # Routstr platform fee constants ROUTSTR_FEE_PERCENT: float = 2.1 @@ -590,6 +592,18 @@ async def pay_for_request( "total_requests": billing_key.total_requests, }, ) + payments_logger.info( + "RESERVE", + extra={ + "event": "reserve", + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "cost_reserved": cost_per_request, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + "total_spent": billing_key.total_spent, + }, + ) return cost_per_request @@ -641,6 +655,17 @@ async def revert_pay_for_request( await session.refresh(billing_key) if billing_key.hashed_key != key.hashed_key: await session.refresh(key) + payments_logger.info( + "REVERT", + extra={ + "event": "revert", + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "cost_reverted": cost_per_request, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + }, + ) return True @@ -745,11 +770,32 @@ async def adjust_payment_for_tokens( }, ) # Finalize by releasing reservation and charging max cost + if billing_key.reserved_balance < deducted_max_cost: + logger.error( + "reserved_balance below deducted_max_cost before MaxCost finalization — clamping to 0", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "reserved_balance": billing_key.reserved_balance, + "deducted_max_cost": deducted_max_cost, + "total_cost_msats": cost.total_msats, + "balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "model": model, + }, + ) + + safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) + finalize_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( - reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, + reserved_balance=safe_reserved, balance=col(ApiKey.balance) - cost.total_msats, total_spent=col(ApiKey.total_spent) + cost.total_msats, ) @@ -758,13 +804,17 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: + child_safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) .values( total_spent=col(ApiKey.total_spent) + cost.total_msats, - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=child_safe_reserved, ) ) await session.exec(child_stmt) # type: ignore[call-overload] @@ -800,6 +850,23 @@ async def adjust_payment_for_tokens( }, ) await _accumulate_fee(cost.total_msats) + payments_logger.info( + "FINALIZE", + extra={ + "event": "finalize", + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + "cost_reserved": deducted_max_cost, + "cost_charged": cost.total_msats, + "input_tokens": cost.input_tokens, + "output_tokens": cost.output_tokens, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + "total_spent": billing_key.total_spent, + "finalize_type": "max_cost", + }, + ) return cost.dict() case CostData() as cost: @@ -833,12 +900,32 @@ async def adjust_payment_for_tokens( "model": model, }, ) + if billing_key.reserved_balance < deducted_max_cost: + logger.error( + "reserved_balance below deducted_max_cost on exact-cost finalization — clamping to 0", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "reserved_balance": billing_key.reserved_balance, + "deducted_max_cost": deducted_max_cost, + "total_cost_msats": total_cost_msats, + "balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "model": model, + }, + ) + + exact_safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) + finalize_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=exact_safe_reserved, balance=col(ApiKey.balance) - total_cost_msats, total_spent=col(ApiKey.total_spent) + total_cost_msats, ) @@ -847,13 +934,17 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: + child_exact_safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) .values( total_spent=col(ApiKey.total_spent) + total_cost_msats, - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=child_exact_safe_reserved, ) ) await session.exec(child_stmt) # type: ignore[call-overload] @@ -863,44 +954,55 @@ async def adjust_payment_for_tokens( if billing_key.hashed_key != key.hashed_key: await session.refresh(key) await _accumulate_fee(total_cost_msats) - return cost.dict() - - # this should never happen why do we handle this??? - if cost_difference > 0: - # Need to charge more than reserved, finalize by releasing reservation and charging total - logger.info( - "Additional charge required for token usage", + payments_logger.info( + "FINALIZE", extra={ + "event": "finalize", "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", - "additional_charge": cost_difference, - "current_balance": billing_key.balance, - "sufficient_balance": billing_key.balance >= cost_difference, "model": model, + "cost_reserved": deducted_max_cost, + "cost_charged": total_cost_msats, + "input_tokens": cost.input_tokens, + "output_tokens": cost.output_tokens, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + "total_spent": billing_key.total_spent, + "finalize_type": "exact", }, ) + return cost.dict() + + # actual cost exceeded discounted reservation (due to tolerance_percentage) + if cost_difference > 0: + # Always release the reservation and charge min(actual_cost, balance). + # Using a CASE expression makes this a single atomic UPDATE — no + # multi-level fallback needed and balance can never go negative. + chargeable = case( + (col(ApiKey.balance) >= total_cost_msats, total_cost_msats), + else_=col(ApiKey.balance), + ) finalize_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, - balance=col(ApiKey.balance) - total_cost_msats, - total_spent=col(ApiKey.total_spent) + total_cost_msats, + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, + balance=col(ApiKey.balance) - chargeable, + total_spent=col(ApiKey.total_spent) + chargeable, ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] - # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( - total_spent=col(ApiKey.total_spent) + total_cost_msats, - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, + total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats), ) ) await session.exec(child_stmt) # type: ignore[call-overload] @@ -908,11 +1010,10 @@ async def adjust_payment_for_tokens( await session.commit() if result.rowcount: - cost.total_msats = total_cost_msats await session.refresh(billing_key) if billing_key.hashed_key != key.hashed_key: await session.refresh(key) - + cost.total_msats = total_cost_msats logger.info( "Finalized payment with additional charge", extra={ @@ -924,9 +1025,28 @@ async def adjust_payment_for_tokens( }, ) await _accumulate_fee(total_cost_msats) + payments_logger.info( + "FINALIZE", + extra={ + "event": "finalize", + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + "cost_reserved": deducted_max_cost, + "cost_charged": total_cost_msats, + "input_tokens": cost.input_tokens, + "output_tokens": cost.output_tokens, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + "total_spent": billing_key.total_spent, + "finalize_type": "overrun", + }, + ) else: + # Guard fired: reservation was already released by a concurrent + # finalization for this key. Nothing left to do. logger.warning( - "Failed to finalize additional charge - releasing reservation", + "Finalization skipped - reservation already released", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", @@ -934,7 +1054,6 @@ async def adjust_payment_for_tokens( "model": model, }, ) - await release_reservation_only() else: # Refund some of the base cost refund = abs(cost_difference) @@ -949,12 +1068,33 @@ async def adjust_payment_for_tokens( }, ) + if billing_key.reserved_balance < deducted_max_cost: + logger.error( + "reserved_balance below deducted_max_cost on refund finalization — clamping to 0", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "reserved_balance": billing_key.reserved_balance, + "deducted_max_cost": deducted_max_cost, + "total_cost_msats": total_cost_msats, + "refund_amount": refund, + "balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "model": model, + }, + ) + + refund_safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) + refund_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=refund_safe_reserved, balance=col(ApiKey.balance) - total_cost_msats, total_spent=col(ApiKey.total_spent) + total_cost_msats, ) @@ -963,13 +1103,17 @@ async def adjust_payment_for_tokens( # Also update total_spent and reserved_balance on the child key if it's different if billing_key.hashed_key != key.hashed_key: + child_refund_safe_reserved = case( + (col(ApiKey.reserved_balance) >= deducted_max_cost, + col(ApiKey.reserved_balance) - deducted_max_cost), + else_=0, + ) child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) .values( total_spent=col(ApiKey.total_spent) + total_cost_msats, - reserved_balance=col(ApiKey.reserved_balance) - - deducted_max_cost, + reserved_balance=child_refund_safe_reserved, ) ) await session.exec(child_stmt) # type: ignore[call-overload] @@ -1007,6 +1151,24 @@ async def adjust_payment_for_tokens( }, ) await _accumulate_fee(total_cost_msats) + payments_logger.info( + "FINALIZE", + extra={ + "event": "finalize", + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + "cost_reserved": deducted_max_cost, + "cost_charged": total_cost_msats, + "refunded": refund, + "input_tokens": cost.input_tokens, + "output_tokens": cost.output_tokens, + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, + "total_spent": billing_key.total_spent, + "finalize_type": "refund", + }, + ) return cost.dict() diff --git a/routstr/balance.py b/routstr/balance.py index 9adbd9b7..ff892ad3 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -5,11 +5,18 @@ from time import monotonic from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException +from fastapi.responses import JSONResponse from pydantic import BaseModel from sqlmodel import select from .auth import get_billing_key, validate_bearer_key -from .core.db import ApiKey, AsyncSession, CashuTransaction, get_session +from .core.db import ( + ApiKey, + AsyncSession, + CashuTransaction, + get_session, + store_cashu_transaction, +) from .core.logging import get_logger from .core.settings import settings from .lightning import lightning_router @@ -204,19 +211,57 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: _refund_cache[key] = (expiry, value) -@router.post("/refund") +@router.post("/refund", response_model=None) async def refund_wallet_endpoint( - authorization: Annotated[str, Header(...)], + authorization: Annotated[str | None, Header()] = None, + x_cashu: Annotated[str | None, Header()] = None, session: AsyncSession = Depends(get_session), -) -> dict[str, str]: - if not authorization.startswith("Bearer "): +) -> JSONResponse | dict[str, str]: + if x_cashu: + # Find the "in" transaction by the original payment token + in_tx_result = await session.exec( + select(CashuTransaction).where( + CashuTransaction.token == x_cashu, + CashuTransaction.type == "in", + ) + ) + in_tx = in_tx_result.first() + if in_tx is None: + raise HTTPException(status_code=404, detail="Refund not found") + + # Use the request_id to find the associated "out" (refund) transaction + if in_tx.request_id is None: + raise HTTPException(status_code=404, detail="Refund not found") + + out_tx_result = await session.exec( + select(CashuTransaction).where( + CashuTransaction.request_id == in_tx.request_id, + CashuTransaction.type == "out", + ) + ) + out_tx = out_tx_result.first() + if out_tx is None: + raise HTTPException(status_code=404, detail="Refund not found") + if out_tx.swept: + raise HTTPException(status_code=410, detail="Refund has been swept") + + out_tx.collected = True + session.add(out_tx) + await session.commit() + body: dict[str, str] = {"token": out_tx.token} + if out_tx.unit == "sat": + body["sats"] = str(out_tx.amount) + else: + body["msats"] = str(out_tx.amount) + return JSONResponse(content=body, headers={"X-Cashu": out_tx.token}) + + if authorization is None or not authorization.startswith("Bearer "): raise HTTPException( status_code=401, detail="Invalid authorization. Use 'Bearer ' or 'Bearer '", ) bearer_value: str = authorization[7:] - key: ApiKey = await validate_bearer_key(bearer_value, session) if key.total_balance <= 0: @@ -229,6 +274,12 @@ async def refund_wallet_endpoint( detail="Cannot refund child key. Please refund the parent key instead.", ) + if key.reserved_balance > 0: + raise HTTPException( + status_code=400, + detail="Cannot refund key. There are ongoing requests for this api key.", + ) + remaining_balance_msats: int = key.total_balance if key.refund_currency == "sat": @@ -265,6 +316,17 @@ async def refund_wallet_endpoint( else: result["msats"] = str(remaining_balance_msats) + if "token" in result: + logger.info( + "refund_wallet_endpoint: cashu token issued", + extra={ + "path": "/v1/wallet/refund", + "token": result["token"], + "amount": remaining_balance, + "currency": key.refund_currency or "sat", + }, + ) + except HTTPException: # Re-raise HTTP exceptions (like 400 for balance too small) raise @@ -283,11 +345,34 @@ async def refund_wallet_endpoint( await _refund_cache_set(bearer_value, result) + previous_reserved_balance = key.reserved_balance key.balance = 0 key.reserved_balance = 0 session.add(key) await session.commit() + if "token" in result: + try: + await store_cashu_transaction( + token=result["token"], + amount=remaining_balance, + unit=key.refund_currency or "sat", + mint_url=key.refund_mint_url, + typ="out", + collected=False, + source="apikey", + ) + except Exception: + pass # store_cashu_transaction already logs + + logger.info( + "refund_wallet_endpoint: refund successful", + extra={ + "refunded_msats": remaining_balance_msats, + "previous_reserved_balance": previous_reserved_balance, + }, + ) + return result diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 46589164..d87120e2 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -295,6 +295,7 @@ class ModelCreate(BaseModel): canonical_slug: str | None = None alias_ids: list[str] | None = None enabled: bool = True + forwarded_model_id: str | None = None @admin_router.post( @@ -339,6 +340,7 @@ async def upsert_provider_model( json.dumps(payload.alias_ids) if payload.alias_ids else None ) existing_row.enabled = payload.enabled + existing_row.forwarded_model_id = payload.forwarded_model_id or payload.id session.add(existing_row) await session.commit() @@ -371,6 +373,7 @@ async def upsert_provider_model( ), upstream_provider_id=provider_id, enabled=payload.enabled, + forwarded_model_id=payload.forwarded_model_id or payload.id, ) session.add(row) await session.commit() diff --git a/routstr/core/db.py b/routstr/core/db.py index c97e6b8e..c5fc990a 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -10,6 +10,7 @@ from alembic import command from alembic.config import Config from alembic.util.exc import CommandError from sqlalchemy import UniqueConstraint +from sqlalchemy.exc import OperationalError from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlmodel import Field, Relationship, SQLModel, col, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -105,6 +106,10 @@ class ModelRow(SQLModel, table=True): # type: ignore default=None, description="JSON array of model alias IDs" ) enabled: bool = Field(default=True, description="Whether this model is enabled") + forwarded_model_id: str | None = Field( + default=None, + description="Model ID to use when forwarding requests to upstream provider. Defaults to id if not set.", + ) upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") @@ -150,6 +155,10 @@ class CashuTransaction(SQLModel, table=True): # type: ignore ) collected: bool = Field(default=False) swept: bool = Field(default=False) + source: str = Field( + default="x-cashu", + description="Payment source: x-cashu or apikey", + ) async def store_cashu_transaction( @@ -161,6 +170,7 @@ async def store_cashu_transaction( request_id: str | None = None, collected: bool = False, created_at: int | None = None, + source: str = "x-cashu", ) -> None: try: async with create_session() as session: @@ -173,6 +183,7 @@ async def store_cashu_transaction( request_id=request_id, collected=collected, created_at=created_at or int(time.time()), + source=source, ) session.add(tx) await session.commit() @@ -371,6 +382,17 @@ def run_migrations() -> None: command.stamp(alembic_cfg, "head") else: raise + except OperationalError as e: + if "duplicate column name" in str(e).lower(): + logger.warning( + "Migration hit a column that already exists (likely added via " + "create_all on another branch). Stamping to current head.", + extra={"error": str(e)}, + ) + _clear_alembic_version() + command.stamp(alembic_cfg, "head") + else: + raise logger.info("Database migrations completed successfully") diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 1810f055..4942f3fd 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -38,11 +38,6 @@ class LoggingMiddleware(BaseHTTPMiddleware): except Exception: pass - # Extract request info - client_host = None - if request.client: - client_host = request.client.host - # Log incoming request logger.info( "Incoming request", @@ -51,7 +46,6 @@ class LoggingMiddleware(BaseHTTPMiddleware): "method": request.method, "path": request.url.path, "query_params": dict(request.query_params), - "client_host": client_host, "headers": { k: v for k, v in request.headers.items() @@ -100,7 +94,6 @@ class LoggingMiddleware(BaseHTTPMiddleware): "path": request.url.path, "status_code": response.status_code, "duration_ms": round(duration * 1000, 2), - "client_host": client_host, }, ) if hasattr(response, "headers"): @@ -120,7 +113,6 @@ class LoggingMiddleware(BaseHTTPMiddleware): "method": request.method, "path": request.url.path, "duration_ms": round(duration * 1000, 2), - "client_host": client_host, "error": str(e), "error_type": type(e).__name__, }, diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 027cc0ba..d05364f4 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -74,7 +74,7 @@ class Settings(BaseSettings): enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") - refund_sweep_ttl_seconds: int = Field(default=86400, env="REFUND_SWEEP_TTL_SECONDS") + refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 2ed7e4b8..44330620 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -109,12 +109,20 @@ async def calculate_cost( # todo: can be sync ) usd_cost = 0.0 + input_usd = 0.0 + output_usd = 0.0 - # Prioritize cost_details.upstream_inference_cost if "cost_details" in usage_data: usd_cost = float( usage_data["cost_details"].get("upstream_inference_cost", 0) or 0 ) + input_usd = float( + usage_data["cost_details"].get("upstream_inference_prompt_cost", 0) or 0 + ) + output_usd = float( + usage_data["cost_details"].get("upstream_inference_completions_cost", 0) + or 0 + ) # Fallback to cost field if upstream_inference_cost is 0 if usd_cost == 0 and "cost" in usage_data: @@ -123,12 +131,34 @@ async def calculate_cost( # todo: can be sync except Exception: pass + MSATS_PER_1K_INPUT_TOKENS: float = ( + float(settings.fixed_per_1k_input_tokens) * 1000.0 + ) + MSATS_PER_1K_OUTPUT_TOKENS: float = ( + float(settings.fixed_per_1k_output_tokens) * 1000.0 + ) + if usd_cost > 0: try: sats_per_usd = 1.0 / sats_usd_price() cost_in_sats = usd_cost * sats_per_usd cost_in_msats = math.ceil(cost_in_sats * 1000) + input_msats = 0 + output_msats = 0 + + if input_usd > 0 or output_usd > 0: + input_msats = int((input_usd * sats_per_usd) * 1000) + output_msats = int((output_usd * sats_per_usd) * 1000) + else: + total_tokens = input_tokens + output_tokens + if total_tokens > 0: + input_ratio = input_tokens / total_tokens + input_msats = int(cost_in_msats * input_ratio) + output_msats = cost_in_msats - input_msats + else: + output_msats = cost_in_msats + logger.info( "Using cost from usage data/details", extra={ @@ -140,9 +170,9 @@ async def calculate_cost( # todo: can be sync ) return CostData( - base_msats=-1, - input_msats=-1, # Cost field doesn't break down by token type - output_msats=-1, + base_msats=0, + input_msats=input_msats, + output_msats=output_msats, total_msats=cost_in_msats, total_usd=usd_cost, input_tokens=input_tokens, @@ -159,13 +189,6 @@ async def calculate_cost( # todo: can be sync ) # Fall through to token-based calculation - MSATS_PER_1K_INPUT_TOKENS: float = ( - float(settings.fixed_per_1k_input_tokens) * 1000.0 - ) - MSATS_PER_1K_OUTPUT_TOKENS: float = ( - float(settings.fixed_per_1k_output_tokens) * 1000.0 - ) - if not settings.fixed_pricing: response_model = response_data.get("model", "") logger.debug( @@ -231,10 +254,10 @@ async def calculate_cost( # todo: can be sync output_tokens=output_tokens, ) - input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3) + calc_input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3) - output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3) - token_based_cost = math.ceil(input_msats + output_msats) + calc_output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3) + token_based_cost = math.ceil(calc_input_msats + calc_output_msats) total_usd = (token_based_cost / 1000.0) * sats_usd_price() logger.info( @@ -242,8 +265,8 @@ async def calculate_cost( # todo: can be sync extra={ "input_tokens": input_tokens, "output_tokens": output_tokens, - "input_cost_msats": input_msats, - "output_cost_msats": output_msats, + "input_cost_msats": calc_input_msats, + "output_cost_msats": calc_output_msats, "total_cost_msats": token_based_cost, "total_usd": total_usd, "model": response_data.get("model", "unknown"), @@ -252,8 +275,8 @@ async def calculate_cost( # todo: can be sync return CostData( base_msats=0, - input_msats=int(input_msats), - output_msats=int(output_msats), + input_msats=int(calc_input_msats), + output_msats=int(calc_output_msats), total_msats=token_based_cost, total_usd=total_usd, input_tokens=input_tokens, diff --git a/routstr/payment/models.py b/routstr/payment/models.py index ae780153..95d7ad89 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -60,6 +60,7 @@ class Model(BaseModel): upstream_provider_id: int | str | None = None canonical_slug: str | None = None alias_ids: list[str] | None = None + forwarded_model_id: str | None = None def __hash__(self) -> int: return hash(self.id) @@ -177,6 +178,7 @@ def _row_to_model( upstream_provider_id=row.upstream_provider_id, canonical_slug=getattr(row, "canonical_slug", None), alias_ids=json.loads(row.alias_ids) if row.alias_ids else None, + forwarded_model_id=getattr(row, "forwarded_model_id", None) or row.id, ) if apply_provider_fee: @@ -329,6 +331,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model: upstream_provider_id=model.upstream_provider_id, canonical_slug=model.canonical_slug, alias_ids=model.alias_ids, + forwarded_model_id=model.forwarded_model_id, ) except Exception as e: logger.error( @@ -411,4 +414,10 @@ async def models(session: AsyncSession = Depends(get_session)) -> dict: from ..proxy import get_unique_models items = get_unique_models() - return {"data": items} + data = [] + for model in items: + m = model.dict() + if model.forwarded_model_id: + m["id"] = model.forwarded_model_id + data.append(m) + return {"data": data} diff --git a/routstr/proxy.py b/routstr/proxy.py index c884c896..ddab1c31 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -207,7 +207,7 @@ async def proxy( elif auth := headers.get("authorization", None): key = await get_bearer_token_key( - headers, path, session, auth, max_cost_for_model + headers, path, session, auth, max_cost_for_model, model_id ) else: @@ -387,7 +387,12 @@ async def proxy( async def get_bearer_token_key( - headers: dict, path: str, session: AsyncSession, auth: str, min_cost: int = 0 + headers: dict, + path: str, + session: AsyncSession, + auth: str, + min_cost: int = 0, + model_id: str = "unknown", ) -> ApiKey: """Handle bearer token authentication proxy requests.""" parts = auth.split() @@ -457,11 +462,13 @@ async def get_bearer_token_key( except Exception as e: key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key logger.error( - f"Bearer token validation failed: {type(e).__name__}: {e} path={path} key={key_preview!r}", + f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}", extra={ "error": str(e), "error_type": type(e).__name__, "path": path, + "model_id": model_id, + "min_cost_msat": min_cost, "bearer_key_preview": key_preview, }, ) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index f7517c4b..b506e388 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -5,6 +5,7 @@ import hashlib import json import re import traceback +import uuid from collections.abc import AsyncGenerator from typing import Mapping @@ -452,6 +453,7 @@ class BaseUpstreamProvider: key: ApiKey, max_cost_for_model: int, background_tasks: BackgroundTasks, + requested_model: str | None = None, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -516,16 +518,32 @@ class BaseUpstreamProvider: continue try: - obj = json.loads(part) - if isinstance(obj, dict): - if obj.get("model"): - last_model_seen = str(obj.get("model")) - - if isinstance(obj.get("usage"), dict): - # Hold this chunk back to merge cost later - usage_chunk_data = obj + # Only parse if it looks like a JSON object to avoid SSE control messages or partials + if part.strip().startswith(b"{") and part.strip().endswith( + b"}" + ): + obj = json.loads(part) + if isinstance(obj, dict): + if obj.get("model"): + last_model_seen = str(obj.get("model")) + if requested_model: + obj["model"] = requested_model + if ( + "id" not in obj + or not isinstance(obj["id"], str) + or obj["id"] == "existing-id" + ): + if not hasattr(self, "_current_stream_id"): + self._current_stream_id = ( + f"chatcmpl-{uuid.uuid4()}" + ) + obj["id"] = self._current_stream_id + if isinstance(obj.get("usage"), dict): + usage_chunk_data = obj + continue + yield b"data: " + json.dumps(obj).encode() + b"\n\n" continue - except json.JSONDecodeError: + except Exception: pass prefix = ( @@ -620,6 +638,7 @@ class BaseUpstreamProvider: key: ApiKey, session: AsyncSession, deducted_max_cost: int, + requested_model: str | None = None, ) -> Response: """Handle non-streaming chat completion responses with token usage tracking and cost adjustment. @@ -655,6 +674,11 @@ class BaseUpstreamProvider: }, ) + if requested_model: + response_json["model"] = requested_model + if "id" not in response_json or not isinstance(response_json["id"], str): + response_json["id"] = f"chatcmpl-{uuid.uuid4()}" + cost_data = await adjust_payment_for_tokens( key, response_json, session, deducted_max_cost ) @@ -714,6 +738,8 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } + if requested_model: + response_json["model"] = requested_model return Response( content=json.dumps(response_json).encode(), status_code=response.status_code, @@ -744,7 +770,11 @@ class BaseUpstreamProvider: raise async def handle_streaming_responses_completion( - self, response: httpx.Response, key: ApiKey, max_cost_for_model: int + self, + response: httpx.Response, + key: ApiKey, + max_cost_for_model: int, + requested_model: str | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -814,6 +844,8 @@ class BaseUpstreamProvider: if isinstance(obj, dict): if obj.get("model"): last_model_seen = str(obj.get("model")) + if requested_model: + obj["model"] = requested_model # Track reasoning tokens for Responses API if usage := obj.get("usage", {}): @@ -934,6 +966,7 @@ class BaseUpstreamProvider: key: ApiKey, session: AsyncSession, deducted_max_cost: int, + requested_model: str | None = None, ) -> Response: """Handle non-streaming Responses API responses with token usage tracking and cost adjustment. @@ -972,6 +1005,11 @@ class BaseUpstreamProvider: }, ) + if requested_model: + response_json["model"] = requested_model + if "id" not in response_json or not isinstance(response_json["id"], str): + response_json["id"] = f"chatcmpl-{uuid.uuid4()}" + cost_data = await adjust_payment_for_tokens( key, response_json, session, deducted_max_cost ) @@ -1031,6 +1069,8 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } + if requested_model: + response_json["model"] = requested_model return Response( content=json.dumps(response_json).encode(), status_code=response.status_code, @@ -1098,6 +1138,174 @@ class BaseUpstreamProvider: }, ) + async def handle_streaming_messages_completion( + self, response: httpx.Response, key: ApiKey, max_cost_for_model: int + ) -> StreamingResponse: + async def stream_with_cost( + max_cost_for_model: int, + ) -> AsyncGenerator[bytes, None]: + stored_chunks: list[bytes] = [] + usage_finalized: bool = False + last_model_seen: str | None = None + input_tokens: int = 0 + output_tokens: int = 0 + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + fallback: dict = { + "model": last_model_seen or "unknown", + "usage": None, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, fallback, new_session, max_cost_for_model + ) + usage_finalized = True + return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception: + usage_finalized = True + return None + + try: + async for chunk in response.aiter_bytes(): + stored_chunks.append(chunk) + try: + decoded_chunk = chunk.decode("utf-8", errors="ignore") + for line in decoded_chunk.split("\n"): + if line.startswith("data: "): + try: + data = json.loads(line[6:]) + if isinstance(data, dict): + msg = data.get("message", {}) + if msg and msg.get("model"): + last_model_seen = str(msg.get("model")) + + if usage := msg.get("usage"): + input_tokens += usage.get("input_tokens", 0) + output_tokens += usage.get( + "output_tokens", 0 + ) + + if usage := data.get("usage"): + input_tokens += usage.get("input_tokens", 0) + output_tokens += usage.get( + "output_tokens", 0 + ) + except json.JSONDecodeError: + pass + except Exception: + pass + + yield chunk + + usage_data = { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + } + + if input_tokens > 0 or output_tokens > 0: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if fresh_key: + try: + combined_data = { + "model": last_model_seen or "unknown", + "usage": usage_data, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, + combined_data, + new_session, + max_cost_for_model, + ) + usage_finalized = True + yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception: + pass + + if not usage_finalized: + maybe_cost_event = await finalize_without_usage() + if maybe_cost_event is not None: + yield maybe_cost_event + + except httpx.ReadError: + if not usage_finalized: + await finalize_without_usage() + # Upstream dropped the connection mid-stream; response already started, swallow silently + except Exception: + if not usage_finalized: + await finalize_without_usage() + raise + finally: + if not usage_finalized: + await finalize_without_usage() + + response_headers = dict(response.headers) + response_headers.pop("content-encoding", None) + response_headers.pop("content-length", None) + + return StreamingResponse( + stream_with_cost(max_cost_for_model), + status_code=response.status_code, + headers=response_headers, + ) + + async def handle_non_streaming_messages_completion( + self, + response: httpx.Response, + key: ApiKey, + session: AsyncSession, + deducted_max_cost: int, + path: str, + ) -> Response: + try: + content = await response.aread() + response_json = json.loads(content) + + if path.endswith("count_tokens") and "usage" not in response_json: + input_tokens = response_json.get("input_tokens", 0) + response_json["usage"] = {"input_tokens": input_tokens} + + cost_data = await adjust_payment_for_tokens( + key, response_json, session, deducted_max_cost + ) + response_json["cost"] = cost_data + + allowed_headers = { + "content-type", + "cache-control", + "date", + "vary", + "access-control-allow-origin", + "access-control-allow-methods", + "access-control-allow-headers", + "access-control-allow-credentials", + "access-control-expose-headers", + "access-control-max-age", + } + + response_headers = { + k: v + for k, v in response.headers.items() + if k.lower() in allowed_headers + } + + return Response( + content=json.dumps(response_json).encode(), + status_code=response.status_code, + headers=response_headers, + media_type="application/json", + ) + except Exception: + raise + async def forward_request( self, request: Request, @@ -1126,6 +1334,10 @@ class BaseUpstreamProvider: path = self.normalize_request_path(path, model_obj) url = self.build_request_url(path, model_obj) + original_model_id = ( + (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + ) + transformed_body = self.prepare_request_body(request_body, model_obj) logger.info( @@ -1197,7 +1409,54 @@ class BaseUpstreamProvider: await client.aclose() return mapped_error - if path.endswith("chat/completions") or path.endswith("embeddings"): + if ( + path.endswith("chat/completions") + or path.endswith("embeddings") + or path.endswith("messages") + or path.endswith("messages/count_tokens") + ): + if path.endswith("messages"): + client_wants_streaming = False + if request_body: + try: + request_data = json.loads(request_body) + client_wants_streaming = request_data.get("stream", False) + except json.JSONDecodeError: + pass + + content_type = response.headers.get("content-type", "") + upstream_is_streaming = "text/event-stream" in content_type + is_streaming = client_wants_streaming and upstream_is_streaming + + if is_streaming and response.status_code == 200: + result = await self.handle_streaming_messages_completion( + response, key, max_cost_for_model + ) + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + result.background = background_tasks + return result + + if response.status_code == 200: + try: + return await self.handle_non_streaming_messages_completion( + response, key, session, max_cost_for_model, path + ) + finally: + await response.aclose() + await client.aclose() + + if path.endswith("messages/count_tokens"): + if response.status_code == 200: + try: + return await self.handle_non_streaming_messages_completion( + response, key, session, max_cost_for_model, path + ) + finally: + await response.aclose() + await client.aclose() + if path.endswith("chat/completions"): client_wants_streaming = False if request_body: @@ -1237,7 +1496,11 @@ class BaseUpstreamProvider: background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) result = await self.handle_streaming_chat_completion( - response, key, max_cost_for_model, background_tasks + response, + key, + max_cost_for_model, + background_tasks, + requested_model=original_model_id, ) result.background = background_tasks return result @@ -1246,7 +1509,11 @@ class BaseUpstreamProvider: if response.status_code == 200: try: return await self.handle_non_streaming_chat_completion( - response, key, session, max_cost_for_model + response, + key, + session, + max_cost_for_model, + requested_model=original_model_id, ) finally: await response.aclose() @@ -1361,6 +1628,10 @@ class BaseUpstreamProvider: path = self.normalize_request_path(path, model_obj) url = self.build_request_url(path, model_obj) + original_model_id = ( + (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + ) + transformed_body = self.prepare_responses_request_body(request_body, model_obj) logger.info( @@ -1447,7 +1718,10 @@ class BaseUpstreamProvider: if is_streaming and response.status_code == 200: result = await self.handle_streaming_responses_completion( - response, key, max_cost_for_model + response, + key, + max_cost_for_model, + requested_model=original_model_id, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -1458,7 +1732,11 @@ class BaseUpstreamProvider: if response.status_code == 200: try: return await self.handle_non_streaming_responses_completion( - response, key, session, max_cost_for_model + response, + key, + session, + max_cost_for_model, + requested_model=original_model_id, ) finally: await response.aclose() @@ -1825,11 +2103,25 @@ class BaseUpstreamProvider: if line.startswith("data: "): try: data_json = json.loads(line[6:]) + # OpenAI format: usage and model at top level if "usage" in data_json: usage_data = data_json["usage"] - model = data_json.get("model") + model = data_json.get("model") or model elif "model" in data_json and not model: model = data_json["model"] + # Anthropic format: model and input usage inside "message" key + if "message" in data_json: + msg = data_json["message"] + if not model and msg.get("model"): + model = msg["model"] + if msg.get("usage") and not usage_data: + usage_data = msg["usage"] + elif msg.get("usage") and usage_data: + # Merge: message_start has input_tokens, message_delta has output_tokens + merged = dict(usage_data) + for k, v in msg["usage"].items(): + merged[k] = merged.get(k, 0) + v + usage_data = merged except json.JSONDecodeError: continue @@ -1870,7 +2162,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -1906,6 +2201,19 @@ class BaseUpstreamProvider: }, ) + if cost_data: + for i, line in enumerate(lines): + if line.startswith("data: "): + try: + data_json = json.loads(line[6:]) + if "usage" in data_json and data_json["usage"]: + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) + lines[i] = "data: " + json.dumps(data_json) + except json.JSONDecodeError: + pass + async def generate() -> AsyncGenerator[bytes, None]: for line in lines: yield (line + "\n").encode("utf-8") @@ -1950,6 +2258,9 @@ class BaseUpstreamProvider: response_json = json.loads(content_str) cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model) + if cost_data and "usage" in response_json: + response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if not cost_data: logger.error( "Failed to calculate cost for response", @@ -1999,7 +2310,10 @@ class BaseUpstreamProvider: if refund_amount > 0: refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -2016,7 +2330,7 @@ class BaseUpstreamProvider: ) return Response( - content=content_str, + content=json.dumps(response_json), status_code=response.status_code, headers=response_headers, media_type="application/json", @@ -2054,7 +2368,6 @@ class BaseUpstreamProvider: extra={ "original_amount": amount, "refund_amount": emergency_refund, - "deduction": 60, }, ) @@ -2229,7 +2542,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - amount - 60, unit, mint, payment_token_hash, + amount, + unit, + mint, + payment_token_hash, request_id=getattr(request.state, "request_id", None), ) @@ -2262,9 +2578,14 @@ class BaseUpstreamProvider: error_response.headers["X-Cashu"] = refund_token return error_response - if path.endswith("chat/completions") or path.endswith("embeddings"): + if ( + path.endswith("chat/completions") + or path.endswith("embeddings") + or path.endswith("messages") + or path.endswith("messages/count_tokens") + ): logger.debug( - "Processing completion/embeddings response", + "Processing completion/embeddings/messages response", extra={"path": path, "amount": amount, "unit": unit}, ) @@ -2516,7 +2837,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - amount - 60, unit, mint, payment_token_hash, + amount, + unit, + mint, + payment_token_hash, request_id=getattr(request.state, "request_id", None), ) @@ -2780,7 +3104,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -2816,6 +3143,19 @@ class BaseUpstreamProvider: }, ) + if cost_data: + for i, line in enumerate(lines): + if line.startswith("data: "): + try: + data_json = json.loads(line[6:]) + if "usage" in data_json and data_json["usage"]: + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) + lines[i] = "data: " + json.dumps(data_json) + except json.JSONDecodeError: + pass + async def generate() -> AsyncGenerator[bytes, None]: for line in lines: yield (line + "\\n").encode("utf-8") @@ -2848,6 +3188,9 @@ class BaseUpstreamProvider: response_json = json.loads(content_str) cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model) + if cost_data and "usage" in response_json: + response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if not cost_data: logger.error( "Failed to calculate cost for Responses API response", @@ -2897,7 +3240,10 @@ class BaseUpstreamProvider: if refund_amount > 0: refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -2914,7 +3260,7 @@ class BaseUpstreamProvider: ) return Response( - content=content_str, + content=json.dumps(response_json), status_code=response.status_code, headers=response_headers, media_type="application/json", @@ -2952,7 +3298,6 @@ class BaseUpstreamProvider: extra={ "original_amount": amount, "refund_amount": emergency_refund, - "deduction": 60, }, ) @@ -3105,6 +3450,7 @@ class BaseUpstreamProvider: upstream_provider_id=model.upstream_provider_id, canonical_slug=model.canonical_slug, alias_ids=model.alias_ids, + forwarded_model_id=model.forwarded_model_id, ) ( @@ -3128,6 +3474,7 @@ class BaseUpstreamProvider: upstream_provider_id=model.upstream_provider_id, canonical_slug=model.canonical_slug, alias_ids=model.alias_ids, + forwarded_model_id=model.forwarded_model_id, ) async def fetch_models(self) -> list[Model]: diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 49a300b9..a6890567 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -1,9 +1,9 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Optional import httpx -from pydantic import BaseModel +from pydantic import BaseModel, Field from ..core.logging import get_logger from ..payment.models import Architecture, Model, Pricing, async_fetch_openrouter_models @@ -16,18 +16,20 @@ logger = get_logger(__name__) class PPQAIModelPricing(BaseModel): - ui: dict[str, float] - api: dict[str, float] + ui: Optional[dict[str, float]] = None + api: Optional[dict[str, float]] = None + input_per_1M_tokens: Optional[float] = Field(None, alias="input_per_1M_tokens") + output_per_1M_tokens: Optional[float] = Field(None, alias="output_per_1M_tokens") class PPQAIModel(BaseModel): id: str - provider: str + provider: Optional[str] = None name: str created_at: int context_length: int pricing: PPQAIModelPricing - popular: bool + popular: bool = False class PPQAIUpstreamProvider(BaseUpstreamProvider): @@ -134,31 +136,54 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) if or_model: - if input_price := ppqai_model.pricing.api.get( - "input_per_1M" - ): + input_price = None + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get( + "input_per_1M" + ) + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens + + if input_price is not None: or_model.pricing.prompt = input_price / 1_000_000 - if output_price := ppqai_model.pricing.api.get( - "output_per_1M" - ): + + output_price = None + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get( + "output_per_1M" + ) + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens + + if output_price is not None: or_model.pricing.completion = output_price / 1_000_000 + if cl := ppqai_model.context_length: or_model.context_length = cl models.append(or_model) else: - input_price = ppqai_model.pricing.api.get( - "input_per_1M", 0.0 - ) - output_price = ppqai_model.pricing.api.get( - "output_per_1M", 0.0 - ) + input_price = 0.0 + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get( + "input_per_1M", 0.0 + ) + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens + + output_price = 0.0 + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get( + "output_per_1M", 0.0 + ) + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens models.append( Model( id=ppqai_model.id, name=ppqai_model.name, created=ppqai_model.created_at // 1000, - description=f"{ppqai_model.provider} model", + description=f"{ppqai_model.provider or 'PPQ.AI'} model", context_length=ppqai_model.context_length, architecture=Architecture( modality="text->text", diff --git a/routstr/wallet.py b/routstr/wallet.py index 0707e2f2..80a823af 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -155,6 +155,20 @@ async def swap_to_primary_mint( raise ValueError("Invalid unit") primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit) + # If the token is already from the primary mint, we don't need to swap + # and we definitely don't want to calculate or pay fees. + if token_obj.mint == settings.primary_mint: + logger.info( + "swap_to_primary_mint: token already on primary mint, skipping swap", + extra={ + "mint": token_obj.mint, + "amount": token_amount, + "unit": token_obj.unit, + }, + ) + await token_wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True) + return token_amount, token_obj.unit, token_obj.mint + minted_amount = await _calculate_swap_amount( amount_msat, token_obj.unit, @@ -535,7 +549,7 @@ async def periodic_refund_sweep() -> None: except Exception as e: error_msg = str(e).lower() if "already spent" in error_msg: - refund.swept = True + refund.collected = True session.add(refund) logger.info( "Refund already spent (client collected), marking swept", diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 1977101f..40ab9d53 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -380,6 +380,13 @@ async def integration_session( yield session +@pytest_asyncio.fixture +async def patched_db_engine(integration_engine: Any) -> AsyncGenerator[None, None]: + """Patch the global db engine so create_session() uses the test engine.""" + with patch("routstr.core.db.engine", integration_engine): + yield + + class DatabaseSnapshot: """Utility to capture and compare database states""" diff --git a/tests/integration/test_balance_negative_on_cost_overrun.py b/tests/integration/test_balance_negative_on_cost_overrun.py new file mode 100644 index 00000000..a304d91d --- /dev/null +++ b/tests/integration/test_balance_negative_on_cost_overrun.py @@ -0,0 +1,369 @@ +""" +Integration tests for the balance-goes-negative bug in adjust_payment_for_tokens. + +Root cause: when actual token cost exceeds the discounted reservation +(cost_difference > 0, caused by tolerance_percentage discounting the reservation), +the finalization UPDATE had no WHERE guard on balance, allowing balance to go negative. + +Fix: added `.where(col(ApiKey.balance) >= total_cost_msats)` so the UPDATE is a no-op +when balance is insufficient, then falls back to charging only deducted_max_cost. +""" + +import uuid +from unittest.mock import patch + +import pytest +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey +from routstr.payment.cost_calculation import CostData + + +def _make_key(balance: int, reserved: int) -> ApiKey: + return ApiKey( + hashed_key=f"test_{uuid.uuid4().hex}", + balance=balance, + reserved_balance=reserved, + total_spent=0, + total_requests=1, + ) + + +async def _refresh(session: AsyncSession, key: ApiKey) -> ApiKey: + await session.refresh(key) + return key + + +# --------------------------------------------------------------------------- +# Helper: build a CostData where token cost > deducted_max_cost +# --------------------------------------------------------------------------- + +def _cost_data(total_msats: int) -> CostData: + return CostData( + base_msats=0, + input_msats=total_msats // 2, + output_msats=total_msats - total_msats // 2, + total_msats=total_msats, + total_usd=0.0, + input_tokens=100, + output_tokens=100, + ) + + +# --------------------------------------------------------------------------- +# Test 1 — exact reproduction of the bug +# +# Setup: balance == deducted_max_cost (user has just enough for the reservation, +# nothing extra). Actual token cost is 1% higher (tolerance_percentage). +# +# Before fix: balance -= total_cost_msats → goes negative. +# After fix: WHERE balance >= total_cost_msats fails → fallback charges +# deducted_max_cost → balance reaches 0, never negative. +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_balance_never_negative_when_cost_exceeds_reservation( + integration_session: AsyncSession, +) -> None: + """Balance must not go negative when actual token cost > discounted reservation.""" + from routstr.auth import adjust_payment_for_tokens + + deducted_max_cost = 990 # reserved (1% below true max of 1000) + actual_token_cost = 1000 # actual cost at true max + + # User has balance exactly equal to the reservation — tight budget + key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost) + integration_session.add(key) + await integration_session.commit() + + response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}} + + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + + await _refresh(integration_session, key) + + assert key.balance >= 0, f"Balance went negative: {key.balance}" + assert key.reserved_balance >= 0, f"Reserved balance went negative: {key.reserved_balance}" + assert key.reserved_balance == 0, "Reservation must be fully released after finalization" + + +# --------------------------------------------------------------------------- +# Test 2 — balance is ZERO after the reservation is accounted for +# (absolute floor case) +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_balance_floor_at_zero_on_overrun( + integration_session: AsyncSession, +) -> None: + """When balance exactly covers deducted_max_cost and cost overruns, balance reaches 0 not negative.""" + from routstr.auth import adjust_payment_for_tokens + + deducted_max_cost = 500 + actual_token_cost = 550 # 10% overrun + + key = _make_key(balance=500, reserved=500) + integration_session.add(key) + await integration_session.commit() + + response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}} + + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + + await _refresh(integration_session, key) + + assert key.balance == 0, ( + f"Expected balance=0 (charged deducted_max_cost fallback), got {key.balance}" + ) + assert key.reserved_balance == 0, f"Reserved balance should be 0, got {key.reserved_balance}" + # Fallback charges deducted_max_cost + assert key.total_spent == deducted_max_cost, ( + f"Expected total_spent={deducted_max_cost}, got {key.total_spent}" + ) + + +# --------------------------------------------------------------------------- +# Test 3 — balance has enough room: full token cost should be charged +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_full_cost_charged_when_balance_sufficient_for_overrun( + integration_session: AsyncSession, +) -> None: + """When balance covers total_cost_msats, the full amount is charged (not just deducted_max_cost).""" + from routstr.auth import adjust_payment_for_tokens + + deducted_max_cost = 990 + actual_token_cost = 1000 + + # User has extra balance beyond the reservation + key = _make_key(balance=2000, reserved=990) + integration_session.add(key) + await integration_session.commit() + + response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}} + + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + + await _refresh(integration_session, key) + + assert key.balance >= 0, f"Balance went negative: {key.balance}" + assert key.reserved_balance == 0, f"Reservation not released: {key.reserved_balance}" + assert key.total_spent == actual_token_cost, ( + f"Expected full charge of {actual_token_cost}, got {key.total_spent}" + ) + assert key.balance == 2000 - actual_token_cost, ( + f"Expected balance={2000 - actual_token_cost}, got {key.balance}" + ) + + +# --------------------------------------------------------------------------- +# Test 4 — concurrent finalizations with cost overrun +# +# Multiple requests finish concurrently. Each has a small overrun. +# None should drive balance negative. +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_concurrent_cost_overruns_never_negative( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """Concurrent finalization with cost overruns must never produce negative balance.""" + import asyncio + + from routstr.auth import adjust_payment_for_tokens, pay_for_request + from routstr.core.db import create_session + + deducted_max_cost = 990 + actual_token_cost = 1000 + n_requests = 5 + + # Fund the key with exactly enough for n_requests reservations + a tiny buffer + starting_balance = deducted_max_cost * n_requests + key_hash = f"test_concurrent_{uuid.uuid4().hex}" + + async with create_session() as session: + key = ApiKey( + hashed_key=key_hash, + balance=starting_balance, + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + session.add(key) + await session.commit() + + # Reserve n_requests slots (sequentially, as pay_for_request is atomic) + async with create_session() as session: + key_to_reserve = await session.get(ApiKey, key_hash) + assert key_to_reserve is not None + for _ in range(n_requests): + await pay_for_request(key_to_reserve, deducted_max_cost, session) + await session.refresh(key_to_reserve) + + # Now finalize all concurrently with cost overrun + async def finalize() -> None: + response_data = { + "model": "test-model", + "usage": {"prompt_tokens": 100, "completion_tokens": 100}, + } + async with create_session() as session: + fresh_key = await session.get(ApiKey, key_hash) + assert fresh_key is not None + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens( + fresh_key, response_data, session, deducted_max_cost + ) + + await asyncio.gather(*[finalize() for _ in range(n_requests)]) + + async with create_session() as session: + final_key = await session.get(ApiKey, key_hash) + assert final_key is not None + + assert final_key.balance >= 0, ( + f"Balance went negative after concurrent overruns: {final_key.balance}" + ) + assert final_key.reserved_balance == 0, ( + f"Reserved balance not fully released: {final_key.reserved_balance}" + ) + assert final_key.total_spent <= starting_balance, ( + f"Total spent ({final_key.total_spent}) exceeds starting balance ({starting_balance})" + ) + # Every request must have been charged at least deducted_max_cost — no free inference. + assert final_key.total_spent == starting_balance, ( + f"Expected total_spent={starting_balance} (all {n_requests} reservations charged), " + f"got {final_key.total_spent} — at least one request got free inference" + ) + + +# --------------------------------------------------------------------------- +# Test 5 — overrun with no balance at all (reserved_balance == balance) +# simulates a user who topped up to exactly the reservation floor +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_zero_free_balance_overrun_is_safe( + integration_session: AsyncSession, +) -> None: + """User with zero free balance (all reserved) should never go negative on overrun.""" + from routstr.auth import adjust_payment_for_tokens + + deducted_max_cost = 1000 + actual_token_cost = 1050 + + # balance == reserved_balance: zero free balance + key = _make_key(balance=1000, reserved=1000) + integration_session.add(key) + await integration_session.commit() + + response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}} + + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + + await _refresh(integration_session, key) + + assert key.balance >= 0, f"Balance went negative: {key.balance}" + assert key.reserved_balance >= 0, f"Reserved balance went negative: {key.reserved_balance}" + + +# --------------------------------------------------------------------------- +# Test 6 — parallel requests: second finalization must not get free inference +# +# Root cause of the bug fixed in auth.py: +# `.where(col(ApiKey.balance) >= total_cost_msats)` ignores other requests' +# reservations, so after Request A charges total_cost_msats, balance can drop +# below deducted_max_cost, causing Request B's fallback to release for free. +# +# Fix: use `balance - reserved_balance + deducted_max_cost >= total_cost_msats` +# so the check accounts for concurrent reservations. +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_parallel_requests_no_free_inference( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """Second parallel finalization must be charged even when first depleted free balance.""" + import asyncio + + from routstr.auth import adjust_payment_for_tokens + from routstr.core.db import create_session + + deducted_max_cost = 100 + actual_token_cost = 150 # overrun: 50 more than reserved + + # Fund the key with exactly 2 * deducted_max_cost. + # Both requests pre-reserved 100 each → balance=200, reserved=200, free=0. + # Old check (balance >= total_cost_msats): + # Request A: 200 >= 150 ✓ → charges 150 → balance=50, reserved=100 + # Request B: 50 >= 150 ✗ → fallback: 50 >= 100 ✗ → releases FREE + # New check (balance - reserved + deducted >= total_cost_msats): + # Both fall to fallback (0 free balance). + # Both charge deducted_max_cost=100 → total_spent=200, balance=0. + starting_balance = deducted_max_cost * 2 + key_hash = f"test_parallel_no_free_{uuid.uuid4().hex}" + + async with create_session() as session: + key = ApiKey( + hashed_key=key_hash, + balance=starting_balance, + reserved_balance=deducted_max_cost * 2, # both slots pre-reserved + total_spent=0, + total_requests=2, + ) + session.add(key) + await session.commit() + + async def finalize() -> None: + response_data = { + "model": "test-model", + "usage": {"prompt_tokens": 50, "completion_tokens": 100}, + } + async with create_session() as session: + fresh_key = await session.get(ApiKey, key_hash) + assert fresh_key is not None + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(actual_token_cost), + ): + await adjust_payment_for_tokens( + fresh_key, response_data, session, deducted_max_cost + ) + + await asyncio.gather(finalize(), finalize()) + + async with create_session() as session: + final_key = await session.get(ApiKey, key_hash) + assert final_key is not None + + assert final_key.balance >= 0, f"Balance went negative: {final_key.balance}" + assert final_key.reserved_balance == 0, ( + f"Reserved balance not released: {final_key.reserved_balance}" + ) + # Both requests must have been charged — no free inference. + assert final_key.total_spent == starting_balance, ( + f"Expected total_spent={starting_balance} (both reservations charged), " + f"got {final_key.total_spent} — one request got free inference" + ) diff --git a/tests/integration/test_insufficient_balance.py b/tests/integration/test_insufficient_balance.py new file mode 100644 index 00000000..489b1dbf --- /dev/null +++ b/tests/integration/test_insufficient_balance.py @@ -0,0 +1,240 @@ +""" +Tests showing how a user hits "Insufficient balance: X mSats required for this model" +when their balance is too low for the model's cost. + +The log line that triggered this: + WARNING Insufficient billing balance during validation + ERROR Bearer token validation failed: HTTPException: 402: + {'error': {'message': 'Insufficient balance: 622888 mSats required + for this model. 20320 available.', ...}} + +This happens in validate_bearer_key (auth.py) when: + billing_key.total_balance < min_cost (model's max cost) + +and also in pay_for_request when the atomic UPDATE finds no available balance. +""" + +import uuid +from unittest.mock import patch + +import pytest +from fastapi import HTTPException +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey + + +def _key(balance: int, reserved: int = 0) -> ApiKey: + return ApiKey( + hashed_key=f"test_{uuid.uuid4().hex}", + balance=balance, + reserved_balance=reserved, + total_spent=0, + ) + + +# --------------------------------------------------------------------------- +# Test 1 — simplest case: balance < model cost → pay_for_request raises 402 +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_pay_for_request_raises_402_when_balance_too_low( + integration_session: AsyncSession, +) -> None: + """ + User has 20_000 msats. Model costs 622_888 msats. + pay_for_request must raise HTTP 402 with a clear message. + """ + from routstr.auth import pay_for_request + + model_cost = 622_888 + user_balance = 20_000 + + key = _key(balance=user_balance) + integration_session.add(key) + await integration_session.commit() + + with pytest.raises(HTTPException) as exc_info: + await pay_for_request(key, model_cost, integration_session) + + assert exc_info.value.status_code == 402 + detail = exc_info.value.detail + assert isinstance(detail, dict) + error = detail["error"] + assert error["code"] == "insufficient_balance" + assert str(model_cost) in error["message"] + assert str(user_balance) in error["message"] + + # Balance must be untouched + await integration_session.refresh(key) + assert key.balance == user_balance + assert key.reserved_balance == 0 + + +# --------------------------------------------------------------------------- +# Test 2 — balance is zero +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_pay_for_request_raises_402_on_zero_balance( + integration_session: AsyncSession, +) -> None: + """User with zero balance cannot make any request.""" + from routstr.auth import pay_for_request + + key = _key(balance=0) + integration_session.add(key) + await integration_session.commit() + + with pytest.raises(HTTPException) as exc_info: + await pay_for_request(key, 1_000, integration_session) + + assert exc_info.value.status_code == 402 + detail = exc_info.value.detail + assert isinstance(detail, dict) + assert detail["error"]["code"] == "insufficient_balance" + + +# --------------------------------------------------------------------------- +# Test 3 — all balance is reserved (total_balance = balance - reserved = 0) +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_pay_for_request_raises_402_when_all_balance_reserved( + integration_session: AsyncSession, +) -> None: + """ + User has 50_000 msats balance but 50_000 is already reserved for in-flight + requests. Free balance (total_balance) = 0. Should get 402. + """ + from routstr.auth import pay_for_request + + key = _key(balance=50_000, reserved=50_000) + integration_session.add(key) + await integration_session.commit() + + with pytest.raises(HTTPException) as exc_info: + await pay_for_request(key, 1_000, integration_session) + + assert exc_info.value.status_code == 402 + # Balance and reserved must be untouched + await integration_session.refresh(key) + assert key.balance == 50_000 + assert key.reserved_balance == 50_000 + + +# --------------------------------------------------------------------------- +# Test 4 — balance just one msat below model cost +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_pay_for_request_raises_402_one_msat_short( + integration_session: AsyncSession, +) -> None: + """Off-by-one: balance is exactly model_cost - 1.""" + from routstr.auth import pay_for_request + + model_cost = 10_000 + key = _key(balance=model_cost - 1) + integration_session.add(key) + await integration_session.commit() + + with pytest.raises(HTTPException) as exc_info: + await pay_for_request(key, model_cost, integration_session) + + assert exc_info.value.status_code == 402 + await integration_session.refresh(key) + assert key.balance == model_cost - 1 # untouched + assert key.reserved_balance == 0 + + +# --------------------------------------------------------------------------- +# Test 5 — balance exactly equal to model cost → succeeds +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_pay_for_request_succeeds_when_balance_equals_cost( + integration_session: AsyncSession, +) -> None: + """Balance == model cost: the request should be reserved successfully.""" + from routstr.auth import pay_for_request + + model_cost = 10_000 + key = _key(balance=model_cost) + integration_session.add(key) + await integration_session.commit() + + # Should not raise + await pay_for_request(key, model_cost, integration_session) + + await integration_session.refresh(key) + assert key.reserved_balance == model_cost + assert key.balance == model_cost # balance unchanged, only reserved goes up + + +# --------------------------------------------------------------------------- +# Test 6 — HTTP layer returns 402 JSON with the right shape +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_http_402_response_shape_on_insufficient_balance( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """ + End-to-end: POST /v1/chat/completions with a key whose balance is far below + the mocked model cost returns HTTP 402 with the expected JSON error body. + + Matches exactly the log snippet in the bug report: + 'Insufficient balance: X mSats required for this model. Y available.' + """ + from unittest.mock import AsyncMock, MagicMock + + model_cost = 622_888 + user_balance = 20_320 + + key = _key(balance=user_balance) + integration_session.add(key) + await integration_session.commit() + + # Minimal model stub so proxy routing doesn't 400 before reaching balance check + mock_model = MagicMock() + mock_model.sats_pricing = None + + # Upstream stub — never reached because balance check fires first + mock_upstream = MagicMock() + mock_upstream.prepare_headers = MagicMock(return_value={}) + + with ( + patch("routstr.proxy.get_model_instance", return_value=mock_model), + patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]), + # Patch where it is used (proxy imports it at module level) + patch( + "routstr.proxy.get_max_cost_for_model", + new=AsyncMock(return_value=model_cost), + ), + ): + response = await integration_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{key.hashed_key}"}, + json={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert response.status_code == 402 + body = response.json() + # FastAPI wraps HTTPException detail under "detail" + error = body["detail"]["error"] + assert error["code"] == "insufficient_balance" + assert error["type"] == "insufficient_quota" + assert str(model_cost) in error["message"] + assert str(user_balance) in error["message"] + + # Balance must be completely untouched + await integration_session.refresh(key) + assert key.balance == user_balance + assert key.reserved_balance == 0 + assert key.total_spent == 0 diff --git a/tests/integration/test_reservation_lifecycle.py b/tests/integration/test_reservation_lifecycle.py new file mode 100644 index 00000000..ee3fc875 --- /dev/null +++ b/tests/integration/test_reservation_lifecycle.py @@ -0,0 +1,268 @@ +""" +Tests for the reservation lifecycle: + + 1. Reserve → reserved_balance increases, available (total_balance) decreases. + 2. Reserve → revert → reserved_balance restored, balance untouched. + 3. Reserve → finalise → reserved_balance released, balance charged. + 4. Two parallel reserves, only one fits → second blocked with 402. + 5. Three parallel reserves, two fit, third blocked with 402. + 6. Sequential reserves until balance exhausted → next request blocked. + +Reservation invariant enforced by the atomic WHERE clause in pay_for_request: + balance - reserved_balance >= cost_per_request +""" + +import asyncio +import uuid + +import pytest +from fastapi import HTTPException +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.auth import pay_for_request, revert_pay_for_request +from routstr.core.db import ApiKey, create_session + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_key(balance: int, reserved: int = 0) -> ApiKey: + return ApiKey( + hashed_key=f"test_{uuid.uuid4().hex}", + balance=balance, + reserved_balance=reserved, + total_spent=0, + total_requests=0, + ) + + +async def _persist(session: AsyncSession, key: ApiKey) -> ApiKey: + session.add(key) + await session.commit() + await session.refresh(key) + return key + + +# --------------------------------------------------------------------------- +# Test 1 — Reserve: reserved_balance increases, available balance decreases +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_reserve_increases_reserved_balance( + integration_session: AsyncSession, +) -> None: + """pay_for_request must increment reserved_balance by cost_per_request.""" + cost = 100 + key = await _persist(integration_session, _make_key(balance=500)) + + await pay_for_request(key, cost, integration_session) + await integration_session.refresh(key) + + assert key.reserved_balance == cost + assert key.balance == 500 # balance column is NOT decremented on reserve + assert key.total_balance == 500 - cost # available = balance - reserved + + +# --------------------------------------------------------------------------- +# Test 2 — Revert: reserved_balance restored, balance untouched +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_revert_releases_reservation( + integration_session: AsyncSession, +) -> None: + """revert_pay_for_request must release the reservation without touching balance.""" + cost = 150 + key = await _persist(integration_session, _make_key(balance=300)) + + await pay_for_request(key, cost, integration_session) + await integration_session.refresh(key) + assert key.reserved_balance == cost + + await revert_pay_for_request(key, integration_session, cost) + await integration_session.refresh(key) + + assert key.reserved_balance == 0 + assert key.balance == 300 # balance unchanged after revert + + +# --------------------------------------------------------------------------- +# Test 3 — Finalise: reservation released + balance charged +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_finalise_releases_reservation_and_charges_balance( + integration_session: AsyncSession, +) -> None: + """adjust_payment_for_tokens must zero reserved_balance and deduct actual cost.""" + from unittest.mock import patch + + from routstr.auth import adjust_payment_for_tokens + from routstr.payment.cost_calculation import CostData + + cost = 100 + actual = 80 # actual < reserved → refund path + key = await _persist(integration_session, _make_key(balance=500)) + + await pay_for_request(key, cost, integration_session) + await integration_session.refresh(key) + assert key.reserved_balance == cost + + cost_data = CostData( + base_msats=0, + input_msats=40, + output_msats=40, + total_msats=actual, + total_usd=0.0, + input_tokens=50, + output_tokens=50, + ) + response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}} + + with patch("routstr.auth.calculate_cost", return_value=cost_data): + await adjust_payment_for_tokens(key, response_data, integration_session, cost) + + await integration_session.refresh(key) + + assert key.reserved_balance == 0 + assert key.balance == 500 - actual + assert key.total_spent == actual + + +# --------------------------------------------------------------------------- +# Test 4 — Concurrent: second parallel reserve blocked when balance exhausted +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_concurrent_second_reserve_blocked_when_balance_exhausted( + patched_db_engine: None, +) -> None: + """When two requests race for the same balance, only one succeeds; the other gets 402.""" + cost = 300 + key_hash = f"test_concurrent_{uuid.uuid4().hex}" + + async with create_session() as session: + key = ApiKey( + hashed_key=key_hash, + balance=300, # exactly enough for ONE reservation + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + session.add(key) + await session.commit() + + results: list[str] = [] + + async def attempt_reserve() -> None: + 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) + results.append("success") + except HTTPException as exc: + assert exc.status_code == 402 + results.append("blocked") + + await asyncio.gather(attempt_reserve(), attempt_reserve()) + + 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, key_hash) + assert final is not None + + # reserved_balance must equal exactly one reservation (not two) + assert final.reserved_balance == cost, ( + f"Expected reserved_balance={cost}, got {final.reserved_balance}" + ) + assert final.balance == 300, "Balance column must not be modified by reservation" + + +# --------------------------------------------------------------------------- +# Test 5 — Concurrent: three requests, two fit, third blocked +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_three_parallel_reserves_third_blocked( + patched_db_engine: None, +) -> None: + """Balance covers two reservations exactly; the third concurrent request must be blocked.""" + cost = 100 + key_hash = f"test_three_parallel_{uuid.uuid4().hex}" + + async with create_session() as session: + key = ApiKey( + hashed_key=key_hash, + balance=200, # fits exactly 2 reservations of 100 + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + session.add(key) + await session.commit() + + results: list[str] = [] + + async def attempt_reserve() -> None: + 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) + results.append("success") + except HTTPException as exc: + assert exc.status_code == 402 + results.append("blocked") + + await asyncio.gather( + attempt_reserve(), + attempt_reserve(), + attempt_reserve(), + ) + + successes = results.count("success") + blocked = results.count("blocked") + + assert successes == 2, f"Expected 2 successes, got {successes}: {results}" + assert blocked == 1, f"Expected 1 blocked, got {blocked}: {results}" + + async with create_session() as session: + final = await session.get(ApiKey, key_hash) + assert final is not None + + assert final.reserved_balance == cost * 2, ( + f"Expected reserved_balance={cost * 2}, got {final.reserved_balance}" + ) + assert final.balance == 200, "Balance column must not be modified by reservation" + + +# --------------------------------------------------------------------------- +# Test 6 — Sequential exhaustion: reserve until empty, next request blocked +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_sequential_reserves_block_when_balance_exhausted( + integration_session: AsyncSession, +) -> None: + """Repeated reservations should block as soon as available balance drops below cost.""" + cost = 100 + key = await _persist(integration_session, _make_key(balance=250)) + + # First two succeed (100 + 100 = 200 ≤ 250) + await pay_for_request(key, cost, integration_session) + await pay_for_request(key, cost, integration_session) + await integration_session.refresh(key) + assert key.reserved_balance == 200 + assert key.total_balance == 50 # 250 - 200 + + # Third: only 50 available, need 100 → blocked + with pytest.raises(HTTPException) as exc_info: + await pay_for_request(key, cost, integration_session) + + assert exc_info.value.status_code == 402 + await integration_session.refresh(key) + assert key.reserved_balance == 200 # unchanged after failed reserve diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py new file mode 100644 index 00000000..ef3d3f91 --- /dev/null +++ b/tests/unit/test_balance.py @@ -0,0 +1,250 @@ +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.responses import JSONResponse + +from routstr.balance import refund_wallet_endpoint +from routstr.core.db import ApiKey, CashuTransaction + + +def _make_cashu_tx( + token: str, + amount: int, + unit: str, + type: str = "out", + request_id: str | None = "req-abc", + swept: bool = False, + collected: bool = False, +) -> CashuTransaction: + tx = CashuTransaction(token=token, amount=amount, unit=unit, type=type, request_id=request_id) + tx.swept = swept + tx.collected = collected + return tx + + +def _exec_result(tx: CashuTransaction | None) -> MagicMock: + result = MagicMock() + result.first.return_value = tx + return result + + +@pytest.mark.asyncio +async def test_refund_x_cashu_returns_token() -> None: + x_cashu_token = "cashuAtest_token_value" + in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc") + out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc") + + session = MagicMock() + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) + session.add = MagicMock() + session.commit = AsyncMock() + + result = await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu=x_cashu_token, + session=session, + ) + + assert isinstance(result, JSONResponse) + body = json.loads(result.body) + assert body["token"] == "cashuArefund_token" + assert body["msats"] == "1000" + assert result.headers["X-Cashu"] == "cashuArefund_token" + assert out_tx.collected is True + + +@pytest.mark.asyncio +async def test_refund_x_cashu_sat_unit() -> None: + x_cashu_token = "cashuAsat_token" + in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat") + out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat") + + session = MagicMock() + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) + session.add = MagicMock() + session.commit = AsyncMock() + + result = await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu=x_cashu_token, + session=session, + ) + + assert isinstance(result, JSONResponse) + body = json.loads(result.body) + assert body["token"] == "cashuArefund_sat" + assert body["sats"] == "500" + assert "msats" not in body + assert result.headers["X-Cashu"] == "cashuArefund_sat" + + +@pytest.mark.asyncio +async def test_refund_x_cashu_not_found_raises_404() -> None: + from fastapi import HTTPException + + session = MagicMock() + session.exec = AsyncMock(return_value=_exec_result(None)) + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu="cashuAmissing_token", + session=session, + ) + + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_refund_x_cashu_swept_raises_410() -> None: + from fastapi import HTTPException + + in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept") + out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True) + + session = MagicMock() + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu="cashuAswept_token", + session=session, + ) + + assert exc_info.value.status_code == 410 + + +# --------------------------------------------------------------------------- +# source field defaults +# --------------------------------------------------------------------------- + + +def test_cashu_transaction_source_defaults_to_x_cashu() -> None: + tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat") + assert tx.source == "x-cashu" + + +def test_cashu_transaction_source_can_be_apikey() -> None: + tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat", source="apikey") + assert tx.source == "apikey" + + +# --------------------------------------------------------------------------- +# apikey-based refund: token logging and CashuTransaction storage +# --------------------------------------------------------------------------- + + +def _make_api_key( + balance: int = 5000, + 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 + key.reserved_balance = 0 + 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 + + +@pytest.mark.asyncio +async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None: + key = _make_api_key(balance=5000, refund_currency="sat") + refund_token = "cashuArefund_apikey_token" + + session = MagicMock() + session.add = MagicMock() + session.commit = AsyncMock() + + with ( + patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), + 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)), + patch("routstr.balance._refund_cache_set", AsyncMock()), + ): + result = await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + assert isinstance(result, dict) + assert result["token"] == refund_token + + mock_store.assert_awaited_once() + call_kwargs = mock_store.call_args.kwargs + assert call_kwargs["source"] == "apikey" + assert call_kwargs["token"] == refund_token + assert call_kwargs["typ"] == "out" + + +@pytest.mark.asyncio +async def test_apikey_refund_logs_token() -> None: + key = _make_api_key(balance=5000, refund_currency="sat") + refund_token = "cashuAlogged_token" + + session = MagicMock() + session.add = MagicMock() + session.commit = AsyncMock() + + with ( + patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), + 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)), + patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.balance.logger") as mock_logger, + ): + await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + calls = [str(c) for c in mock_logger.info.call_args_list] + assert any("cashu token issued" in c for c in calls) + + +@pytest.mark.asyncio +async def test_apikey_refund_log_includes_path() -> None: + key = _make_api_key(balance=5000, refund_currency="sat") + refund_token = "cashuApath_token" + + session = MagicMock() + session.add = MagicMock() + session.commit = AsyncMock() + + with ( + patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), + 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)), + patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.balance.logger") as mock_logger, + ): + await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + # Find the "cashu token issued" call and verify extra contains the path + token_issued_calls = [ + c for c in mock_logger.info.call_args_list + if c.args and "cashu token issued" in c.args[0] + ] + assert len(token_issued_calls) == 1 + extra = token_issued_calls[0].kwargs.get("extra", {}) + assert extra.get("path") == "/v1/wallet/refund" diff --git a/tests/unit/test_stream_id_injection.py b/tests/unit/test_stream_id_injection.py new file mode 100644 index 00000000..30c8323e --- /dev/null +++ b/tests/unit/test_stream_id_injection.py @@ -0,0 +1,109 @@ +import json +from collections.abc import AsyncGenerator +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from routstr.core.db import ApiKey +from routstr.upstream.base import BaseUpstreamProvider + + +@pytest.mark.asyncio +async def test_stream_with_id_injection() -> None: + """Test that stream_with_cost correctly injects IDs into complete JSON chunks but skips partials.""" + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test_key" + ) + + # Mock response with mixed chunks: + # 1. Complete JSON without ID + # 2. Partial JSON (should be passed through) + # 3. Complete JSON with ID (should be preserved or updated if requested_model is set) + # 4. [DONE] message + chunks = [ + b'data: {"choices": [{"delta": {"content": "Hello"}}]}\n\n', + b'data: {"choices": [{"delta": {"content": "', # Partial + b'world"}}]}\n\n', + b'data: {"id": "existing-id", "choices": [{"delta": {"content": "!"}}]}\n\n', + b"data: [DONE]\n\n", + ] + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + for chunk in chunks: + yield chunk + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "text/event-stream"} + mock_response.aiter_bytes = aiter_bytes + + key = MagicMock(spec=ApiKey) + key.hashed_key = "test_hash" + key.balance = 1000 + + background_tasks = MagicMock() + + # We need to mock adjust_payment_for_tokens since it's called at the end + with MagicMock(): + from routstr.upstream import base + + # Mocking the module-level function used in the generator + base.adjust_payment_for_tokens = AsyncMock( + return_value={"total_usd": 0.1, "total_msats": 100} + ) + base.create_session = MagicMock() + + streaming_response = await provider.handle_streaming_chat_completion( + response=mock_response, + key=key, + max_cost_for_model=100, + background_tasks=background_tasks, + requested_model="test-model", + ) + + results = [] + async for chunk in streaming_response.body_iterator: + results.append(chunk) + + # Parse results + parsed_results = [] + for r in results: + if isinstance(r, bytes) and r.startswith(b"data: "): + data = r[6:].decode().strip() + if data == "[DONE]": + parsed_results.append(data) + else: + try: + parsed_results.append(json.loads(data)) + except (json.JSONDecodeError, UnicodeDecodeError): + parsed_results.append( + data + ) # Keep as string if it failed to parse + + # Verifications + # 1. First chunk should have an injected ID and the requested model + assert isinstance(parsed_results[0], dict) + assert "id" in parsed_results[0] + assert parsed_results[0]["id"].startswith("chatcmpl-") + assert parsed_results[0]["model"] == "test-model" + + # 2. Second chunk was partial, should be passed as-is + # In current implementation, re.split(b"data: ", b'data: {...') gives ['', '{...'] + # The first empty part is skipped. The second part is processed. + + # Check that we have results + assert len(parsed_results) >= 4 + + # Find the chunk that was "existing-id" + id_chunk = next( + r + for r in parsed_results + if isinstance(r, dict) + and "choices" in r + and r["choices"][0]["delta"].get("content") == "!" + ) + assert id_chunk["id"] == parsed_results[0]["id"] + assert id_chunk["model"] == "test-model" + + # 4. [DONE] should be there + assert "[DONE]" in parsed_results diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 6635b24f..35a5694b 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -220,6 +220,33 @@ async def test_recieve_token_untrusted_mint() -> None: @pytest.mark.asyncio +@pytest.mark.asyncio +async def test_swap_to_primary_mint_already_on_primary() -> None: + from routstr.core.settings import settings + from routstr.wallet import swap_to_primary_mint + + mock_token = Mock() + mock_token.mint = settings.primary_mint + mock_token.amount = 1000 + mock_token.unit = "sat" + mock_token.proofs = [] + + mock_token_wallet = Mock() + mock_token_wallet.split = AsyncMock(return_value=None) + mock_token_wallet.request_mint = AsyncMock() + mock_token_wallet.melt_quote = AsyncMock() + + with patch("routstr.wallet.get_wallet", AsyncMock(return_value=mock_token_wallet)): + amount, unit, mint = await swap_to_primary_mint(mock_token, mock_token_wallet) + + assert amount == 1000 + assert unit == "sat" + assert mint == settings.primary_mint + mock_token_wallet.split.assert_called_once() + mock_token_wallet.request_mint.assert_not_called() + mock_token_wallet.melt_quote.assert_not_called() + + async def test_swap_to_primary_mint_success() -> None: """Test successful swap with dynamic fee calculation.""" from routstr.wallet import swap_to_primary_mint diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py new file mode 100644 index 00000000..d37d7acc --- /dev/null +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -0,0 +1,216 @@ +import json +import os +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.payment.cost_calculation import CostData # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + + +def _make_provider() -> BaseUpstreamProvider: + return BaseUpstreamProvider(base_url="http://test", api_key="test-key") + + +def _make_httpx_response(status_code: int = 200) -> httpx.Response: + return httpx.Response(status_code, headers={}) + + +def _make_cost_data(total_msats: int = 5000) -> CostData: + return CostData( + base_msats=0, + input_msats=3000, + output_msats=2000, + total_msats=total_msats, + total_usd=0.00025, + input_tokens=100, + output_tokens=50, + ) + + +# --------------------------------------------------------------------------- +# Non-streaming (chat completions) +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_non_streaming_includes_cost_sats() -> None: + provider = _make_provider() + cost_data = _make_cost_data(total_msats=5000) + + response_body = { + "model": "gpt-4o", + "usage": { + "prompt_tokens": 100, + "completion_tokens": 50, + "total_tokens": 150, + "cost": 0.00025, + }, + } + content_str = json.dumps(response_body) + httpx_response = _make_httpx_response() + + with ( + patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)), + patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")), + ): + response = await provider.handle_x_cashu_non_streaming_response( + content_str=content_str, + response=httpx_response, + amount=10000, + unit="msat", + max_cost_for_model=10000, + mint=None, + payment_token_hash=None, + ) + + body = json.loads(response.body) + assert "cost_sats" in body["usage"] + assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000 + + +@pytest.mark.asyncio +async def test_non_streaming_cost_sats_value_rounds_down() -> None: + provider = _make_provider() + cost_data = _make_cost_data(total_msats=1999) + + response_body = {"model": "gpt-4o", "usage": {"prompt_tokens": 10}} + content_str = json.dumps(response_body) + + with ( + patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)), + patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")), + ): + response = await provider.handle_x_cashu_non_streaming_response( + content_str=content_str, + response=_make_httpx_response(), + amount=10000, + unit="msat", + max_cost_for_model=10000, + ) + + body = json.loads(response.body) + assert body["usage"]["cost_sats"] == 1 # 1999 // 1000 + + +@pytest.mark.asyncio +async def test_non_streaming_preserves_existing_usage_fields() -> None: + provider = _make_provider() + cost_data = _make_cost_data(total_msats=3000) + + response_body = { + "model": "gpt-4o", + "usage": { + "prompt_tokens": 100, + "completion_tokens": 50, + "total_tokens": 150, + "cost": 0.00015, + }, + } + + with ( + patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)), + patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")), + ): + response = await provider.handle_x_cashu_non_streaming_response( + content_str=json.dumps(response_body), + response=_make_httpx_response(), + amount=10000, + unit="msat", + max_cost_for_model=10000, + ) + + body = json.loads(response.body) + usage = body["usage"] + assert usage["prompt_tokens"] == 100 + assert usage["completion_tokens"] == 50 + assert usage["total_tokens"] == 150 + assert usage["cost"] == 0.00015 + assert usage["cost_sats"] == 3 + + +# --------------------------------------------------------------------------- +# Streaming (chat completions) +# --------------------------------------------------------------------------- + +async def _collect_streaming(response: object) -> list[str]: + chunks: list[str] = [] + async for chunk in response.body_iterator: # type: ignore[attr-defined] + if isinstance(chunk, bytes): + chunks.append(chunk.decode("utf-8")) + else: + chunks.append(str(chunk)) + return chunks + + +@pytest.mark.asyncio +async def test_streaming_includes_cost_sats_in_usage_chunk() -> None: + provider = _make_provider() + cost_data = _make_cost_data(total_msats=7000) + + usage_chunk = { + "id": "chatcmpl-123", + "model": "gpt-4o", + "usage": {"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150}, + } + content_str = "\n".join([ + 'data: {"id":"chatcmpl-123","model":"gpt-4o","choices":[]}', + f"data: {json.dumps(usage_chunk)}", + "data: [DONE]", + ]) + + with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)): + response = await provider.handle_x_cashu_streaming_response( + content_str=content_str, + response=_make_httpx_response(), + amount=10000, + unit="msat", + max_cost_for_model=10000, + mint=None, + payment_token_hash=None, + ) + + chunks = await _collect_streaming(response) + + full_output = "".join(chunks) + usage_line = next( + line for line in full_output.split("\n") if '"usage"' in line and "cost_sats" in line + ) + data_json = json.loads(usage_line.lstrip("data: ").strip()) + assert data_json["usage"]["cost_sats"] == 7 # 7000 // 1000 + + +@pytest.mark.asyncio +async def test_streaming_non_usage_chunks_unmodified() -> None: + provider = _make_provider() + cost_data = _make_cost_data(total_msats=2000) + + regular_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "choices": [{"delta": {"content": "hi"}}]} + usage_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "usage": {"prompt_tokens": 10}} + content_str = "\n".join([ + f"data: {json.dumps(regular_chunk)}", + f"data: {json.dumps(usage_chunk)}", + "data: [DONE]", + ]) + + with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)): + response = await provider.handle_x_cashu_streaming_response( + content_str=content_str, + response=_make_httpx_response(), + amount=10000, + unit="msat", + max_cost_for_model=10000, + ) + + chunks = await _collect_streaming(response) + + lines = [ + line for line in "".join(chunks).split("\n") + if line.startswith("data: ") and line != "data: [DONE]" + ] + regular_line_data = json.loads(lines[0][6:]) + # regular chunk should not have cost_sats injected + assert "cost_sats" not in regular_line_data.get("usage", {}) diff --git a/ui/app/transactions/page.tsx b/ui/app/transactions/page.tsx index 07dc0e2a..19e0f032 100644 --- a/ui/app/transactions/page.tsx +++ b/ui/app/transactions/page.tsx @@ -22,6 +22,7 @@ import { SelectValue, } from '@/components/ui/select'; import { Badge } from '@/components/ui/badge'; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; import { Table, TableBody, @@ -47,6 +48,8 @@ import { Copy, Check, Receipt, + Key, + Zap, } from 'lucide-react'; import { AdminService, type Transaction } from '@/lib/api/services/admin'; import { format } from 'date-fns'; @@ -54,6 +57,118 @@ import { toast } from 'sonner'; const STORAGE_KEY = 'routstr-transaction-filters'; +function TransactionTable({ + transactions, + copiedId, + onCopy, + getStatusBadge, +}: { + transactions: Transaction[]; + copiedId: string | null; + onCopy: (text: string, id: string) => void; + getStatusBadge: (tx: Transaction) => React.ReactNode; +}) { + if (transactions.length === 0) { + return ( + + + + + + No transactions found + + Try adjusting your filters or check back later. + + + + ); + } + + return ( + + + + + Type + Amount + Status + Request ID + Mint + Date + Actions + + + + {transactions.map((tx) => ( + + +
+ {tx.type === 'in' ? ( + + ) : ( + + )} + {tx.type} +
+
+ + {tx.amount} {tx.unit} + + {getStatusBadge(tx)} + + {tx.request_id ? ( +
+ + {tx.request_id} + + +
+ ) : ( + — + )} +
+ +
+ {tx.mint_url} +
+
+ + {format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')} + + + + +
+ ))} +
+
+
+ ); +} + export default function TransactionsPage() { const [search, setSearch] = useState(''); const [type, setType] = useState('all'); @@ -145,12 +260,41 @@ export default function TransactionsPage() { .filter(Boolean) .join(' • '); + const xcashuTxs = + data?.transactions.filter((tx) => !tx.source || tx.source === 'x-cashu') ?? + []; + const apikeyTxs = + data?.transactions.filter((tx) => tx.source === 'apikey') ?? []; + + const renderCardContent = (txs: Transaction[]) => { + if (isLoading) { + return ( +
+ {Array.from({ length: 8 }).map((_, index) => ( + + ))} +
+ ); + } + return ( + + ); + }; + return (
refetch()} @@ -228,138 +372,64 @@ export default function TransactionsPage() { - - -
- Transaction History + + + + + X-Cashu {data && ( - - {data.transactions.length} entries + + {xcashuTxs.length} )} -
- {hasActiveFilters && ( - - Showing transactions filtered by {activeFilterDescription} - - )} -
- - {isLoading ? ( -
- {Array.from({ length: 8 }).map((_, index) => ( - - ))} -
- ) : data?.transactions && data.transactions.length > 0 ? ( - - - - - Type - Amount - Status - Request ID - Mint - Date - Actions - - - - {data.transactions.map((tx) => ( - - -
- {tx.type === 'in' ? ( - - ) : ( - - )} - {tx.type} -
-
- - {tx.amount} {tx.unit} - - {getStatusBadge(tx)} - - {tx.request_id ? ( -
- - {tx.request_id} - - -
- ) : ( - - — - - )} -
- -
- {tx.mint_url} -
-
- - {format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')} - - - - -
- ))} -
-
-
- ) : ( - - - - - - No transactions found - - Try adjusting your filters or check back later. - - - - )} -
-
+ + + + API Key Refunds + {data && ( + + {apikeyTxs.length} + + )} + + + + + + +
+ X-Cashu Transaction History + {hasActiveFilters && ( + + Filtered by {activeFilterDescription} + + )} +
+
+ + {renderCardContent(xcashuTxs)} + +
+
+ + + + +
+ API Key Refund History + {hasActiveFilters && ( + + Filtered by {activeFilterDescription} + + )} +
+
+ + {renderCardContent(apikeyTxs)} + +
+
+
); diff --git a/ui/components/add-provider-model-dialog.tsx b/ui/components/add-provider-model-dialog.tsx index 009eb942..a7fc8a6f 100644 --- a/ui/components/add-provider-model-dialog.tsx +++ b/ui/components/add-provider-model-dialog.tsx @@ -1,6 +1,6 @@ 'use client'; -import React, { useEffect, useMemo, useState } from 'react'; +import React, { useCallback, useEffect, useMemo, useState } from 'react'; import { useForm } from 'react-hook-form'; import { z } from 'zod'; import { zodResolver } from '@hookform/resolvers/zod'; @@ -39,7 +39,7 @@ import { FormMessage, } from '@/components/ui/form'; import { Switch } from '@/components/ui/switch'; -import { Loader2, Plus } from 'lucide-react'; +import { Check, Copy, Loader2, Plus } from 'lucide-react'; import { toast } from 'sonner'; import { AdminService, type AdminModel } from '@/lib/api/services/admin'; @@ -64,6 +64,7 @@ const FormSchema = z.object({ instruct_type: z.string().default(''), canonical_slug: z.string().default(''), alias_ids_raw: z.string().default(''), + forwarded_model_id: z.string().default(''), upstream_provider_id: z.string().default(''), input_cost: z.coerce.number().min(0).default(0), output_cost: z.coerce.number().min(0).default(0), @@ -104,6 +105,7 @@ export function AddProviderModelDialog({ const [isPresetOpen, setIsPresetOpen] = useState(false); const [selectedPresetLabel, setSelectedPresetLabel] = useState('Select a preset'); + const [forwardedModelIdCopied, setForwardedModelIdCopied] = useState(false); const form = useForm({ resolver: zodResolver(FormSchema) as never, @@ -119,6 +121,7 @@ export function AddProviderModelDialog({ instruct_type: '', canonical_slug: '', alias_ids_raw: '', + forwarded_model_id: '', upstream_provider_id: '', input_cost: 0, output_cost: 0, @@ -180,6 +183,7 @@ export function AddProviderModelDialog({ : '', canonical_slug: initialData.canonical_slug || '', alias_ids_raw: listToString(initialData.alias_ids), + forwarded_model_id: initialData.forwarded_model_id || initialData.id, upstream_provider_id: typeof initialData.upstream_provider_id === 'string' ? initialData.upstream_provider_id @@ -223,6 +227,7 @@ export function AddProviderModelDialog({ instruct_type: '', canonical_slug: '', alias_ids_raw: '', + forwarded_model_id: '', upstream_provider_id: '', input_cost: 0, output_cost: 0, @@ -280,6 +285,7 @@ export function AddProviderModelDialog({ ); form.setValue('canonical_slug', model.canonical_slug || ''); form.setValue('alias_ids_raw', listToString(model.alias_ids)); + form.setValue('forwarded_model_id', model.forwarded_model_id || model.id); form.setValue( 'upstream_provider_id', typeof model.upstream_provider_id === 'string' @@ -385,6 +391,7 @@ export function AddProviderModelDialog({ canonical_slug: data.canonical_slug?.trim() || null, alias_ids: listFromString(data.alias_ids_raw || ''), enabled: data.enabled, + forwarded_model_id: data.forwarded_model_id?.trim() || data.id, }; if (isEdit) { @@ -520,6 +527,53 @@ export function AddProviderModelDialog({ )} /> + { + const handleCopy = () => { + const value = field.value || form.getValues('id'); + if (!value) return; + navigator.clipboard.writeText(value); + setForwardedModelIdCopied(true); + setTimeout(() => setForwardedModelIdCopied(false), 1500); + }; + return ( + + Upstream Model ID + +
+ + +
+
+ + Model ID sent to the upstream provider. Defaults to the + model's own ID. + + +
+ ); + }} + /> + { + onApiKeyChange: (apiKey: string) => void; +} + +export function ApiKeyInput({ + value, + onApiKeyChange, + ...props +}: ApiKeyInputProps) { + const [internalValue, setInternalValue] = useState(value || ''); + + useEffect(() => { + setInternalValue(value || ''); + }, [value]); + + useEffect(() => { + const handler = setTimeout(() => { + onApiKeyChange(internalValue as string); + }, 300); + + return () => clearTimeout(handler); + }, [internalValue, onApiKeyChange]); + + return ( + setInternalValue(e.target.value)} + placeholder='sk-...' + className='font-mono text-sm' + {...props} + /> + ); +} diff --git a/ui/components/child-key-creator.tsx b/ui/components/child-key-creator.tsx index 1c67b442..2f462179 100644 --- a/ui/components/child-key-creator.tsx +++ b/ui/components/child-key-creator.tsx @@ -1,7 +1,9 @@ 'use client'; import { useState } from 'react'; +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, @@ -14,17 +16,8 @@ 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, - RotateCcw, - Plus, - Trash2, -} from 'lucide-react'; +import { Key, Copy, Check, Loader2, Plus, Trash2 } from 'lucide-react'; import { toast } from 'sonner'; -import { Badge } from '@/components/ui/badge'; import { KeyOptions } from './key-options'; interface KeyConfig { @@ -42,6 +35,14 @@ interface ChildKeyCreatorProps { 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, @@ -50,6 +51,7 @@ export function ChildKeyCreator({ }: ChildKeyCreatorProps) { const [internalApiKey, setInternalApiKey] = useState(''); const [loading, setLoading] = useState(false); + const [error, setError] = useState(null); const [configs, setConfigs] = useState([ { id: crypto.randomUUID(), @@ -59,15 +61,15 @@ export function ChildKeyCreator({ validityDate: '', }, ]); - const [childKeyToCheck, setChildKeyToCheck] = useState(''); - const [checking, setChecking] = useState(false); - const [keyStatus, setKeyStatus] = useState<{ - total_spent: number; - balance_limit: number | null; - validity_date: number | null; - is_expired: boolean; - is_drained: boolean; - } | null>(null); + + 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; @@ -75,13 +77,6 @@ export function ChildKeyCreator({ } | null>(null); const [copiedKey, setCopiedKey] = useState(null); - const activeApiKey = propApiKey ?? internalApiKey; - - const handleApiKeyChange = (val: string) => { - setInternalApiKey(val); - onApiKeyChange?.(val); - }; - const addConfig = () => { setConfigs([ ...configs, @@ -112,6 +107,7 @@ export function ChildKeyCreator({ } setLoading(true); + setError(null); try { let allNewKeys: string[] = []; let totalCost = 0; @@ -152,55 +148,21 @@ export function ChildKeyCreator({ ); } catch (error) { console.error('Failed to create child key:', error); - toast.error( - error instanceof Error ? error.message : 'Failed to create child key' - ); + 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 handleCheckKey = async () => { - if (!childKeyToCheck) { - toast.error('Please provide a Child API key to check'); - return; - } - - setChecking(true); - setKeyStatus(null); - try { - const baseUrlToUse = baseUrl || ''; - const response = await fetch(`${baseUrlToUse}/v1/balance/info`, { - headers: { - Authorization: `Bearer ${childKeyToCheck}`, - }, - }); - - if (!response.ok) { - throw new Error('Failed to fetch key info'); - } - - const info = await response.json(); - const now = Math.floor(Date.now() / 1000); - - setKeyStatus({ - total_spent: info.total_spent, - balance_limit: info.balance_limit, - validity_date: info.validity_date, - is_expired: info.validity_date ? now > info.validity_date : false, - is_drained: info.balance_limit - ? info.total_spent >= info.balance_limit - : false, - }); - } catch (error) { - toast.error( - error instanceof Error ? error.message : 'Failed to check child key' - ); - } finally { - setChecking(false); - } - }; - const copyToClipboard = (key: string) => { navigator.clipboard.writeText(key); setCopiedKey(key); @@ -243,12 +205,55 @@ export function ChildKeyCreator({ - handleApiKeyChange(e.target.value)} - placeholder='sk-...' - className='font-mono text-sm' - /> +
+
+ +
+ +
+ {walletInfo && ( +
+
+ + Spendable Balance + + + {formatSats(walletInfo.balanceMsats)} sats + +
+
+ + Total Requests + + + {walletInfo.totalRequests} + +
+
+ + Total Spent + +
+

+ {formatSats(walletInfo.totalSpent)} sats +

+

+ {formatMsats(walletInfo.totalSpent)} msats +

+
+
+
+ )} )} @@ -362,11 +367,18 @@ export function ChildKeyCreator({ )} - -

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

+ {error && ( + + Error + {error} + + )} + +

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

+ {newKeys.length > 0 && (
@@ -458,89 +470,6 @@ export function ChildKeyCreator({
- - - - Check Child Key Status - - View the current spending, limit, and expiration status of any child - key. - - - -
-
- - setChildKeyToCheck(e.target.value)} - placeholder='sk-...' - className='font-mono text-sm' - /> -
- - - {keyStatus && ( -
-
- Total Spent: - - {keyStatus.total_spent} mSats - -
- {keyStatus.balance_limit !== null && ( -
- Limit: - - {keyStatus.balance_limit} mSats - -
- )} - {keyStatus.validity_date !== null && ( -
- Expires: - - {new Date( - keyStatus.validity_date * 1000 - ).toLocaleDateString()} - -
- )} -
- {keyStatus.is_drained && ( - Drained - )} - {keyStatus.is_expired && ( - Expired - )} - {!keyStatus.is_drained && !keyStatus.is_expired && ( - Active - )} -
-
- )} -
-
-
); } diff --git a/ui/components/edit-model-form.tsx b/ui/components/edit-model-form.tsx deleted file mode 100644 index b565ef5f..00000000 --- a/ui/components/edit-model-form.tsx +++ /dev/null @@ -1,577 +0,0 @@ -'use client'; - -import React, { useState, useEffect, useCallback } from 'react'; -import { useForm } from 'react-hook-form'; -import { zodResolver } from '@hookform/resolvers/zod'; -import { z } from 'zod'; -import { type Model } from '@/lib/api/schemas/models'; -import { AdminService, type AdminModel } from '@/lib/api/services/admin'; -import { Button } from '@/components/ui/button'; -import { Input } from '@/components/ui/input'; -import { Textarea } from '@/components/ui/textarea'; -import { - Dialog, - DialogContent, - DialogDescription, - DialogHeader, - DialogTitle, -} from '@/components/ui/dialog'; -import { - Form, - FormControl, - FormDescription, - FormField, - FormItem, - FormLabel, - FormMessage, -} from '@/components/ui/form'; -import { Edit3, Loader2 } from 'lucide-react'; -import { toast } from 'sonner'; -import { Switch } from '@/components/ui/switch'; - -const EditModelFormSchema = z.object({ - name: z.string().min(1, 'Name is required'), - description: z.string().optional(), - context_length: z.number().min(0), - prompt: z.number().min(0), - completion: z.number().min(0), - enabled: z.boolean(), -}); - -type EditModelFormData = z.infer; - -const roundToFiveDecimals = (value: number | undefined | null): number => { - if (value === undefined || value === null || isNaN(value)) { - return 0; - } - return Math.round(value * 100000) / 100000; -}; - -const toNumber = (value: unknown, fallback = 0): number => { - if (typeof value === 'number' && Number.isFinite(value)) { - return value; - } - if (typeof value === 'string') { - const parsed = Number(value); - if (Number.isFinite(parsed)) { - return parsed; - } - } - return fallback; -}; - -const toStringArray = (value: unknown, fallback: string[]): string[] => { - if (!Array.isArray(value)) { - return fallback; - } - const filtered = value.filter( - (item): item is string => typeof item === 'string' - ); - return filtered.length > 0 ? filtered : fallback; -}; - -interface EditModelFormProps { - model: Model; - providerId?: number; - onModelUpdate?: () => void; - onCancel?: () => void; - isOpen: boolean; -} - -interface AdminModelData { - id: string; - name: string; - description?: string; - created: number; - context_length: number; - architecture: { - modality: string; - input_modalities: string[]; - output_modalities: string[]; - tokenizer: string; - instruct_type: string | null; - }; - pricing: { - prompt: number; - completion: number; - request: number; - image: number; - web_search: number; - internal_reasoning: number; - }; - per_request_limits: null | undefined; - top_provider: null | undefined; - upstream_provider_id: number; - enabled: boolean; -} - -const normalizeAdminModelData = ( - adminModel: AdminModel, - fallbackModel: Model, - providerId: number -): AdminModelData => { - const pricingRecord = - adminModel.pricing && typeof adminModel.pricing === 'object' - ? (adminModel.pricing as Record) - : {}; - - const architectureRecord = - adminModel.architecture && typeof adminModel.architecture === 'object' - ? (adminModel.architecture as Record) - : {}; - - return { - id: adminModel.id, - name: adminModel.name, - description: adminModel.description || '', - created: toNumber(adminModel.created, Math.floor(Date.now() / 1000)), - context_length: Math.max( - 0, - Math.trunc( - toNumber(adminModel.context_length, fallbackModel.contextLength || 4096) - ) - ), - architecture: { - modality: - typeof architectureRecord.modality === 'string' - ? architectureRecord.modality - : fallbackModel.modelType || 'text', - input_modalities: toStringArray(architectureRecord.input_modalities, [ - fallbackModel.modelType || 'text', - ]), - output_modalities: toStringArray(architectureRecord.output_modalities, [ - fallbackModel.modelType || 'text', - ]), - tokenizer: - typeof architectureRecord.tokenizer === 'string' - ? architectureRecord.tokenizer - : '', - instruct_type: - typeof architectureRecord.instruct_type === 'string' - ? architectureRecord.instruct_type - : null, - }, - pricing: { - prompt: roundToFiveDecimals( - toNumber(pricingRecord.prompt, fallbackModel.input_cost) - ), - completion: roundToFiveDecimals( - toNumber(pricingRecord.completion, fallbackModel.output_cost) - ), - request: toNumber(pricingRecord.request, 0), - image: toNumber(pricingRecord.image, 0), - web_search: toNumber(pricingRecord.web_search, 0), - internal_reasoning: toNumber(pricingRecord.internal_reasoning, 0), - }, - per_request_limits: - adminModel.per_request_limits === null || - adminModel.per_request_limits === undefined - ? adminModel.per_request_limits - : null, - top_provider: - adminModel.top_provider === null || adminModel.top_provider === undefined - ? adminModel.top_provider - : null, - upstream_provider_id: - typeof adminModel.upstream_provider_id === 'number' - ? adminModel.upstream_provider_id - : providerId, - enabled: adminModel.enabled !== false, - }; -}; - -export function EditModelForm({ - model, - providerId, - onModelUpdate, - onCancel, - isOpen, -}: EditModelFormProps) { - const [isSubmitting, setIsSubmitting] = useState(false); - const [adminModelData, setAdminModelData] = useState( - null - ); - const [isNewOverride, setIsNewOverride] = useState(false); - - const form = useForm({ - resolver: zodResolver(EditModelFormSchema), - defaultValues: { - name: model.name, - description: model.description || '', - context_length: model.contextLength || 4096, - prompt: roundToFiveDecimals(model.input_cost), - completion: roundToFiveDecimals(model.output_cost), - enabled: model.isEnabled !== false, - }, - }); - - const loadAdminModel = useCallback(async () => { - if (!providerId) { - console.error('loadAdminModel called without providerId'); - return; - } - - try { - const adminModel = await AdminService.getProviderModel( - providerId, - model.id - ); - - const normalizedAdminModel = normalizeAdminModelData( - adminModel, - model, - providerId - ); - setAdminModelData(normalizedAdminModel); - setIsNewOverride(false); - - form.reset({ - name: normalizedAdminModel.name, - description: normalizedAdminModel.description || '', - context_length: normalizedAdminModel.context_length, - prompt: normalizedAdminModel.pricing.prompt, - completion: normalizedAdminModel.pricing.completion, - enabled: normalizedAdminModel.enabled !== false, - }); - } catch { - setIsNewOverride(true); - setAdminModelData({ - id: model.full_name, - name: model.name, - description: model.description || '', - created: Math.floor(Date.now() / 1000), - context_length: model.contextLength || 4096, - architecture: { - modality: model.modelType || 'text', - input_modalities: [model.modelType || 'text'], - output_modalities: [model.modelType || 'text'], - tokenizer: '', - instruct_type: null, - }, - pricing: { - prompt: roundToFiveDecimals(model.input_cost), - completion: roundToFiveDecimals(model.output_cost), - request: 0, - image: 0, - web_search: 0, - internal_reasoning: 0, - }, - per_request_limits: null, - top_provider: null, - upstream_provider_id: providerId, - enabled: model.isEnabled !== false, - }); - - form.reset({ - name: model.name, - description: model.description || '', - context_length: model.contextLength || 4096, - prompt: roundToFiveDecimals(model.input_cost), - completion: roundToFiveDecimals(model.output_cost), - enabled: model.isEnabled !== false, - }); - } - }, [providerId, model, form]); - - useEffect(() => { - if (isOpen && providerId) { - loadAdminModel(); - } else if (isOpen && !providerId) { - console.error('EditModelForm opened without providerId', { - model, - providerId, - }); - toast.error('Missing provider information for this model'); - } - }, [isOpen, providerId, model, loadAdminModel]); - - const onSubmit = async (data: EditModelFormData) => { - if (!providerId) { - console.error('onSubmit called without providerId', { - model, - providerId, - }); - toast.error('Missing provider ID - cannot update model'); - return; - } - - if (!adminModelData) { - console.error('onSubmit called without adminModelData', { - model, - providerId, - adminModelData, - }); - toast.error('Model data not loaded - please try reopening the form'); - return; - } - - setIsSubmitting(true); - try { - const payload = { - id: adminModelData.id, - name: data.name, - description: data.description || '', - created: adminModelData.created || Math.floor(Date.now() / 1000), - context_length: data.context_length, - architecture: adminModelData.architecture || { - modality: 'text', - input_modalities: ['text'], - output_modalities: ['text'], - tokenizer: '', - instruct_type: null, - }, - pricing: { - prompt: roundToFiveDecimals(data.prompt), - completion: roundToFiveDecimals(data.completion), - request: 0, - image: 0, - web_search: 0, - internal_reasoning: 0, - }, - per_request_limits: adminModelData.per_request_limits, - top_provider: adminModelData.top_provider, - upstream_provider_id: providerId, - enabled: data.enabled, - }; - - if (isNewOverride) { - await AdminService.createProviderModel(providerId, payload); - toast.success('Model override created successfully!'); - } else { - await AdminService.updateProviderModel( - providerId, - adminModelData.id, - payload - ); - toast.success('Model updated successfully!'); - } - - onModelUpdate?.(); - onCancel?.(); - } catch (error) { - const action = isNewOverride ? 'create' : 'update'; - toast.error(`Failed to ${action} model. Please try again.`); - console.error(`Error ${action}ing model:`, error); - } finally { - setIsSubmitting(false); - } - }; - - const handleClose = () => { - if (!isSubmitting) { - onCancel?.(); - } - }; - - return ( - - - - - - {isNewOverride ? 'Create Model Override' : 'Edit Model Override'} - - - {isNewOverride - ? `Create an override for "${model.name}"` - : `Update the model override for "${model.name}"`} - - - -
- -
- ( - - Display Name * - - - - - Custom display name for the model - - - - )} - /> - - ( - - Context Length * - - { - const value = e.target.value; - field.onChange( - value === '' ? 0 : parseInt(value, 10) || 0 - ); - }} - onBlur={(e) => { - const value = parseInt(e.target.value, 10); - field.onChange( - Number.isNaN(value) ? 0 : Math.max(0, value) - ); - }} - className='w-full' - /> - - - Maximum context window size - - - - )} - /> -
- - ( - - Description - -