From fa313e4534562fe6db1851af082a06917db43f49 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 28 Nov 2025 13:40:30 +0900 Subject: [PATCH] the refactor part 01 --- .../versions/5f6ac1e4fa9f_the_refactor.py | 57 ++ routstr/auth.py | 705 ++-------------- routstr/balance.py | 53 +- routstr/core/admin.py | 40 +- routstr/core/db.py | 51 +- routstr/core/main.py | 276 +------ routstr/core/tasks.py | 134 +++ routstr/core/ui.py | 160 ++++ routstr/models/__init__.py | 5 + routstr/models/algorithm.py | 300 +++++++ routstr/models/crud.py | 180 +++++ routstr/models/metadata.py | 89 ++ routstr/{payment => models}/models.py | 261 +----- routstr/nostr/__init__.py | 0 routstr/{ => nostr}/discovery.py | 4 +- routstr/{nip91.py => nostr/listing.py} | 6 +- routstr/payment/__init__.py | 2 +- routstr/payment/cashu.py | 74 ++ routstr/payment/cost.py | 426 ++++++++++ routstr/payment/cost_calculation.py | 156 ---- routstr/payment/helpers.py | 760 +++++++++--------- routstr/payment/lnurl.py | 9 + routstr/{ => payment}/wallet.py | 55 +- routstr/proxy.py | 190 ++--- routstr/upstream/anthropic.py | 2 +- routstr/upstream/base.py | 28 +- routstr/upstream/helpers.py | 2 +- routstr/upstream/ollama.py | 12 +- routstr/upstream/openai.py | 2 +- routstr/upstream/openrouter.py | 2 +- routstr/upstream/perplexity.py | 2 +- routstr/upstream/xai.py | 2 +- 32 files changed, 2053 insertions(+), 1992 deletions(-) create mode 100644 migrations/versions/5f6ac1e4fa9f_the_refactor.py create mode 100644 routstr/core/tasks.py create mode 100644 routstr/core/ui.py create mode 100644 routstr/models/__init__.py create mode 100644 routstr/models/algorithm.py create mode 100644 routstr/models/crud.py create mode 100644 routstr/models/metadata.py rename routstr/{payment => models}/models.py (63%) create mode 100644 routstr/nostr/__init__.py rename routstr/{ => nostr}/discovery.py (99%) rename routstr/{nip91.py => nostr/listing.py} (99%) create mode 100644 routstr/payment/cashu.py create mode 100644 routstr/payment/cost.py delete mode 100644 routstr/payment/cost_calculation.py rename routstr/{ => payment}/wallet.py (87%) diff --git a/migrations/versions/5f6ac1e4fa9f_the_refactor.py b/migrations/versions/5f6ac1e4fa9f_the_refactor.py new file mode 100644 index 00000000..cf4fb8a2 --- /dev/null +++ b/migrations/versions/5f6ac1e4fa9f_the_refactor.py @@ -0,0 +1,57 @@ +"""the refactor + +Revision ID: 5f6ac1e4fa9f +Revises: a1a1a1a1a1a1 +Create Date: 2025-11-28 13:31:49.796461 +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "5f6ac1e4fa9f" +down_revision = "a1a1a1a1a1a1" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Rename the table + op.rename_table("api_keys", "temporary_credit") + + # Perform column modifications in batch mode for SQLite support + with op.batch_alter_table("temporary_credit", schema=None) as batch_op: + batch_op.add_column(sa.Column("created", sa.DateTime(), nullable=True)) + batch_op.add_column( + sa.Column("refund_expiration_time", sa.Integer(), nullable=True) + ) + batch_op.drop_column("key_expiry_time") + batch_op.drop_column("total_spent") + batch_op.drop_column("total_requests") + + +def downgrade() -> None: + # Revert column modifications + with op.batch_alter_table("temporary_credit", schema=None) as batch_op: + batch_op.add_column( + sa.Column( + "total_requests", + sa.Integer(), + nullable=False, + server_default=sa.text("0"), + ) + ) + batch_op.add_column( + sa.Column( + "total_spent", + sa.Integer(), + nullable=False, + server_default=sa.text("0"), + ) + ) + batch_op.add_column(sa.Column("key_expiry_time", sa.Integer(), nullable=True)) + batch_op.drop_column("refund_expiration_time") + batch_op.drop_column("created") + + # Revert table rename + op.rename_table("temporary_credit", "api_keys") diff --git a/routstr/auth.py b/routstr/auth.py index d5b8ed87..9a739b05 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -1,661 +1,90 @@ import hashlib -import math -from typing import Optional +from typing import Annotated -from fastapi import HTTPException +from fastapi import Depends, Header, HTTPException from sqlalchemy.exc import IntegrityError -from sqlmodel import col, update from .core import get_logger -from .core.db import ApiKey, AsyncSession +from .core.db import AsyncSession, TemporaryCredit, get_session from .core.settings import settings -from .payment.cost_calculation import ( - CostData, - CostDataError, - MaxCostData, - calculate_cost, -) -from .wallet import credit_balance, deserialize_token_from_string +from .payment.wallet import credit_balance, deserialize_token_from_string logger = get_logger(__name__) -# TODO: implement prepaid api key (not like it was before) -# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None) -# PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats + +async def api_key_to_credit(key: str, db_session: AsyncSession) -> TemporaryCredit: + key = key[3:] if key.startswith("sk-") else key + if existing_credit := await db_session.get(TemporaryCredit, key): + return existing_credit + + raise HTTPException(status_code=401, detail="Invalid API-KEY") -async def validate_bearer_key( - bearer_key: str, - session: AsyncSession, - refund_address: Optional[str] = None, - key_expiry_time: Optional[int] = None, -) -> ApiKey: - """ - Validates the provided API key using SQLModel. - If it's a cashu key, it redeems it and stores its hash and balance. - Otherwise checks if the hash of the key exists. - """ - logger.debug( - "Starting bearer key validation", - extra={ - "key_preview": bearer_key[:20] + "..." - if len(bearer_key) > 20 - else bearer_key, - "has_refund_address": bool(refund_address), - "has_expiry_time": bool(key_expiry_time), - }, +async def cashu_token_to_credit( + cashu_token: str, db_session: AsyncSession +) -> TemporaryCredit: + try: + token_hash = hashlib.sha256(cashu_token.encode()).hexdigest() + token_obj = deserialize_token_from_string(cashu_token) + except Exception: + raise HTTPException(status_code=400, detail="Invalid token format") + + if existing_credit := await db_session.get(TemporaryCredit, token_hash): + return existing_credit + + # create new TemporaryCredit + if token_obj.mint in settings.cashu_mints: + refund_currency = token_obj.unit + refund_mint_url = token_obj.mint + else: + refund_mint_url = settings.primary_mint + refund_currency = "sat" + + new_credit = TemporaryCredit( + hashed_key=token_hash, + refund_currency=refund_currency, + refund_mint_url=refund_mint_url, ) + db_session.add(new_credit) - if not bearer_key: - logger.error("Empty bearer key provided") - raise HTTPException( - status_code=401, - detail={ - "error": { - "message": "API key or Cashu token required", - "type": "invalid_request_error", - "code": "missing_api_key", - } - }, - ) + try: + await db_session.flush() + except IntegrityError: # fallback to api key in case of race condition + await db_session.rollback() + return await api_key_to_credit(f"sk-{token_hash}", db_session) - if bearer_key.startswith("sk-"): - logger.debug( - "Processing sk- prefixed API key", - extra={"key_preview": bearer_key[:10] + "..."}, - ) + msats = await credit_balance(cashu_token, new_credit, db_session) - if existing_key := await session.get(ApiKey, bearer_key[3:]): - logger.info( - "Existing sk- API key found", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "balance": existing_key.balance, - "total_requests": existing_key.total_requests, - }, - ) - - if key_expiry_time is not None: - existing_key.key_expiry_time = key_expiry_time - logger.debug( - "Updated key expiry time", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "expiry_time": key_expiry_time, - }, - ) - - if refund_address is not None: - existing_key.refund_address = refund_address - logger.debug( - "Updated refund address", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "refund_address_preview": refund_address[:20] + "..." - if len(refund_address) > 20 - else refund_address, - }, - ) - - return existing_key - else: - logger.warning( - "sk- API key not found in database", - extra={"key_preview": bearer_key[:10] + "..."}, - ) - - if bearer_key.startswith("cashu"): - logger.debug( - "Processing Cashu token", - extra={ - "token_preview": bearer_key[:20] + "...", - "token_type": bearer_key[:6] if len(bearer_key) >= 6 else bearer_key, - }, - ) - - try: - hashed_key = hashlib.sha256(bearer_key.encode()).hexdigest() - token_obj = deserialize_token_from_string(bearer_key) - logger.debug( - "Generated token hash", extra={"hash_preview": hashed_key[:16] + "..."} - ) - - if existing_key := await session.get(ApiKey, hashed_key): - logger.info( - "Existing Cashu token found", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "balance": existing_key.balance, - "total_requests": existing_key.total_requests, - }, - ) - - if key_expiry_time is not None: - existing_key.key_expiry_time = key_expiry_time - logger.debug( - "Updated key expiry time for existing Cashu key", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "expiry_time": key_expiry_time, - }, - ) - - if refund_address is not None: - existing_key.refund_address = refund_address - logger.debug( - "Updated refund address for existing Cashu key", - extra={ - "key_hash": existing_key.hashed_key[:8] + "...", - "refund_address_preview": refund_address[:20] + "..." - if len(refund_address) > 20 - else refund_address, - }, - ) - - return existing_key - - logger.info( - "Creating new Cashu token entry", - extra={ - "hash_preview": hashed_key[:16] + "...", - "has_refund_address": bool(refund_address), - "has_expiry_time": bool(key_expiry_time), - }, - ) - if token_obj.mint in settings.cashu_mints: - refund_currency = token_obj.unit - refund_mint_url = token_obj.mint - else: - refund_currency = "sat" - refund_mint_url = settings.primary_mint - - new_key = ApiKey( - hashed_key=hashed_key, - balance=0, - refund_address=refund_address, - key_expiry_time=key_expiry_time, - refund_currency=refund_currency, - refund_mint_url=refund_mint_url, - ) - session.add(new_key) - - try: - await session.flush() - except IntegrityError: - await session.rollback() - logger.info( - "Concurrent key creation detected, fetching existing key", - extra={"key_hash": hashed_key[:8] + "..."}, - ) - existing_key = await session.get(ApiKey, hashed_key) - if not existing_key: - raise Exception("Failed to fetch existing key after IntegrityError") - - if key_expiry_time is not None: - existing_key.key_expiry_time = key_expiry_time - if refund_address is not None: - existing_key.refund_address = refund_address - - return existing_key - - logger.debug( - "New key created, starting token redemption", - extra={"key_hash": hashed_key[:8] + "..."}, - ) - - logger.info( - "AUTH: About to call credit_balance", - extra={"token_preview": bearer_key[:50]}, - ) - try: - msats = await credit_balance(bearer_key, new_key, session) - logger.info( - "AUTH: credit_balance returned successfully", extra={"msats": msats} - ) - except Exception as credit_error: - logger.error( - "AUTH: credit_balance failed", - extra={ - "error": str(credit_error), - "error_type": type(credit_error).__name__, - }, - ) - raise credit_error - - if msats <= 0: - logger.error( - "Token redemption returned zero or negative amount", - extra={"msats": msats, "key_hash": hashed_key[:8] + "..."}, - ) - raise Exception("Token redemption failed") - - await session.refresh(new_key) - await session.commit() - - logger.info( - "New Cashu token successfully redeemed and stored", - extra={ - "key_hash": hashed_key[:8] + "...", - "redeemed_msats": msats, - "final_balance": new_key.balance, - }, - ) - - return new_key - except Exception as e: - logger.error( - "Cashu token redemption failed", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "token_preview": bearer_key[:20] + "..." - if len(bearer_key) > 20 - else bearer_key, - }, - ) - raise HTTPException( - status_code=401, - detail={ - "error": { - "message": f"Invalid or expired Cashu key: {str(e)}", - "type": "invalid_request_error", - "code": "invalid_api_key", - } - }, - ) - - logger.error( - "Invalid API key format", - extra={ - "key_preview": bearer_key[:10] + "..." - if len(bearer_key) > 10 - else bearer_key, - "key_length": len(bearer_key), - }, - ) - - raise HTTPException( - status_code=401, - detail={ - "error": { - "message": "Invalid API key", - "type": "invalid_request_error", - "code": "invalid_api_key", - } - }, - ) - - -async def pay_for_request( - key: ApiKey, cost_per_request: int, session: AsyncSession -) -> int: - """Process payment for a request.""" + await db_session.refresh(new_credit) + await db_session.commit() logger.info( - "Processing payment for request", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "current_balance": key.balance, - "required_cost": cost_per_request, - "sufficient_balance": key.balance >= cost_per_request, - }, + "New TemporaryCredit created from CashuToken", + extra={"key_hash": token_hash[:8] + "...", "amount_msats": msats}, ) - if key.total_balance < cost_per_request: - logger.warning( - "Insufficient balance for request", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "balance": key.balance, - "reserved_balance": key.reserved_balance, - "required": cost_per_request, - "shortfall": cost_per_request - key.total_balance, - }, - ) + return new_credit - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", - "type": "insufficient_quota", - "code": "insufficient_balance", - } - }, - ) - logger.debug( - "Charging base cost for request", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "cost": cost_per_request, - "balance_before": key.balance, - }, +async def nwc_to_credit( + connection_string: str, db_session: AsyncSession +) -> TemporaryCredit: + raise NotImplementedError + + +async def get_credit( + authorization: Annotated[str, Header(...)], + db_session: AsyncSession = Depends(get_session), +) -> TemporaryCredit: + authorization = ( + authorization[7:] if authorization.startswith("Bearer ") else authorization ) - # Charge the base cost for the request atomically to avoid race conditions - stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) >= cost_per_request) - .values( - reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, - total_requests=col(ApiKey.total_requests) + 1, - ) - ) - result = await session.exec(stmt) # type: ignore[call-overload] - await session.commit() - - if result.rowcount == 0: - logger.error( - "Concurrent request depleted balance", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "required_cost": cost_per_request, - "current_balance": key.balance, - }, - ) - - # Another concurrent request spent the balance first - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", - "type": "insufficient_quota", - "code": "insufficient_balance", - } - }, - ) - - await session.refresh(key) - - logger.info( - "Payment processed successfully", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "charged_amount": cost_per_request, - "new_balance": key.balance, - "total_spent": key.total_spent, - "total_requests": key.total_requests, - }, - ) - - return cost_per_request - - -async def revert_pay_for_request( - key: ApiKey, session: AsyncSession, cost_per_request: int -) -> None: - stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, - total_requests=col(ApiKey.total_requests) - 1, - ) - ) - - result = await session.exec(stmt) # type: ignore[call-overload] - await session.commit() - if result.rowcount == 0: - logger.error( - "Failed to revert payment - insufficient reserved balance", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "cost_to_revert": cost_per_request, - "current_reserved_balance": key.reserved_balance, - }, - ) - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.", - "type": "payment_error", - "code": "payment_error", - } - }, - ) - await session.refresh(key) - - -async def adjust_payment_for_tokens( - key: ApiKey, response_data: dict, session: AsyncSession, deducted_max_cost: int -) -> dict: - """ - Adjusts the payment based on token usage in the response. - This is called after the initial payment and the upstream request is complete. - Returns cost data to be included in the response. - """ - model = response_data.get("model", "unknown") - - logger.debug( - "Starting payment adjustment for tokens", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": model, - "deducted_max_cost": deducted_max_cost, - "current_balance": key.balance, - "has_usage": "usage" in response_data, - }, - ) - - match await calculate_cost(response_data, deducted_max_cost, session): - case MaxCostData() as cost: - logger.debug( - "Using max cost data (no token adjustment)", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": model, - "max_cost": cost.total_msats, - }, - ) - # Finalize by releasing reservation and charging max cost - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, - balance=col(ApiKey.balance) - cost.total_msats, - total_spent=col(ApiKey.total_spent) + cost.total_msats, - ) - ) - result = await session.exec(finalize_stmt) # type: ignore[call-overload] - await session.commit() - if result.rowcount == 0: - logger.error( - "Failed to finalize max-cost payment - insufficient reserved balance", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, - "total_cost": cost.total_msats, - "model": model, - }, - ) - else: - await session.refresh(key) - logger.info( - "Max cost payment finalized", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "charged_amount": cost.total_msats, - "new_balance": key.balance, - "model": model, - }, - ) - return cost.dict() - - case CostData() as cost: - # If token-based pricing is enabled and base cost is 0, use token-based cost - # Otherwise, token cost is additional to the base cost - cost_difference = cost.total_msats - deducted_max_cost - total_cost_msats: int = math.ceil(cost.total_msats) - - logger.info( - "Calculated token-based cost", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": model, - "token_cost": cost.total_msats, - "deducted_max_cost": deducted_max_cost, - "cost_difference": cost_difference, - "input_msats": cost.input_msats, - "output_msats": cost.output_msats, - }, - ) - - if cost_difference == 0: - logger.debug( - "Finalizing with exact reserved cost", - extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, - ) - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .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, - ) - ) - await session.exec(finalize_stmt) # type: ignore[call-overload] - await session.commit() - await session.refresh(key) - 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", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "additional_charge": cost_difference, - "current_balance": key.balance, - "sufficient_balance": key.balance >= cost_difference, - "model": model, - }, - ) - - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .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, - ) - ) - result = await session.exec(finalize_stmt) # type: ignore[call-overload] - await session.commit() - - if result.rowcount: - cost.total_msats = total_cost_msats - await session.refresh(key) - - logger.info( - "Finalized payment with additional charge", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "charged_amount": total_cost_msats, - "new_balance": key.balance, - "model": model, - }, - ) - else: - logger.warning( - "Failed to finalize additional charge (concurrent operation)", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "attempted_charge": total_cost_msats, - "model": model, - }, - ) - else: - # Refund some of the base cost - refund = abs(cost_difference) - logger.info( - "Refunding excess payment", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "refund_amount": refund, - "current_balance": key.balance, - "model": model, - }, - ) - - refund_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .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, - ) - ) - result = await session.exec(refund_stmt) # type: ignore[call-overload] - await session.commit() - - if result.rowcount == 0: - logger.error( - "Failed to finalize payment - insufficient reserved balance", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, - "total_cost": total_cost_msats, - "model": model, - }, - ) - # Still return the cost data even if we couldn't properly finalize - # The reservation was already made, so the user has paid - - cost.total_msats = total_cost_msats - await session.refresh(key) - - logger.info( - "Refund processed successfully", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "refunded_amount": refund, - "new_balance": key.balance, - "final_cost": cost.total_msats, - "model": model, - }, - ) - - return cost.dict() - - case CostDataError() as error: - logger.error( - "Cost calculation error during payment adjustment", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": model, - "error_message": error.message, - "error_code": error.code, - }, - ) - - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": error.message, - "type": "invalid_request_error", - "code": error.code, - } - }, - ) - # Fallback return to satisfy type checker; execution should not reach here - return { - "base_msats": deducted_max_cost, - "input_msats": 0, - "output_msats": 0, - "total_msats": deducted_max_cost, - } + if authorization.startswith("cashu"): + return await cashu_token_to_credit(authorization, db_session) + elif authorization.startswith("sk-"): + return await api_key_to_credit(authorization, db_session) + elif authorization.startswith("nwc"): + return await nwc_to_credit(authorization, db_session) + else: + raise HTTPException(status_code=401, detail="Unable to parse bearer token") diff --git a/routstr/balance.py b/routstr/balance.py index 883c4673..e2e7310a 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -6,11 +6,12 @@ from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel -from .auth import validate_bearer_key -from .core.db import ApiKey, AsyncSession, get_session +from .auth import get_credit +from .core.db import AsyncSession, TemporaryCredit, get_session from .core.logging import get_logger from .core.settings import settings -from .wallet import credit_balance, recieve_token, send_to_lnurl, send_token +from .payment.lnurl import send_to_lnurl +from .payment.wallet import credit_balance, recieve_token, send_token router = APIRouter() balance_router = APIRouter(prefix="/v1/balance") @@ -18,26 +19,15 @@ balance_router = APIRouter(prefix="/v1/balance") logger = get_logger(__name__) -async def get_key_from_header( - authorization: Annotated[str, Header(...)], - session: AsyncSession = Depends(get_session), -) -> ApiKey: - if authorization.startswith("Bearer "): - return await validate_bearer_key(authorization[7:], session) - - raise HTTPException( - status_code=401, - detail="Invalid authorization. Use 'Bearer ' or 'Bearer '", - ) - - # TODO: remove this endpoint when frontend is updated @router.get("/", include_in_schema=False) -async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: +async def account_info( + credit: TemporaryCredit = Depends(get_credit), +) -> dict: return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, + "api_key": credit.api_key, + "balance": credit.balance, + "reserved": credit.reserved_balance, } @@ -57,19 +47,18 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: async def create_balance( initial_balance_token: str, session: AsyncSession = Depends(get_session) ) -> dict: - key = await validate_bearer_key(initial_balance_token, session) - return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - } + credit = await get_credit(initial_balance_token, session) + return {"api_key": credit.api_key, "balance": credit.balance} @router.get("/info") -async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: +async def wallet_info( + credit: TemporaryCredit = Depends(get_credit), +) -> dict: return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, + "api_key": credit.api_key, + "balance": credit.balance, + "reserved": credit.reserved_balance, } @@ -81,7 +70,7 @@ class TopupRequest(BaseModel): async def topup_wallet_endpoint( cashu_token: str | None = None, topup_request: TopupRequest | None = None, - key: ApiKey = Depends(get_key_from_header), + credit: TemporaryCredit = Depends(get_credit), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: if topup_request is not None: @@ -93,7 +82,7 @@ async def topup_wallet_endpoint( if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") try: - amount_msats = await credit_balance(cashu_token, key, session) + amount_msats = await credit_balance(cashu_token, credit, session) except ValueError as e: error_msg = str(e) if "already spent" in error_msg.lower(): @@ -152,7 +141,7 @@ async def refund_wallet_endpoint( if cached := await _refund_cache_get(bearer_value): return cached - key: ApiKey = await validate_bearer_key(bearer_value, session) + key: TemporaryCredit = await get_credit(bearer_value, session) remaining_balance_msats: int = key.balance if key.refund_currency == "sat": diff --git a/routstr/core/admin.py b/routstr/core/admin.py index c9357b60..1ccdedcc 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -8,16 +8,16 @@ from fastapi.responses import HTMLResponse, RedirectResponse from pydantic import BaseModel from sqlmodel import select -from ..payment.models import _row_to_model, list_models -from ..proxy import refresh_model_maps, reinitialize_upstreams -from ..wallet import ( +from ..models.crud import _row_to_model, list_models +from ..payment.wallet import ( fetch_all_balances, get_proofs_per_mint_and_unit, get_wallet, send_token, slow_filter_spend_proofs, ) -from .db import ApiKey, ModelRow, UpstreamProviderRow, create_session +from ..proxy import refresh_model_maps, reinitialize_upstreams +from .db import ModelRow, TemporaryCredit, UpstreamProviderRow, create_session from .log_manager import log_manager from .logging import get_logger from .settings import SettingsService, settings @@ -107,13 +107,13 @@ async def partial_balances(request: Request) -> str: @admin_router.get( - "/partials/apikeys", + "/partials/TemporaryCredits", dependencies=[Depends(require_admin_api)], response_class=HTMLResponse, ) -async def partial_apikeys(request: Request) -> str: +async def partial_TemporaryCredits(request: Request) -> str: async with create_session() as session: - result = await session.exec(select(ApiKey)) + result = await session.exec(select(TemporaryCredit)) api_keys = result.all() def fmt_time(ts: int | None) -> str: @@ -124,7 +124,7 @@ async def partial_apikeys(request: Request) -> str: rows = "".join( [ - f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" + f"{key.hashed_key}{key.balance}{key.refund_address}{key.refund_mint_url}{key.refund_currency}{fmt_time(key.refund_expiration_time)}" for key in api_keys ] ) @@ -147,17 +147,17 @@ async def partial_apikeys(request: Request) -> str: @admin_router.get("/api/temporary-balances", dependencies=[Depends(require_admin_api)]) async def get_temporary_balances_api(request: Request) -> list[dict[str, object]]: async with create_session() as session: - result = await session.exec(select(ApiKey)) + result = await session.exec(select(TemporaryCredit)) api_keys = result.all() return [ { "hashed_key": key.hashed_key, "balance": key.balance, - "total_spent": key.total_spent, - "total_requests": key.total_requests, "refund_address": key.refund_address, - "key_expiry_time": key.key_expiry_time, + "refund_mint_url": key.refund_mint_url, + "refund_currency": key.refund_currency, + "refund_expiration_time": key.refund_expiration_time, } for key in api_keys ] @@ -777,8 +777,8 @@ async def dashboard(request: Request) -> str:

Save this token! It represents your withdrawn balance.

-

Temporary Balances

@@ -1763,9 +1763,9 @@ UPSTREAM_PROVIDERS_JS: str = """ provider_fee: feeValue ? parseFloat(feeValue) : defaultFee, }; - const apiKey = document.getElementById('provider-api-key').value; - if (apiKey) { - payload.api_key = apiKey; + const TemporaryCredit = document.getElementById('provider-api-key').value; + if (TemporaryCredit) { + payload.api_key = TemporaryCredit; } const saveBtn = document.getElementById('provider-save-btn'); @@ -1783,7 +1783,7 @@ UPSTREAM_PROVIDERS_JS: str = """ body: JSON.stringify(payload) }); } else { - if (!apiKey) { + if (!TemporaryCredit) { throw new Error('API Key is required for new providers'); } resp = await fetch('/admin/api/upstream-providers', { @@ -2525,7 +2525,7 @@ async def update_provider_model( await session.refresh(row) if was_disabled and payload.enabled: - from ..payment.models import _cleanup_enabled_models_once + from ..models.models import _cleanup_enabled_models_once try: await _cleanup_enabled_models_once() @@ -2796,7 +2796,7 @@ async def get_provider_models(provider_id: int) -> dict[str, object]: dependencies=[Depends(require_admin_api)], ) async def get_openrouter_presets() -> list[dict[str, object]]: - from ..payment.models import async_fetch_openrouter_models + from ..models import async_fetch_openrouter_models models_data = await async_fetch_openrouter_models() return models_data diff --git a/routstr/core/db.py b/routstr/core/db.py index c6effbe3..666ae43a 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -1,5 +1,6 @@ import os from contextlib import asynccontextmanager +from datetime import datetime from typing import AsyncGenerator from alembic import command @@ -18,39 +19,40 @@ DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL -class ApiKey(SQLModel, table=True): # type: ignore - __tablename__ = "api_keys" +class TemporaryCredit(SQLModel, table=True): # type: ignore + __tablename__ = "temporary_credit" hashed_key: str = Field(primary_key=True) - balance: int = Field(default=0, description="Balance in millisatoshis (msats)") - reserved_balance: int = Field( - default=0, description="Reserved balance in millisatoshis (msats)" - ) + balance: int = Field(default=0, description="Balance in msats") + reserved_balance: int = Field(default=0, description="Blocked balance in msats") + created: datetime | None = Field(None, description="Timestamp of creation") refund_address: str | None = Field( - default=None, - description="Lightning address to refund remaining balance after key expires", + None, description="Address to refund on expiration" ) - key_expiry_time: int | None = Field( - default=None, - description="Unix-timestamp after which the cashu-token's balance gets refunded to the refund_address", - ) - total_spent: int = Field( - default=0, description="Total spent in millisatoshis (msats)" - ) - total_requests: int = Field(default=0) refund_mint_url: str | None = Field( - default=None, - description="URL of the mint used to create the cashu-token", + default=None, description="Mint used to issue the refund token" ) - refund_currency: str | None = Field( - default=None, - description="Currency of the cashu-token", + refund_currency: str | None = Field(None, description="Currency of the cashu-token") + refund_expiration_time: int | None = Field( + None, description="Refund not allowed after timeout" ) @property - def total_balance(self) -> int: + def total_balance_msat(self) -> int: return self.balance - self.reserved_balance + @property + def total_balance_sat(self) -> int: + return self.total_balance_msat // 1000 + + @property + def total_balance(self) -> int: + return self.total_balance_msat + + @property + def api_key(self) -> str: + return "sk-" + self.hashed_key + class ModelRow(SQLModel, table=True): # type: ignore __tablename__ = "models" @@ -95,8 +97,9 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore async def balances_for_mint_and_unit( db_session: AsyncSession, mint_url: str, unit: str ) -> int: - query = select(func.sum(ApiKey.balance)).where( - ApiKey.refund_mint_url == mint_url, ApiKey.refund_currency == unit + query = select(func.sum(TemporaryCredit.balance)).where( + TemporaryCredit.refund_mint_url == mint_url, + TemporaryCredit.refund_currency == unit, ) result = await db_session.exec(query) return result.one() or 0 diff --git a/routstr/core/main.py b/routstr/core/main.py index c3d4e104..81c8f429 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -1,33 +1,21 @@ -import asyncio import os -from contextlib import asynccontextmanager -from pathlib import Path -from typing import AsyncGenerator from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -from fastapi.responses import FileResponse, RedirectResponse -from fastapi.staticfiles import StaticFiles +from fastapi.responses import RedirectResponse from starlette.exceptions import HTTPException from ..balance import balance_router, deprecated_wallet_router -from ..discovery import providers_cache_refresher, providers_router -from ..nip91 import announce_provider -from ..payment.models import ( - cleanup_enabled_models_periodically, - models_router, - update_sats_pricing, -) -from ..payment.price import update_prices_periodically -from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically -from ..wallet import periodic_payout +from ..models.models import models_router +from ..nostr.discovery import providers_router +from ..proxy import proxy_router from .admin import admin_router -from .db import create_session, init_db, run_migrations from .exceptions import general_exception_handler, http_exception_handler from .logging import get_logger, setup_logging from .middleware import LoggingMiddleware -from .settings import SettingsService from .settings import settings as global_settings +from .tasks import lifespan +from .ui import setup_ui_routes # Initialize logging first setup_logging() @@ -39,119 +27,6 @@ else: __version__ = "0.2.1" -@asynccontextmanager -async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: - logger.info("Application startup initiated", extra={"version": __version__}) - - btc_price_task = None - pricing_task = None - payout_task = None - nip91_task = None - providers_task = None - models_refresh_task = None - models_cleanup_task = None - model_maps_refresh_task = None - - try: - # Run database migrations on startup - # This ensures the database schema is always up-to-date in production - # Migrations are idempotent - running them multiple times is safe - logger.info("Running database migrations") - run_migrations() - - # Initialize database connection pools - # This creates any tables that might not be tracked by migrations yet - await init_db() - - # Initialize application settings (env -> computed -> DB precedence) - async with create_session() as session: - s = await SettingsService.initialize(session) - - # Apply app metadata from settings - try: - app.title = s.name - app.description = s.description - except Exception: - pass - - # await ensure_models_bootstrapped() - - from ..payment.price import _update_prices - from ..proxy import get_upstreams - from ..upstream.helpers import refresh_upstreams_models_periodically - - await _update_prices() - await initialize_upstreams() - - btc_price_task = asyncio.create_task(update_prices_periodically()) - pricing_task = asyncio.create_task(update_sats_pricing()) - if global_settings.models_refresh_interval_seconds > 0: - models_refresh_task = asyncio.create_task( - refresh_upstreams_models_periodically(get_upstreams()) - ) - models_cleanup_task = asyncio.create_task(cleanup_enabled_models_periodically()) - model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) - payout_task = asyncio.create_task(periodic_payout()) - nip91_task = asyncio.create_task(announce_provider()) - providers_task = asyncio.create_task(providers_cache_refresher()) - - yield - - except Exception as e: - logger.error( - "Application startup failed", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - raise - finally: - logger.info("Application shutdown initiated") - - if btc_price_task is not None: - btc_price_task.cancel() - if pricing_task is not None: - pricing_task.cancel() - if payout_task is not None: - payout_task.cancel() - if nip91_task is not None: - nip91_task.cancel() - if providers_task is not None: - providers_task.cancel() - if models_refresh_task is not None: - models_refresh_task.cancel() - if models_cleanup_task is not None: - models_cleanup_task.cancel() - if model_maps_refresh_task is not None: - model_maps_refresh_task.cancel() - - try: - tasks_to_wait = [] - if btc_price_task is not None: - tasks_to_wait.append(btc_price_task) - if pricing_task is not None: - tasks_to_wait.append(pricing_task) - if payout_task is not None: - tasks_to_wait.append(payout_task) - if nip91_task is not None: - tasks_to_wait.append(nip91_task) - if providers_task is not None: - tasks_to_wait.append(providers_task) - if models_refresh_task is not None: - tasks_to_wait.append(models_refresh_task) - if models_cleanup_task is not None: - tasks_to_wait.append(models_cleanup_task) - if model_maps_refresh_task is not None: - tasks_to_wait.append(model_maps_refresh_task) - - if tasks_to_wait: - await asyncio.gather(*tasks_to_wait, return_exceptions=True) - logger.info("Background tasks stopped successfully") - except Exception as e: - logger.error( - "Error stopping background tasks", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - - app = FastAPI(version=__version__, lifespan=lifespan) @@ -190,144 +65,7 @@ async def providers() -> RedirectResponse: return RedirectResponse("/v1/providers/") -UI_DIST_PATH = Path(__file__).parent.parent.parent / "ui_out" - -if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): - logger.info(f"Serving static UI from {UI_DIST_PATH}") - - app.mount( - "/_next", - StaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True), - name="next-static", - ) - - @app.get("/", include_in_schema=False) - async def serve_root_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "index.html") - - # Add explicit route for /index.txt to redirect to / - @app.get("/index.txt", include_in_schema=False) - async def redirect_index_txt() -> RedirectResponse: - return RedirectResponse("/") - - @app.get("/admin") - async def admin_redirect() -> FileResponse: - return FileResponse(UI_DIST_PATH / "index.html") - - @app.get("/dashboard", include_in_schema=False) - async def serve_dashboard_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "index.html") - - @app.get("/login", include_in_schema=False) - async def serve_login_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "login" / "index.html") - - # Add explicit route for /login/index.txt to redirect to /login - @app.get("/login/index.txt", include_in_schema=False) - async def redirect_login_index_txt() -> RedirectResponse: - return RedirectResponse("/login") - - @app.get("/model", include_in_schema=False) - async def serve_models_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "model" / "index.html") - - # Add explicit route for /model/index.txt to redirect to /model - @app.get("/model/index.txt", include_in_schema=False) - async def redirect_model_index_txt() -> RedirectResponse: - return RedirectResponse("/model") - - @app.get("/providers", include_in_schema=False) - async def serve_providers_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "providers" / "index.html") - - # Add explicit route for /providers/index.txt to redirect to /providers - @app.get("/providers/index.txt", include_in_schema=False) - async def redirect_providers_index_txt() -> RedirectResponse: - return RedirectResponse("/providers") - - @app.get("/settings", include_in_schema=False) - async def serve_settings_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "settings" / "index.html") - - # Add explicit route for /settings/index.txt to redirect to /settings - @app.get("/settings/index.txt", include_in_schema=False) - async def redirect_settings_index_txt() -> RedirectResponse: - return RedirectResponse("/settings") - - @app.get("/transactions", include_in_schema=False) - async def serve_transactions_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "transactions" / "index.html") - - # Add explicit route for /transactions/index.txt to redirect to /transactions - @app.get("/transactions/index.txt", include_in_schema=False) - async def redirect_transactions_index_txt() -> RedirectResponse: - return RedirectResponse("/transactions") - - @app.get("/balances", include_in_schema=False) - async def serve_balances_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "balances" / "index.html") - - # Add explicit route for /balances/index.txt to redirect to /balances - @app.get("/balances/index.txt", include_in_schema=False) - async def redirect_balances_index_txt() -> RedirectResponse: - return RedirectResponse("/balances") - - @app.get("/logs", include_in_schema=False) - async def serve_logs_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "logs" / "index.html") - - # Add explicit route for /logs/index.txt to redirect to /logs - @app.get("/logs/index.txt", include_in_schema=False) - async def redirect_logs_index_txt() -> RedirectResponse: - return RedirectResponse("/logs") - - @app.get("/usage", include_in_schema=False) - async def serve_usage_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "usage" / "index.html") - - # Add explicit route for /usage/index.txt to redirect to /usage - @app.get("/usage/index.txt", include_in_schema=False) - async def redirect_usage_index_txt() -> RedirectResponse: - return RedirectResponse("/usage") - - @app.get("/unauthorized", include_in_schema=False) - async def serve_unauthorized_ui() -> FileResponse: - return FileResponse(UI_DIST_PATH / "unauthorized" / "index.html") - - # Add explicit route for /unauthorized/index.txt to redirect to /unauthorized - @app.get("/unauthorized/index.txt", include_in_schema=False) - async def redirect_unauthorized_index_txt() -> RedirectResponse: - return RedirectResponse("/unauthorized") - - @app.get("/favicon.ico", include_in_schema=False) - async def serve_favicon() -> FileResponse: - icon_path = UI_DIST_PATH / "icon.ico" - if icon_path.exists(): - return FileResponse(icon_path) - return FileResponse(UI_DIST_PATH / "favicon.ico") - - @app.get("/icon.ico", include_in_schema=False) - async def serve_icon() -> FileResponse: - return FileResponse(UI_DIST_PATH / "icon.ico") - - app.mount( - "/static", StaticFiles(directory=UI_DIST_PATH, check_dir=True), name="ui-static" - ) -else: - logger.warning( - f"UI dist directory not found at {UI_DIST_PATH}, skipping static file serving" - ) - - @app.get("/", include_in_schema=False) - async def root_fallback() -> dict: - return { - "name": global_settings.name, - "description": global_settings.description, - "version": __version__, - "status": "running", - "ui": "not available", - } - +setup_ui_routes(app) app.include_router(models_router) app.include_router(admin_router) diff --git a/routstr/core/tasks.py b/routstr/core/tasks.py new file mode 100644 index 00000000..c3207b43 --- /dev/null +++ b/routstr/core/tasks.py @@ -0,0 +1,134 @@ +import asyncio +from contextlib import asynccontextmanager +from typing import AsyncGenerator + +from fastapi import FastAPI + +from ..models.models import ( + cleanup_enabled_models_periodically, + update_sats_pricing, +) +from ..nostr.discovery import providers_cache_refresher +from ..nostr.listing import announce_provider +from ..payment.price import _update_prices, update_prices_periodically +from ..payment.wallet import periodic_payout +from ..proxy import get_upstreams, initialize_upstreams, refresh_model_maps_periodically +from ..upstream.helpers import refresh_upstreams_models_periodically +from .db import create_session, init_db, run_migrations +from .logging import get_logger +from .settings import SettingsService +from .settings import settings as global_settings + +logger = get_logger(__name__) + + +@asynccontextmanager +async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: + # Extract version from app if available, or use default/log without it + version = getattr(app, "version", "unknown") + + logger.info("Application startup initiated", extra={"version": version}) + + btc_price_task = None + pricing_task = None + payout_task = None + nip91_task = None + providers_task = None + models_refresh_task = None + models_cleanup_task = None + model_maps_refresh_task = None + + try: + # Run database migrations on startup + # This ensures the database schema is always up-to-date in production + # Migrations are idempotent - running them multiple times is safe + logger.info("Running database migrations") + run_migrations() + + # Initialize database connection pools + # This creates any tables that might not be tracked by migrations yet + await init_db() + + # Initialize application settings (env -> computed -> DB precedence) + async with create_session() as session: + s = await SettingsService.initialize(session) + + # Apply app metadata from settings + try: + app.title = s.name + app.description = s.description + except Exception: + pass + + # await ensure_models_bootstrapped() + + await _update_prices() + await initialize_upstreams() + + btc_price_task = asyncio.create_task(update_prices_periodically()) + pricing_task = asyncio.create_task(update_sats_pricing()) + if global_settings.models_refresh_interval_seconds > 0: + models_refresh_task = asyncio.create_task( + refresh_upstreams_models_periodically(get_upstreams()) + ) + models_cleanup_task = asyncio.create_task(cleanup_enabled_models_periodically()) + model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) + payout_task = asyncio.create_task(periodic_payout()) + nip91_task = asyncio.create_task(announce_provider()) + providers_task = asyncio.create_task(providers_cache_refresher()) + + yield + + except Exception as e: + logger.error( + "Application startup failed", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + raise + finally: + logger.info("Application shutdown initiated") + + if btc_price_task is not None: + btc_price_task.cancel() + if pricing_task is not None: + pricing_task.cancel() + if payout_task is not None: + payout_task.cancel() + if nip91_task is not None: + nip91_task.cancel() + if providers_task is not None: + providers_task.cancel() + if models_refresh_task is not None: + models_refresh_task.cancel() + if models_cleanup_task is not None: + models_cleanup_task.cancel() + if model_maps_refresh_task is not None: + model_maps_refresh_task.cancel() + + try: + tasks_to_wait = [] + if btc_price_task is not None: + tasks_to_wait.append(btc_price_task) + if pricing_task is not None: + tasks_to_wait.append(pricing_task) + if payout_task is not None: + tasks_to_wait.append(payout_task) + if nip91_task is not None: + tasks_to_wait.append(nip91_task) + if providers_task is not None: + tasks_to_wait.append(providers_task) + if models_refresh_task is not None: + tasks_to_wait.append(models_refresh_task) + if models_cleanup_task is not None: + tasks_to_wait.append(models_cleanup_task) + if model_maps_refresh_task is not None: + tasks_to_wait.append(model_maps_refresh_task) + + if tasks_to_wait: + await asyncio.gather(*tasks_to_wait, return_exceptions=True) + logger.info("Background tasks stopped successfully") + except Exception as e: + logger.error( + "Error stopping background tasks", + extra={"error": str(e), "error_type": type(e).__name__}, + ) diff --git a/routstr/core/ui.py b/routstr/core/ui.py new file mode 100644 index 00000000..0737d0fb --- /dev/null +++ b/routstr/core/ui.py @@ -0,0 +1,160 @@ +from pathlib import Path + +from fastapi import APIRouter, FastAPI +from fastapi.responses import FileResponse, RedirectResponse +from fastapi.staticfiles import StaticFiles + +from .logging import get_logger +from .settings import settings as global_settings + +logger = get_logger(__name__) + + +def setup_ui_routes(app: FastAPI) -> None: + UI_DIST_PATH = Path(__file__).parent.parent.parent / "ui_out" + + if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): + logger.info(f"Serving static UI from {UI_DIST_PATH}") + + app.mount( + "/_next", + StaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True), + name="next-static", + ) + + router = APIRouter() + + @router.get("/", include_in_schema=False) + async def serve_root_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "index.html") + + # Add explicit route for /index.txt to redirect to / + @router.get("/index.txt", include_in_schema=False) + async def redirect_index_txt() -> RedirectResponse: + return RedirectResponse("/") + + @router.get("/admin") + async def admin_redirect() -> FileResponse: + return FileResponse(UI_DIST_PATH / "index.html") + + @router.get("/dashboard", include_in_schema=False) + async def serve_dashboard_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "index.html") + + @router.get("/login", include_in_schema=False) + async def serve_login_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "login" / "index.html") + + # Add explicit route for /login/index.txt to redirect to /login + @router.get("/login/index.txt", include_in_schema=False) + async def redirect_login_index_txt() -> RedirectResponse: + return RedirectResponse("/login") + + @router.get("/model", include_in_schema=False) + async def serve_models_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "model" / "index.html") + + # Add explicit route for /model/index.txt to redirect to /model + @router.get("/model/index.txt", include_in_schema=False) + async def redirect_model_index_txt() -> RedirectResponse: + return RedirectResponse("/model") + + @router.get("/providers", include_in_schema=False) + async def serve_providers_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "providers" / "index.html") + + # Add explicit route for /providers/index.txt to redirect to /providers + @router.get("/providers/index.txt", include_in_schema=False) + async def redirect_providers_index_txt() -> RedirectResponse: + return RedirectResponse("/providers") + + @router.get("/settings", include_in_schema=False) + async def serve_settings_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "settings" / "index.html") + + # Add explicit route for /settings/index.txt to redirect to /settings + @router.get("/settings/index.txt", include_in_schema=False) + async def redirect_settings_index_txt() -> RedirectResponse: + return RedirectResponse("/settings") + + @router.get("/transactions", include_in_schema=False) + async def serve_transactions_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "transactions" / "index.html") + + # Add explicit route for /transactions/index.txt to redirect to /transactions + @router.get("/transactions/index.txt", include_in_schema=False) + async def redirect_transactions_index_txt() -> RedirectResponse: + return RedirectResponse("/transactions") + + @router.get("/balances", include_in_schema=False) + async def serve_balances_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "balances" / "index.html") + + # Add explicit route for /balances/index.txt to redirect to /balances + @router.get("/balances/index.txt", include_in_schema=False) + async def redirect_balances_index_txt() -> RedirectResponse: + return RedirectResponse("/balances") + + @router.get("/logs", include_in_schema=False) + async def serve_logs_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "logs" / "index.html") + + # Add explicit route for /logs/index.txt to redirect to /logs + @router.get("/logs/index.txt", include_in_schema=False) + async def redirect_logs_index_txt() -> RedirectResponse: + return RedirectResponse("/logs") + + @router.get("/usage", include_in_schema=False) + async def serve_usage_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "usage" / "index.html") + + # Add explicit route for /usage/index.txt to redirect to /usage + @router.get("/usage/index.txt", include_in_schema=False) + async def redirect_usage_index_txt() -> RedirectResponse: + return RedirectResponse("/usage") + + @router.get("/unauthorized", include_in_schema=False) + async def serve_unauthorized_ui() -> FileResponse: + return FileResponse(UI_DIST_PATH / "unauthorized" / "index.html") + + # Add explicit route for /unauthorized/index.txt to redirect to /unauthorized + @router.get("/unauthorized/index.txt", include_in_schema=False) + async def redirect_unauthorized_index_txt() -> RedirectResponse: + return RedirectResponse("/unauthorized") + + @router.get("/favicon.ico", include_in_schema=False) + async def serve_favicon() -> FileResponse: + icon_path = UI_DIST_PATH / "icon.ico" + if icon_path.exists(): + return FileResponse(icon_path) + return FileResponse(UI_DIST_PATH / "favicon.ico") + + @router.get("/icon.ico", include_in_schema=False) + async def serve_icon() -> FileResponse: + return FileResponse(UI_DIST_PATH / "icon.ico") + + app.include_router(router) + + app.mount( + "/static", + StaticFiles(directory=UI_DIST_PATH, check_dir=True), + name="ui-static", + ) + else: + logger.warning( + f"UI dist directory not found at {UI_DIST_PATH}, skipping static file serving" + ) + + router = APIRouter() + + @router.get("/", include_in_schema=False) + async def root_fallback() -> dict: + return { + "name": global_settings.name, + "description": global_settings.description, + "version": app.version, + "status": "running", + "ui": "not available", + } + + app.include_router(router) diff --git a/routstr/models/__init__.py b/routstr/models/__init__.py new file mode 100644 index 00000000..77e395b4 --- /dev/null +++ b/routstr/models/__init__.py @@ -0,0 +1,5 @@ +from .models import Model, async_fetch_openrouter_models + +__all__ = ["Model", "async_fetch_openrouter_models"] + +# specifically ai models that the node sells and not meaning database models diff --git a/routstr/models/algorithm.py b/routstr/models/algorithm.py new file mode 100644 index 00000000..9e6e3811 --- /dev/null +++ b/routstr/models/algorithm.py @@ -0,0 +1,300 @@ +"""Model prioritization algorithm for selecting cheapest upstream providers.""" + +from typing import TYPE_CHECKING + +from ..core import get_logger +from ..models.crud import _row_to_model +from ..upstream.helpers import resolve_model_alias + +if TYPE_CHECKING: + from ..upstream import BaseUpstreamProvider + from .models import Model + +logger = get_logger(__name__) + + +def calculate_model_cost_score(model: "Model") -> float: + """Calculate a representative cost score for a model. + + This score is used to compare models when multiple providers offer the same model. + Lower scores indicate cheaper models. + + The score is calculated as a weighted average of: + - Input token cost (weighted by typical input usage) + - Output token cost (weighted by typical output usage) + - Fixed request cost + + Args: + model: Model instance with pricing information + + Returns: + Float representing the cost score. Lower is better. + """ + pricing = model.pricing + + # Weight costs by typical usage patterns + # Assume average request: 1000 input tokens, 500 output tokens + TYPICAL_INPUT_TOKENS = 1000.0 + TYPICAL_OUTPUT_TOKENS = 500.0 + + # Calculate weighted cost in USD + input_cost = pricing.prompt * (TYPICAL_INPUT_TOKENS / 1000.0) + output_cost = pricing.completion * (TYPICAL_OUTPUT_TOKENS / 1000.0) + request_cost = pricing.request + + # Include additional costs if present + image_cost = ( + getattr(pricing, "image", 0.0) * 0.1 + ) # Weight lower as not every request uses images + web_search_cost = getattr(pricing, "web_search", 0.0) * 0.1 + reasoning_cost = getattr(pricing, "internal_reasoning", 0.0) * 0.2 + + total_cost = ( + input_cost + + output_cost + + request_cost + + image_cost + + web_search_cost + + reasoning_cost + ) + + return total_cost + + +def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: + """Calculate a penalty multiplier for certain providers. + + This allows applying policy-based adjustments beyond pure cost. + For example, preferring certain providers for reliability or features. + + Args: + provider: UpstreamProvider instance + + Returns: + Float multiplier to apply to cost (1.0 = no penalty, >1.0 = penalize) + """ + # Default: no penalty + penalty = 1.0 + + # Check if this is OpenRouter (can be identified by base URL) + base_url = getattr(provider, "base_url", "") + if "openrouter.ai" in base_url.lower(): + # Small penalty for OpenRouter to prefer other providers when costs are very close + # This maintains the original behavior of preferring non-OpenRouter providers + penalty = 1.001 # 0.1% penalty + + return penalty + + +def should_prefer_model( + candidate_model: "Model", + candidate_provider: "BaseUpstreamProvider", + current_model: "Model", + current_provider: "BaseUpstreamProvider", + alias: str, +) -> bool: + """Determine if candidate model should replace current model for an alias. + + This is the core decision function for model prioritization. It considers: + 1. Alias matching quality (exact match vs. canonical slug match) + 2. Model cost (lower is better) + 3. Provider penalties (e.g., slight preference against OpenRouter) + + Args: + candidate_model: The new model being considered + candidate_provider: Provider offering the candidate model + current_model: The currently selected model for this alias + current_provider: Provider offering the current model + alias: The model alias being mapped + + Returns: + True if candidate should replace current, False otherwise + """ + + def get_base_model_id(model_id: str) -> str: + """Get base model ID by removing provider prefix.""" + return model_id.split("/", 1)[1] if "/" in model_id else model_id + + def alias_priority(model: "Model") -> int: + """Rank how strong the mapping of alias->model is. + + Highest priority when alias exactly equals the model ID without provider prefix. + Next when alias equals canonical slug without prefix. Otherwise lowest. + """ + model_base = get_base_model_id(model.id) + if model_base == alias: + return 3 + if model.canonical_slug: + canonical_base = get_base_model_id(model.canonical_slug) + if canonical_base == alias: + return 2 + return 1 + + candidate_alias_priority = alias_priority(candidate_model) + current_alias_priority = alias_priority(current_model) + + # If candidate has better alias match, prefer it regardless of cost + if candidate_alias_priority > current_alias_priority: + return True + + # If current has better alias match, keep it regardless of cost + if current_alias_priority > candidate_alias_priority: + return False + + # Same alias priority - compare costs + candidate_cost = calculate_model_cost_score(candidate_model) + current_cost = calculate_model_cost_score(current_model) + + # Apply provider penalties + candidate_adjusted = candidate_cost * get_provider_penalty(candidate_provider) + current_adjusted = current_cost * get_provider_penalty(current_provider) + + # Prefer lower adjusted cost + should_replace = candidate_adjusted < current_adjusted + + # Log provider changes when candidate wins + if should_replace: + candidate_provider_name = getattr( + candidate_provider, "upstream_name", "unknown" + ) + current_provider_name = getattr(current_provider, "upstream_name", "unknown") + logger.debug( + f"Model selection for alias '{alias}': choosing {candidate_provider_name} " + f"(cost: ${candidate_adjusted:.6f}) over {current_provider_name} " + f"(cost: ${current_adjusted:.6f})" + ) + + return should_replace + + +def create_model_mappings( + upstreams: list["BaseUpstreamProvider"], + overrides_by_id: dict[str, tuple], + disabled_model_ids: set[str], +) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]: + """Create optimal model mappings based on cost and provider preferences. + + This is the main entry point for the algorithm. It processes all upstream providers + and creates three mappings based on cost optimization: + + 1. model_instances: alias -> Model (all model aliases mapped to their Model objects) + 2. provider_map: alias -> UpstreamProvider (which provider to use for each alias) + 3. unique_models: base_id -> Model (unique models without provider prefixes) + + The algorithm: + - Processes non-OpenRouter providers first (they're typically cheaper) + - Then processes OpenRouter models (they can still win if cheaper) + - For each model alias, uses should_prefer_model() to select the best provider + + Args: + upstreams: List of all upstream provider instances + overrides_by_id: Dict of model overrides from database {model_id: (ModelRow, fee)} + disabled_model_ids: Set of model IDs that should be excluded + + Returns: + Tuple of (model_instances, provider_map, unique_models) + """ + + model_instances: dict[str, "Model"] = {} + provider_map: dict[str, "BaseUpstreamProvider"] = {} + unique_models: dict[str, "Model"] = {} + + # Separate OpenRouter from other providers + openrouter: "BaseUpstreamProvider" | None = None + other_upstreams: list["BaseUpstreamProvider"] = [] + + for upstream in upstreams: + base_url = getattr(upstream, "base_url", "") + if base_url == "https://openrouter.ai/api/v1": + openrouter = upstream + else: + other_upstreams.append(upstream) + + def get_base_model_id(model_id: str) -> str: + """Get base model ID by removing provider prefix.""" + return model_id.split("/", 1)[1] if "/" in model_id else model_id + + def _maybe_set_alias( + alias: str, model: "Model", provider: "BaseUpstreamProvider" + ) -> None: + """Set alias to model/provider if not set or if new model is preferred.""" + existing_model = model_instances.get(alias) + if not existing_model: + # No existing mapping, set it + model_instances[alias] = model + provider_map[alias] = provider + else: + # Check if candidate should replace existing + existing_provider = provider_map[alias] + if should_prefer_model( + model, provider, existing_model, existing_provider, alias + ): + model_instances[alias] = model + provider_map[alias] = provider + + def process_provider_models( + upstream: "BaseUpstreamProvider", is_openrouter: bool = False + ) -> None: + """Process all models from a given provider.""" + upstream_prefix = getattr(upstream, "upstream_name", None) + + for model in upstream.get_cached_models(): + if not model.enabled or model.id in disabled_model_ids: + continue + + # Apply overrides if present + if model.id in overrides_by_id: + override_row, provider_fee = overrides_by_id[model.id] + model_to_use = _row_to_model( + override_row, apply_provider_fee=True, provider_fee=provider_fee + ) + else: + model_to_use = model + + # Add to unique models + base_id = get_base_model_id(model_to_use.id) + if not is_openrouter or base_id not in unique_models: + unique_model = model_to_use.copy(update={"id": base_id}) + unique_models[base_id] = unique_model + + # Get all aliases for this model + aliases = resolve_model_alias( + model_to_use.id, + model_to_use.canonical_slug, + alias_ids=model_to_use.alias_ids, + ) + + # Add prefixed alias if applicable + if upstream_prefix and "/" not in model_to_use.id: + prefixed_id = f"{upstream_prefix}/{model_to_use.id}" + if prefixed_id not in aliases: + aliases.append(prefixed_id) + + # Try to set each alias + for alias in aliases: + _maybe_set_alias(alias, model_to_use, upstream) + + # Process non-OpenRouter providers first (they're typically cheaper) + for upstream in other_upstreams: + process_provider_models(upstream, is_openrouter=False) + + # Process OpenRouter last - models only win if they're cheaper or better matched + if openrouter: + process_provider_models(openrouter, is_openrouter=True) + + # Log provider distribution + provider_counts: dict[str, int] = {} + for provider in provider_map.values(): + provider_name = getattr(provider, "upstream_name", "unknown") + provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 + + logger.debug( + "Created model mappings", + extra={ + "unique_model_count": len(unique_models), + "total_alias_count": len(model_instances), + "provider_distribution": provider_counts, + }, + ) + + return model_instances, provider_map, unique_models diff --git a/routstr/models/crud.py b/routstr/models/crud.py new file mode 100644 index 00000000..c18c8ec7 --- /dev/null +++ b/routstr/models/crud.py @@ -0,0 +1,180 @@ +import json +from pathlib import Path + +from ..core import get_logger +from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow +from ..core.settings import settings +from ..payment.price import sats_usd_price +from .metadata import fetch_openrouter_models +from .models import ( + Architecture, + Model, + Pricing, + TopProvider, + _calculate_usd_max_costs, + _update_model_sats_pricing, + is_openrouter_upstream, +) + +logger = get_logger(__name__) + + +def load_models() -> list[Model]: + """Load model definitions from a JSON file or auto-generate from OpenRouter API. + + The file path can be specified via the ``MODELS_PATH`` environment variable. + If a user-provided models.json exists, it will be used. Otherwise, models are + automatically fetched from OpenRouter API in memory. If the example file exists + and no user file is provided, it will be used as a fallback. + """ + + try: + models_path = Path(settings.models_path) + except Exception: + models_path = Path("models.json") + + # Check if user has actively provided a models.json file + if models_path.exists(): + logger.info(f"Loading models from user-provided file: {models_path}") + try: + with models_path.open("r") as f: + data = json.load(f) + return [Model(**model) for model in data.get("models", [])] # type: ignore + except Exception as e: + logger.error(f"Error loading models from {models_path}: {e}") + # Fall through to auto-generation + + # Only auto-generate from OpenRouter when upstream is OpenRouter + if not is_openrouter_upstream(): + logger.info( + "Skipping auto-generation from OpenRouter because upstream_base_url is not https://openrouter.ai/api/v1" + ) + return [] + + logger.info("Auto-generating models from OpenRouter API") + try: + source_filter = settings.source or None + except Exception: + source_filter = None + source_filter = source_filter if source_filter and source_filter.strip() else None + + models_data = fetch_openrouter_models(source_filter=source_filter) + if not models_data: + logger.error("Failed to fetch models from OpenRouter API") + return [] + + logger.info(f"Successfully fetched {len(models_data)} models from OpenRouter API") + return [Model(**model) for model in models_data] # type: ignore + + +def _row_to_model( + row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 +) -> Model: + architecture = json.loads(row.architecture) + pricing = json.loads(row.pricing) + per_request_limits = ( + json.loads(row.per_request_limits) if row.per_request_limits else None + ) + top_provider_dict = json.loads(row.top_provider) if row.top_provider else None + + if apply_provider_fee and isinstance(pricing, dict): + pricing = {k: float(v) * provider_fee for k, v in pricing.items()} + + if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0: + pricing["request"] = max(pricing.get("request", 0.0), 0.0) + + parsed_pricing = Pricing.parse_obj(pricing) + model = Model( + id=row.id, + name=row.name, + created=row.created, + description=row.description, + context_length=row.context_length, + architecture=Architecture.parse_obj(architecture), + pricing=parsed_pricing, + sats_pricing=None, + per_request_limits=per_request_limits, + top_provider=TopProvider.parse_obj(top_provider_dict) + if top_provider_dict + else None, + enabled=row.enabled, + upstream_provider_id=row.upstream_provider_id, + canonical_slug=getattr(row, "canonical_slug", None), + ) + + if apply_provider_fee: + ( + parsed_pricing.max_prompt_cost, + parsed_pricing.max_completion_cost, + parsed_pricing.max_cost, + ) = _calculate_usd_max_costs(model) + + try: + sats_to_usd = sats_usd_price() + model = _update_model_sats_pricing(model, sats_to_usd) + except Exception as e: + logger.warning(f"Could not calculate sats pricing: {e}") + + return model + + +def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]: + return { + "id": model.id, + "name": model.name, + "created": model.created, + "description": model.description, + "context_length": model.context_length, + "architecture": json.dumps(model.architecture.dict()), + "pricing": json.dumps(model.pricing.dict()), + "sats_pricing": json.dumps(model.sats_pricing.dict()) + if model.sats_pricing + else None, + "per_request_limits": json.dumps(model.per_request_limits) + if model.per_request_limits is not None + else None, + "top_provider": json.dumps(model.top_provider.dict()) + if model.top_provider is not None + else None, + "enabled": model.enabled, + "upstream_provider_id": model.upstream_provider_id, + } + + +async def list_models( + session: AsyncSession, + upstream_id: int, + include_disabled: bool = False, +) -> list[Model]: + from sqlmodel import select + + query = select(ModelRow) + if upstream_id is not None: + query = query.where(ModelRow.upstream_provider_id == upstream_id) + if not include_disabled: + query = query.where(ModelRow.enabled) + + rows = (await session.exec(query)).all() # type: ignore + provider_result = await session.exec(select(UpstreamProviderRow)) + providers_by_id = {p.id: p for p in provider_result.all()} + return [ + _row_to_model( + r, + apply_provider_fee=True, + provider_fee=providers_by_id[r.upstream_provider_id].provider_fee + if r.upstream_provider_id in providers_by_id + else 1.01, + ) + for r in rows + ] + + +async def get_model_by_id( + model_id: str, provider_id: int, session: AsyncSession +) -> Model | None: + row = await session.get(ModelRow, (model_id, provider_id)) + if not row or not row.enabled: + return None + provider = await session.get(UpstreamProviderRow, provider_id) + provider_fee = provider.provider_fee if provider else 1.01 + return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) diff --git a/routstr/models/metadata.py b/routstr/models/metadata.py new file mode 100644 index 00000000..0fbc64eb --- /dev/null +++ b/routstr/models/metadata.py @@ -0,0 +1,89 @@ +import json +from typing import Final +from urllib.request import urlopen + +import httpx + +from ..core import get_logger + +logger = get_logger(__name__) + +DEFAULT_EXCLUDED_MODEL_IDS: Final[set[str]] = { + "openrouter/auto", + "google/gemini-2.5-pro-exp-03-25", + "opengvlab/internvl3-78b", + "openrouter/sonoma-dusk-alpha", + "openrouter/sonoma-sky-alpha", +} + + +def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: + """Fetches model information from OpenRouter API.""" + base_url = "https://openrouter.ai/api/v1" + + try: + with urlopen(f"{base_url}/models") as response: + data = json.loads(response.read().decode("utf-8")) + + models_data: list[dict] = [] + for model in data.get("data", []): + model_id = model.get("id", "") + + if source_filter: + source_prefix = f"{source_filter}/" + if not model_id.startswith(source_prefix): + continue + + model = dict(model) + model["id"] = model_id[len(source_prefix) :] + model_id = model["id"] + + if ( + "(free)" in model.get("name", "") + or model_id in DEFAULT_EXCLUDED_MODEL_IDS + ): + continue + + models_data.append(model) + + return models_data + except Exception as e: + logger.error(f"Error fetching models from OpenRouter API: {e}") + return [] + + +async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: + """Asynchronously fetch model information from OpenRouter API.""" + base_url = "https://openrouter.ai/api/v1" + + try: + async with httpx.AsyncClient() as client: + response = await client.get(f"{base_url}/models", timeout=30) + response.raise_for_status() + data = response.json() + + models_data: list[dict] = [] + for model in data.get("data", []): + model_id = model.get("id", "") + + if source_filter: + source_prefix = f"{source_filter}/" + if not model_id.startswith(source_prefix): + continue + + model = dict(model) + model["id"] = model_id[len(source_prefix) :] + model_id = model["id"] + + if ( + "(free)" in model.get("name", "") + or model_id in DEFAULT_EXCLUDED_MODEL_IDS + ): + continue + + models_data.append(model) + + return models_data + except Exception as e: + logger.error(f"Error (async) fetching models from OpenRouter API: {e}") + return [] diff --git a/routstr/payment/models.py b/routstr/models/models.py similarity index 63% rename from routstr/payment/models.py rename to routstr/models/models.py index 2e3fbd9c..2331df5a 100644 --- a/routstr/payment/models.py +++ b/routstr/models/models.py @@ -2,10 +2,7 @@ import asyncio import json import random from pathlib import Path -from typing import Final -from urllib.request import urlopen -import httpx from fastapi import APIRouter, Depends from pydantic.v1 import BaseModel from sqlmodel import select @@ -14,20 +11,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..core.db import ModelRow, create_session, get_session from ..core.logging import get_logger from ..core.settings import settings -from .price import sats_usd_price +from ..payment.price import sats_usd_price +from .metadata import async_fetch_openrouter_models logger = get_logger(__name__) models_router = APIRouter() -DEFAULT_EXCLUDED_MODEL_IDS: Final[set[str]] = { - "openrouter/auto", - "google/gemini-2.5-pro-exp-03-25", - "opengvlab/internvl3-78b", - "openrouter/sonoma-dusk-alpha", - "openrouter/sonoma-sky-alpha", -} - class Architecture(BaseModel): modality: str @@ -75,78 +65,6 @@ class Model(BaseModel): return hash(self.id) -def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: - """Fetches model information from OpenRouter API.""" - base_url = "https://openrouter.ai/api/v1" - - try: - with urlopen(f"{base_url}/models") as response: - data = json.loads(response.read().decode("utf-8")) - - models_data: list[dict] = [] - for model in data.get("data", []): - model_id = model.get("id", "") - - if source_filter: - source_prefix = f"{source_filter}/" - if not model_id.startswith(source_prefix): - continue - - model = dict(model) - model["id"] = model_id[len(source_prefix) :] - model_id = model["id"] - - if ( - "(free)" in model.get("name", "") - or model_id in DEFAULT_EXCLUDED_MODEL_IDS - ): - continue - - models_data.append(model) - - return models_data - except Exception as e: - logger.error(f"Error fetching models from OpenRouter API: {e}") - return [] - - -async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: - """Asynchronously fetch model information from OpenRouter API.""" - base_url = "https://openrouter.ai/api/v1" - - try: - async with httpx.AsyncClient() as client: - response = await client.get(f"{base_url}/models", timeout=30) - response.raise_for_status() - data = response.json() - - models_data: list[dict] = [] - for model in data.get("data", []): - model_id = model.get("id", "") - - if source_filter: - source_prefix = f"{source_filter}/" - if not model_id.startswith(source_prefix): - continue - - model = dict(model) - model["id"] = model_id[len(source_prefix) :] - model_id = model["id"] - - if ( - "(free)" in model.get("name", "") - or model_id in DEFAULT_EXCLUDED_MODEL_IDS - ): - continue - - models_data.append(model) - - return models_data - except Exception as e: - logger.error(f"Error (async) fetching models from OpenRouter API: {e}") - return [] - - def is_openrouter_upstream() -> bool: try: base = (settings.upstream_base_url or "").strip().rstrip("/") @@ -155,171 +73,6 @@ def is_openrouter_upstream() -> bool: return base.lower() == "https://openrouter.ai/api/v1" -def load_models() -> list[Model]: - """Load model definitions from a JSON file or auto-generate from OpenRouter API. - - The file path can be specified via the ``MODELS_PATH`` environment variable. - If a user-provided models.json exists, it will be used. Otherwise, models are - automatically fetched from OpenRouter API in memory. If the example file exists - and no user file is provided, it will be used as a fallback. - """ - - try: - models_path = Path(settings.models_path) - except Exception: - models_path = Path("models.json") - - # Check if user has actively provided a models.json file - if models_path.exists(): - logger.info(f"Loading models from user-provided file: {models_path}") - try: - with models_path.open("r") as f: - data = json.load(f) - return [Model(**model) for model in data.get("models", [])] # type: ignore - except Exception as e: - logger.error(f"Error loading models from {models_path}: {e}") - # Fall through to auto-generation - - # Only auto-generate from OpenRouter when upstream is OpenRouter - if not is_openrouter_upstream(): - logger.info( - "Skipping auto-generation from OpenRouter because upstream_base_url is not https://openrouter.ai/api/v1" - ) - return [] - - logger.info("Auto-generating models from OpenRouter API") - try: - source_filter = settings.source or None - except Exception: - source_filter = None - source_filter = source_filter if source_filter and source_filter.strip() else None - - models_data = fetch_openrouter_models(source_filter=source_filter) - if not models_data: - logger.error("Failed to fetch models from OpenRouter API") - return [] - - logger.info(f"Successfully fetched {len(models_data)} models from OpenRouter API") - return [Model(**model) for model in models_data] # type: ignore - - -def _row_to_model( - row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 -) -> Model: - architecture = json.loads(row.architecture) - pricing = json.loads(row.pricing) - per_request_limits = ( - json.loads(row.per_request_limits) if row.per_request_limits else None - ) - top_provider_dict = json.loads(row.top_provider) if row.top_provider else None - - if apply_provider_fee and isinstance(pricing, dict): - pricing = {k: float(v) * provider_fee for k, v in pricing.items()} - - if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0: - pricing["request"] = max(pricing.get("request", 0.0), 0.0) - - parsed_pricing = Pricing.parse_obj(pricing) - model = Model( - id=row.id, - name=row.name, - created=row.created, - description=row.description, - context_length=row.context_length, - architecture=Architecture.parse_obj(architecture), - pricing=parsed_pricing, - sats_pricing=None, - per_request_limits=per_request_limits, - top_provider=TopProvider.parse_obj(top_provider_dict) - if top_provider_dict - else None, - enabled=row.enabled, - upstream_provider_id=row.upstream_provider_id, - canonical_slug=getattr(row, "canonical_slug", None), - ) - - if apply_provider_fee: - ( - parsed_pricing.max_prompt_cost, - parsed_pricing.max_completion_cost, - parsed_pricing.max_cost, - ) = _calculate_usd_max_costs(model) - - try: - sats_to_usd = sats_usd_price() - model = _update_model_sats_pricing(model, sats_to_usd) - except Exception as e: - logger.warning(f"Could not calculate sats pricing: {e}") - - return model - - -def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]: - return { - "id": model.id, - "name": model.name, - "created": model.created, - "description": model.description, - "context_length": model.context_length, - "architecture": json.dumps(model.architecture.dict()), - "pricing": json.dumps(model.pricing.dict()), - "sats_pricing": json.dumps(model.sats_pricing.dict()) - if model.sats_pricing - else None, - "per_request_limits": json.dumps(model.per_request_limits) - if model.per_request_limits is not None - else None, - "top_provider": json.dumps(model.top_provider.dict()) - if model.top_provider is not None - else None, - "enabled": model.enabled, - "upstream_provider_id": model.upstream_provider_id, - } - - -async def list_models( - session: AsyncSession, - upstream_id: int, - include_disabled: bool = False, -) -> list[Model]: - from sqlmodel import select - - from ..core.db import UpstreamProviderRow - - query = select(ModelRow) - if upstream_id is not None: - query = query.where(ModelRow.upstream_provider_id == upstream_id) - if not include_disabled: - query = query.where(ModelRow.enabled) - - rows = (await session.exec(query)).all() # type: ignore - provider_result = await session.exec(select(UpstreamProviderRow)) - providers_by_id = {p.id: p for p in provider_result.all()} - return [ - _row_to_model( - r, - apply_provider_fee=True, - provider_fee=providers_by_id[r.upstream_provider_id].provider_fee - if r.upstream_provider_id in providers_by_id - else 1.01, - ) - for r in rows - ] - - -async def get_model_by_id( - model_id: str, provider_id: int, session: AsyncSession -) -> Model | None: - from ..core.db import UpstreamProviderRow - - row = await session.get(ModelRow, (model_id, provider_id)) - if not row or not row.enabled: - return None - provider = await session.get(UpstreamProviderRow, provider_id) - provider_fee = provider.provider_fee if provider else 1.01 - return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) - - def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]: """Calculate max costs in USD based on model context/token limits. @@ -432,6 +185,8 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model: async def ensure_models_bootstrapped() -> None: + from .crud import _model_to_row_payload + async with create_session() as s: existing = (await s.exec(select(ModelRow.id).limit(1))).all() # type: ignore if existing: @@ -462,7 +217,9 @@ async def ensure_models_bootstrapped() -> None: source_filter = src if src and src.strip() else None except Exception: pass - models_to_insert = fetch_openrouter_models(source_filter=source_filter) + models_to_insert = await async_fetch_openrouter_models( + source_filter=source_filter + ) elif not models_to_insert: logger.info( "No models.json found and upstream is not OpenRouter; skipping bootstrap" @@ -650,6 +407,8 @@ async def refresh_models_periodically() -> None: - Does not overwrite existing rows - Sleeps according to settings.models_refresh_interval_seconds; disabled when 0 """ + from .crud import _model_to_row_payload + interval = getattr(settings, "models_refresh_interval_seconds", 0) if not interval or interval <= 0: return @@ -672,7 +431,7 @@ async def refresh_models_periodically() -> None: except Exception: source_filter = None - models = fetch_openrouter_models(source_filter=source_filter) + models = await async_fetch_openrouter_models(source_filter=source_filter) if not models: await asyncio.sleep(interval) continue diff --git a/routstr/nostr/__init__.py b/routstr/nostr/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/routstr/discovery.py b/routstr/nostr/discovery.py similarity index 99% rename from routstr/discovery.py rename to routstr/nostr/discovery.py index a2680145..0d715fc5 100644 --- a/routstr/discovery.py +++ b/routstr/nostr/discovery.py @@ -8,8 +8,8 @@ import httpx import websockets from fastapi import APIRouter -from .core.logging import get_logger -from .core.settings import settings +from ..core.logging import get_logger +from ..core.settings import settings logger = get_logger(__name__) diff --git a/routstr/nip91.py b/routstr/nostr/listing.py similarity index 99% rename from routstr/nip91.py rename to routstr/nostr/listing.py index 3a0a40d2..0e1c3ff1 100644 --- a/routstr/nip91.py +++ b/routstr/nostr/listing.py @@ -18,15 +18,15 @@ from nostr.key import PrivateKey from nostr.message_type import ClientMessageType from nostr.relay_manager import RelayManager -from .core import get_logger -from .core.settings import settings +from ..core import get_logger +from ..core.settings import settings logger = get_logger(__name__) def get_app_version() -> str | None: try: - from .core.main import __version__ as imported_version + from ..core.main import __version__ as imported_version return imported_version except Exception: diff --git a/routstr/payment/__init__.py b/routstr/payment/__init__.py index 0ca1ed03..fd3c495a 100644 --- a/routstr/payment/__init__.py +++ b/routstr/payment/__init__.py @@ -1,4 +1,4 @@ -from .cost_calculation import CostData, CostDataError, MaxCostData, calculate_cost +from .cost import CostData, CostDataError, MaxCostData, calculate_cost __all__ = [ "CostData", diff --git a/routstr/payment/cashu.py b/routstr/payment/cashu.py new file mode 100644 index 00000000..1b5774ed --- /dev/null +++ b/routstr/payment/cashu.py @@ -0,0 +1,74 @@ +from cashu.wallet.helpers import deserialize_token_from_string +from fastapi import HTTPException + +from ..core import get_logger + +logger = get_logger(__name__) + + +def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None: + if x_cashu := headers.get("x-cashu", None): + cashu_token = x_cashu + logger.debug( + "Using X-Cashu token", + extra={ + "token_preview": cashu_token[:20] + "..." + if len(cashu_token) > 20 + else cashu_token + }, + ) + elif auth := headers.get("authorization", None): + cashu_token = auth.split(" ")[1] if len(auth.split(" ")) > 1 else "" + logger.debug( + "Using Authorization header token", + extra={ + "token_preview": cashu_token[:20] + "..." + if len(cashu_token) > 20 + else cashu_token + }, + ) + else: + logger.error("No authentication token provided") + raise HTTPException(status_code=401, detail="Unauthorized") + + # Handle empty token + if not cashu_token: + logger.error("Empty token provided") + raise HTTPException( + status_code=401, + detail={ + "error": { + "message": "API key or Cashu token required", + "type": "invalid_request_error", + "code": "missing_api_key", + } + }, + ) + + # Handle regular API keys (sk-*) + if cashu_token.startswith("sk-"): + return + + try: + token_obj = deserialize_token_from_string(cashu_token) + except Exception: + # Invalid token format - let the auth system handle it + raise HTTPException( + status_code=401, + detail="Invalid authentication token format", + ) + + amount_msat = ( + token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000 + ) + + if max_cost_for_model > amount_msat: + raise HTTPException( + status_code=413, + detail={ + "reason": "Insufficient balance", + "amount_required_msat": max_cost_for_model, + "model": body.get("model", "unknown"), + "type": "minimum_balance_required", + }, + ) diff --git a/routstr/payment/cost.py b/routstr/payment/cost.py new file mode 100644 index 00000000..6ac72695 --- /dev/null +++ b/routstr/payment/cost.py @@ -0,0 +1,426 @@ +import base64 +import math +from io import BytesIO +from typing import Any + +import httpx +from PIL import Image +from pydantic.v1 import BaseModel + +from ..core import get_logger +from ..core.db import AsyncSession +from ..core.settings import settings +from ..models.models import Model + +logger = get_logger(__name__) + + +class CostData(BaseModel): + base_msats: int + input_msats: int + output_msats: int + total_msats: int + + +class MaxCostData(CostData): + pass + + +class CostDataError(BaseModel): + message: str + code: str + + +async def calculate_cost( # todo: can be sync + response_data: dict, max_cost: int, session: AsyncSession +) -> CostData | MaxCostData | CostDataError: + """ + Calculate the cost of an API request based on token usage. + + Args: + response_data: Response data containing usage information + max_cost: Maximum cost in millisats + + Returns: + Cost data or error information + """ + logger.debug( + "Starting cost calculation", + extra={ + "max_cost_msats": max_cost, + "has_usage_data": "usage" in response_data, + "response_model": response_data.get("model", "unknown"), + }, + ) + + cost_data = MaxCostData( + base_msats=max_cost, + input_msats=0, + output_msats=0, + total_msats=max_cost, + ) + + if "usage" not in response_data or response_data["usage"] is None: + logger.warning( + "No usage data in response, using base cost only", + extra={ + "max_cost_msats": max_cost, + "model": response_data.get("model", "unknown"), + }, + ) + return cost_data + + 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( + "Using model-based pricing", + extra={"model": response_model}, + ) + + from ..proxy import get_model_instance + + model_obj = get_model_instance(response_model) + + if not model_obj: + logger.error( + "Invalid model in response", + extra={"response_model": response_model}, + ) + return CostDataError( + message=f"Invalid model in response: {response_model}", + code="model_not_found", + ) + + if not model_obj.sats_pricing: + logger.error( + "Model pricing not defined", + extra={"model": response_model, "model_id": response_model}, + ) + return CostDataError( + message="Model pricing not defined", code="pricing_not_found" + ) + + try: + mspp = float(model_obj.sats_pricing.prompt) + mspc = float(model_obj.sats_pricing.completion) + except Exception: + return CostDataError(message="Invalid pricing data", code="pricing_invalid") + + MSATS_PER_1K_INPUT_TOKENS = mspp * 1_000_000.0 + MSATS_PER_1K_OUTPUT_TOKENS = mspc * 1_000_000.0 + + logger.info( + "Applied model-specific pricing", + extra={ + "model": response_model, + "input_price_msats_per_1k": MSATS_PER_1K_INPUT_TOKENS, + "output_price_msats_per_1k": MSATS_PER_1K_OUTPUT_TOKENS, + }, + ) + + if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS): + logger.warning( + "No token pricing configured, using base cost", + extra={ + "base_cost_msats": max_cost, + "model": response_data.get("model", "unknown"), + }, + ) + return cost_data + + input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0) + output_tokens = response_data.get("usage", {}).get("completion_tokens", 0) + + 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) + + logger.info( + "Calculated token-based cost", + extra={ + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "input_cost_msats": input_msats, + "output_cost_msats": output_msats, + "total_cost_msats": token_based_cost, + "model": response_data.get("model", "unknown"), + }, + ) + + return CostData( + base_msats=0, + input_msats=int(input_msats), + output_msats=int(output_msats), + total_msats=token_based_cost, + ) + + +def get_max_cost_for_model(model_obj: Model) -> int: + """Get the maximum cost for a specific model from providers with overrides.""" + if settings.fixed_pricing: + default_cost_msats = settings.fixed_cost_per_request * 1000 + return max(settings.min_request_msat, default_cost_msats) + + if model_obj.sats_pricing: + max_cost = ( + model_obj.sats_pricing.max_cost + * 1000 + * (1 - settings.tolerance_percentage / 100) + ) + calculated_msats = int(max_cost) + return max(settings.min_request_msat, calculated_msats) + + logger.warning( + "Model pricing not found, using fixed cost", + extra={ + "model": model_obj.id, + "default_cost_msats": settings.fixed_cost_per_request * 1000, + }, + ) + return max(settings.min_request_msat, settings.fixed_cost_per_request * 1000) + + +async def calculate_discounted_max_cost( + max_cost_for_model: int, + body: dict, + model_obj: Any | None = None, +) -> int: + """Calculate the discounted max cost for a request using model pricing when available.""" + if settings.fixed_pricing: + return max_cost_for_model + + model = body.get("model", "unknown") + + model_pricing = model_obj.sats_pricing if model_obj else None + if not model_pricing: + return max_cost_for_model + + tol = settings.tolerance_percentage + tol_factor = max(0.0, 1 - float(tol) / 100.0) + max_prompt_allowed_sats = model_pricing.max_prompt_cost * tol_factor + max_completion_allowed_sats = model_pricing.max_completion_cost * tol_factor + + adjusted = max_cost_for_model + + if messages := body.get("messages"): + prompt_tokens = _estimate_tokens(messages) + + image_tokens = await _estimate_image_tokens_in_messages(messages) + if image_tokens > 0: + logger.debug( + "Found images in request", + extra={ + "model": model, + "image_tokens": image_tokens, + }, + ) + prompt_tokens += image_tokens + + estimated_prompt_delta_sats = ( + max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt + ) + if estimated_prompt_delta_sats > 0: + adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000) + + max_tokens_raw = body.get("max_tokens", None) + if max_tokens_raw is not None: + try: + max_tokens_int = int(max_tokens_raw) + except (TypeError, ValueError): + logger.warning( + "Invalid max_tokens; ignoring in cost adjustment", + extra={"max_tokens": str(max_tokens_raw)[:64], "model": model}, + ) + else: + estimated_completion_delta_sats = ( + max_completion_allowed_sats - max_tokens_int * model_pricing.completion + ) + if estimated_completion_delta_sats > 0: + adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000) + + logger.debug( + "Discounted max cost computed", + extra={ + "model": model, + "original_msats": max_cost_for_model, + "adjusted_msats": adjusted, + "tolerance_pct": tol, + }, + ) + + return max(0, adjusted) + + +def _estimate_tokens(messages: list) -> int: + """Estimate tokens for text content, excluding image_url fields.""" + total = 0 + for msg in messages: + if isinstance(msg, dict): + content = msg.get("content") + if isinstance(content, str): + total += len(content) + elif isinstance(content, list): + total += sum( + len(item.get("text", "")) + for item in content + if isinstance(item, dict) and item.get("type") == "text" + ) + return total // 3 + + +def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: + """Extract image dimensions from image bytes.""" + try: + img = Image.open(BytesIO(image_data)) + return img.size + except Exception as e: + logger.warning( + "Failed to get image dimensions, using default", + extra={"error": str(e)}, + ) + return (512, 512) + + +async def _fetch_image_from_url(url: str) -> bytes | None: + """Fetch image from URL.""" + try: + async with httpx.AsyncClient(timeout=10.0) as client: + response = await client.get(url) + response.raise_for_status() + return response.content + except Exception as e: + logger.warning( + "Failed to fetch image from URL", + extra={"error": str(e), "url": url[:100]}, + ) + return None + + +def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int: + """Calculate image tokens based on OpenAI's vision pricing. + + For low detail: 85 tokens + For high detail/auto: 85 base tokens + 170 tokens per 512px tile + """ + if detail == "low": + return 85 + + if width > 2048 or height > 2048: + aspect_ratio = width / height + if width > height: + width = 2048 + height = int(width / aspect_ratio) + else: + height = 2048 + width = int(height * aspect_ratio) + + if width > 768 or height > 768: + aspect_ratio = width / height + if width > height: + width = 768 + height = int(width / aspect_ratio) + else: + height = 768 + width = int(height * aspect_ratio) + + tiles_width = (width + 511) // 512 + tiles_height = (height + 511) // 512 + num_tiles = tiles_width * tiles_height + + return 85 + (170 * num_tiles) + + +async def _estimate_image_tokens_in_messages(messages: list) -> int: + """Estimate total tokens for all images in messages. + + Supports both base64 encoded images and image URLs. + """ + total_image_tokens = 0 + + for message in messages: + if not isinstance(message, dict): + continue + + content = message.get("content") + if not content: + continue + + if isinstance(content, str): + continue + + if not isinstance(content, list): + continue + + for content_item in content: + if not isinstance(content_item, dict): + continue + + content_type = content_item.get("type") + if content_type not in ("image_url", "input_image"): + continue + + image_url_data = content_item.get("image_url") + if not image_url_data: + continue + + if isinstance(image_url_data, str): + url = image_url_data + detail = "auto" + elif isinstance(image_url_data, dict): + url = image_url_data.get("url", "") + detail = image_url_data.get("detail", "auto") + else: + continue + + if not url: + continue + + if url.startswith("data:image/"): + try: + header, base64_data = url.split(",", 1) + image_bytes = base64.b64decode(base64_data) + width, height = _get_image_dimensions(image_bytes) + tokens = _calculate_image_tokens(width, height, detail) + total_image_tokens += tokens + logger.debug( + "Calculated tokens for base64 image", + extra={ + "width": width, + "height": height, + "detail": detail, + "tokens": tokens, + }, + ) + except Exception as e: + logger.warning( + "Failed to process base64 image", + extra={"error": str(e)}, + ) + total_image_tokens += 85 + else: + image_bytes_or_none = await _fetch_image_from_url(url) + if image_bytes_or_none: + width, height = _get_image_dimensions(image_bytes_or_none) + tokens = _calculate_image_tokens(width, height, detail) + total_image_tokens += tokens + logger.debug( + "Calculated tokens for URL image", + extra={ + "url": url[:100], + "width": width, + "height": height, + "detail": detail, + "tokens": tokens, + }, + ) + else: + total_image_tokens += 85 + + return total_image_tokens diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py deleted file mode 100644 index f8eb4ffb..00000000 --- a/routstr/payment/cost_calculation.py +++ /dev/null @@ -1,156 +0,0 @@ -import math - -from pydantic.v1 import BaseModel - -from ..core import get_logger -from ..core.db import AsyncSession -from ..core.settings import settings - -logger = get_logger(__name__) - - -class CostData(BaseModel): - base_msats: int - input_msats: int - output_msats: int - total_msats: int - - -class MaxCostData(CostData): - pass - - -class CostDataError(BaseModel): - message: str - code: str - - -async def calculate_cost( # todo: can be sync - response_data: dict, max_cost: int, session: AsyncSession -) -> CostData | MaxCostData | CostDataError: - """ - Calculate the cost of an API request based on token usage. - - Args: - response_data: Response data containing usage information - max_cost: Maximum cost in millisats - - Returns: - Cost data or error information - """ - logger.debug( - "Starting cost calculation", - extra={ - "max_cost_msats": max_cost, - "has_usage_data": "usage" in response_data, - "response_model": response_data.get("model", "unknown"), - }, - ) - - cost_data = MaxCostData( - base_msats=max_cost, - input_msats=0, - output_msats=0, - total_msats=max_cost, - ) - - if "usage" not in response_data or response_data["usage"] is None: - logger.warning( - "No usage data in response, using base cost only", - extra={ - "max_cost_msats": max_cost, - "model": response_data.get("model", "unknown"), - }, - ) - return cost_data - - 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( - "Using model-based pricing", - extra={"model": response_model}, - ) - - from ..proxy import get_model_instance - - model_obj = get_model_instance(response_model) - - if not model_obj: - logger.error( - "Invalid model in response", - extra={"response_model": response_model}, - ) - return CostDataError( - message=f"Invalid model in response: {response_model}", - code="model_not_found", - ) - - if not model_obj.sats_pricing: - logger.error( - "Model pricing not defined", - extra={"model": response_model, "model_id": response_model}, - ) - return CostDataError( - message="Model pricing not defined", code="pricing_not_found" - ) - - try: - mspp = float(model_obj.sats_pricing.prompt) - mspc = float(model_obj.sats_pricing.completion) - except Exception: - return CostDataError(message="Invalid pricing data", code="pricing_invalid") - - MSATS_PER_1K_INPUT_TOKENS = mspp * 1_000_000.0 - MSATS_PER_1K_OUTPUT_TOKENS = mspc * 1_000_000.0 - - logger.info( - "Applied model-specific pricing", - extra={ - "model": response_model, - "input_price_msats_per_1k": MSATS_PER_1K_INPUT_TOKENS, - "output_price_msats_per_1k": MSATS_PER_1K_OUTPUT_TOKENS, - }, - ) - - if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS): - logger.warning( - "No token pricing configured, using base cost", - extra={ - "base_cost_msats": max_cost, - "model": response_data.get("model", "unknown"), - }, - ) - return cost_data - - input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0) - output_tokens = response_data.get("usage", {}).get("completion_tokens", 0) - - 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) - - logger.info( - "Calculated token-based cost", - extra={ - "input_tokens": input_tokens, - "output_tokens": output_tokens, - "input_cost_msats": input_msats, - "output_cost_msats": output_msats, - "total_cost_msats": token_based_cost, - "model": response_data.get("model", "unknown"), - }, - ) - - return CostData( - base_msats=0, - input_msats=int(input_msats), - output_msats=int(output_msats), - total_msats=token_based_cost, - ) diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index f1027bd6..826ae56a 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,396 +1,18 @@ -import base64 import json import math -from io import BytesIO -from typing import Any -import httpx from fastapi import HTTPException, Response from fastapi.requests import Request -from PIL import Image -from sqlmodel.ext.asyncio.session import AsyncSession +from sqlmodel import col, update from ..core import get_logger -from ..core.settings import settings -from ..wallet import deserialize_token_from_string +from ..core.db import AsyncSession, TemporaryCredit +from ..payment.cost import CostData, CostDataError, MaxCostData, calculate_cost logger = get_logger(__name__) -def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None: - if x_cashu := headers.get("x-cashu", None): - cashu_token = x_cashu - logger.debug( - "Using X-Cashu token", - extra={ - "token_preview": cashu_token[:20] + "..." - if len(cashu_token) > 20 - else cashu_token - }, - ) - elif auth := headers.get("authorization", None): - cashu_token = auth.split(" ")[1] if len(auth.split(" ")) > 1 else "" - logger.debug( - "Using Authorization header token", - extra={ - "token_preview": cashu_token[:20] + "..." - if len(cashu_token) > 20 - else cashu_token - }, - ) - else: - logger.error("No authentication token provided") - raise HTTPException(status_code=401, detail="Unauthorized") - - # Handle empty token - if not cashu_token: - logger.error("Empty token provided") - raise HTTPException( - status_code=401, - detail={ - "error": { - "message": "API key or Cashu token required", - "type": "invalid_request_error", - "code": "missing_api_key", - } - }, - ) - - # Handle regular API keys (sk-*) - if cashu_token.startswith("sk-"): - return - - try: - token_obj = deserialize_token_from_string(cashu_token) - except Exception: - # Invalid token format - let the auth system handle it - raise HTTPException( - status_code=401, - detail="Invalid authentication token format", - ) - - amount_msat = ( - token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000 - ) - - if max_cost_for_model > amount_msat: - raise HTTPException( - status_code=413, - detail={ - "reason": "Insufficient balance", - "amount_required_msat": max_cost_for_model, - "model": body.get("model", "unknown"), - "type": "minimum_balance_required", - }, - ) - - -async def get_max_cost_for_model( - model: str, - session: AsyncSession, - model_obj: Any | None = None, -) -> int: - """Get the maximum cost for a specific model from providers with overrides.""" - logger.debug( - "Getting max cost for model", - extra={ - "model": model, - "fixed_pricing": settings.fixed_pricing, - }, - ) - - if settings.fixed_pricing: - default_cost_msats = settings.fixed_cost_per_request * 1000 - logger.debug( - "Using fixed cost pricing", - extra={"cost_msats": default_cost_msats, "model": model}, - ) - return max(settings.min_request_msat, default_cost_msats) - - if not model_obj: - from ..proxy import get_model_instance - - model_obj = get_model_instance(model) - - if not model_obj: - fallback_msats = settings.fixed_cost_per_request * 1000 - logger.warning( - "Model not found in providers or overrides", - extra={ - "requested_model": model, - "using_default_cost": fallback_msats, - }, - ) - return max(settings.min_request_msat, fallback_msats) - - if model_obj.sats_pricing: - try: - max_cost = ( - model_obj.sats_pricing.max_cost - * 1000 - * (1 - settings.tolerance_percentage / 100) - ) - logger.debug( - "Found model-specific max cost", - extra={"model": model, "max_cost_msats": max_cost}, - ) - calculated_msats = int(max_cost) - return max(settings.min_request_msat, calculated_msats) - except Exception as e: - logger.error( - "Error calculating max cost from model pricing", - extra={"model": model, "error": str(e)}, - ) - - logger.warning( - "Model pricing not found, using fixed cost", - extra={ - "model": model, - "default_cost_msats": settings.fixed_cost_per_request * 1000, - }, - ) - return max(settings.min_request_msat, settings.fixed_cost_per_request * 1000) - - -async def calculate_discounted_max_cost( - max_cost_for_model: int, - body: dict, - model_obj: Any | None = None, -) -> int: - """Calculate the discounted max cost for a request using model pricing when available.""" - if settings.fixed_pricing: - return max_cost_for_model - - model = body.get("model", "unknown") - - model_pricing = model_obj.sats_pricing if model_obj else None - if not model_pricing: - return max_cost_for_model - - tol = settings.tolerance_percentage - tol_factor = max(0.0, 1 - float(tol) / 100.0) - max_prompt_allowed_sats = model_pricing.max_prompt_cost * tol_factor - max_completion_allowed_sats = model_pricing.max_completion_cost * tol_factor - - adjusted = max_cost_for_model - - if messages := body.get("messages"): - prompt_tokens = estimate_tokens(messages) - - image_tokens = await estimate_image_tokens_in_messages(messages) - if image_tokens > 0: - logger.debug( - "Found images in request", - extra={ - "model": model, - "image_tokens": image_tokens, - }, - ) - prompt_tokens += image_tokens - - estimated_prompt_delta_sats = ( - max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt - ) - if estimated_prompt_delta_sats > 0: - adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000) - - max_tokens_raw = body.get("max_tokens", None) - if max_tokens_raw is not None: - try: - max_tokens_int = int(max_tokens_raw) - except (TypeError, ValueError): - logger.warning( - "Invalid max_tokens; ignoring in cost adjustment", - extra={"max_tokens": str(max_tokens_raw)[:64], "model": model}, - ) - else: - estimated_completion_delta_sats = ( - max_completion_allowed_sats - max_tokens_int * model_pricing.completion - ) - if estimated_completion_delta_sats > 0: - adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000) - - logger.debug( - "Discounted max cost computed", - extra={ - "model": model, - "original_msats": max_cost_for_model, - "adjusted_msats": adjusted, - "tolerance_pct": tol, - }, - ) - - return max(0, adjusted) - - -def estimate_tokens(messages: list) -> int: - """Estimate tokens for text content, excluding image_url fields.""" - total = 0 - for msg in messages: - if isinstance(msg, dict): - content = msg.get("content") - if isinstance(content, str): - total += len(content) - elif isinstance(content, list): - total += sum( - len(item.get("text", "")) - for item in content - if isinstance(item, dict) and item.get("type") == "text" - ) - return total // 3 - - -def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: - """Extract image dimensions from image bytes.""" - try: - img = Image.open(BytesIO(image_data)) - return img.size - except Exception as e: - logger.warning( - "Failed to get image dimensions, using default", - extra={"error": str(e)}, - ) - return (512, 512) - - -async def _fetch_image_from_url(url: str) -> bytes | None: - """Fetch image from URL.""" - try: - async with httpx.AsyncClient(timeout=10.0) as client: - response = await client.get(url) - response.raise_for_status() - return response.content - except Exception as e: - logger.warning( - "Failed to fetch image from URL", - extra={"error": str(e), "url": url[:100]}, - ) - return None - - -def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int: - """Calculate image tokens based on OpenAI's vision pricing. - - For low detail: 85 tokens - For high detail/auto: 85 base tokens + 170 tokens per 512px tile - """ - if detail == "low": - return 85 - - if width > 2048 or height > 2048: - aspect_ratio = width / height - if width > height: - width = 2048 - height = int(width / aspect_ratio) - else: - height = 2048 - width = int(height * aspect_ratio) - - if width > 768 or height > 768: - aspect_ratio = width / height - if width > height: - width = 768 - height = int(width / aspect_ratio) - else: - height = 768 - width = int(height * aspect_ratio) - - tiles_width = (width + 511) // 512 - tiles_height = (height + 511) // 512 - num_tiles = tiles_width * tiles_height - - return 85 + (170 * num_tiles) - - -async def estimate_image_tokens_in_messages(messages: list) -> int: - """Estimate total tokens for all images in messages. - - Supports both base64 encoded images and image URLs. - """ - total_image_tokens = 0 - - for message in messages: - if not isinstance(message, dict): - continue - - content = message.get("content") - if not content: - continue - - if isinstance(content, str): - continue - - if not isinstance(content, list): - continue - - for content_item in content: - if not isinstance(content_item, dict): - continue - - content_type = content_item.get("type") - if content_type not in ("image_url", "input_image"): - continue - - image_url_data = content_item.get("image_url") - if not image_url_data: - continue - - if isinstance(image_url_data, str): - url = image_url_data - detail = "auto" - elif isinstance(image_url_data, dict): - url = image_url_data.get("url", "") - detail = image_url_data.get("detail", "auto") - else: - continue - - if not url: - continue - - if url.startswith("data:image/"): - try: - header, base64_data = url.split(",", 1) - image_bytes = base64.b64decode(base64_data) - width, height = _get_image_dimensions(image_bytes) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for base64 image", - extra={ - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - except Exception as e: - logger.warning( - "Failed to process base64 image", - extra={"error": str(e)}, - ) - total_image_tokens += 85 - else: - image_bytes_or_none = await _fetch_image_from_url(url) - if image_bytes_or_none: - width, height = _get_image_dimensions(image_bytes_or_none) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for URL image", - extra={ - "url": url[:100], - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - else: - total_image_tokens += 85 - - return total_image_tokens - - +# TODO: remove and replace with custom HTTPException def create_error_response( error_type: str, message: str, @@ -414,3 +36,377 @@ def create_error_response( media_type="application/json", headers={"X-Cashu": token} if token else {}, ) + + +# Request payment handlers, todo: maybe mode somewhere else... + + +async def pay_for_request( + key: TemporaryCredit, cost_per_request: int, session: AsyncSession +) -> int: + """Process payment for a request.""" + + logger.info( + "Processing payment for request", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "current_balance": key.balance, + "required_cost": cost_per_request, + "sufficient_balance": key.balance >= cost_per_request, + }, + ) + + if key.total_balance_msat < cost_per_request: + logger.warning( + "Insufficient balance for request", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "balance": key.balance, + "reserved_balance": key.reserved_balance, + "required": cost_per_request, + "shortfall": cost_per_request - key.total_balance_msat, + }, + ) + + raise HTTPException( + status_code=402, + detail={ + "error": { + "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", + "type": "insufficient_quota", + "code": "insufficient_balance", + } + }, + ) + + logger.debug( + "Charging base cost for request", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "cost": cost_per_request, + "balance_before": key.balance, + }, + ) + + # Charge the base cost for the request atomically to avoid race conditions + stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .where(col(TemporaryCredit.balance) >= cost_per_request) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) + cost_per_request + ) + ) + result = await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + + if result.rowcount == 0: + logger.error( + "Concurrent request depleted balance", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "required_cost": cost_per_request, + "current_balance": key.balance, + }, + ) + + # Another concurrent request spent the balance first + raise HTTPException( + status_code=402, + detail={ + "error": { + "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", + "type": "insufficient_quota", + "code": "insufficient_balance", + } + }, + ) + + await session.refresh(key) + + logger.info( + "Payment processed successfully", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "charged_amount": cost_per_request, + "new_balance": key.balance, + }, + ) + + return cost_per_request + + +async def revert_pay_for_request( + key: TemporaryCredit, session: AsyncSession, cost_per_request: int +) -> None: + stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) - cost_per_request + ) + ) + + result = await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + if result.rowcount == 0: + logger.error( + "Failed to revert payment - insufficient reserved balance", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "cost_to_revert": cost_per_request, + "current_reserved_balance": key.reserved_balance, + }, + ) + raise HTTPException( + status_code=402, + detail={ + "error": { + "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.", + "type": "payment_error", + "code": "payment_error", + } + }, + ) + await session.refresh(key) + + +async def adjust_payment_for_tokens( + key: TemporaryCredit, + response_data: dict, + session: AsyncSession, + deducted_max_cost: int, +) -> dict: + """ + Adjusts the payment based on token usage in the response. + This is called after the initial payment and the upstream request is complete. + Returns cost data to be included in the response. + """ + model = response_data.get("model", "unknown") + + logger.debug( + "Starting payment adjustment for tokens", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": model, + "deducted_max_cost": deducted_max_cost, + "current_balance": key.balance, + "has_usage": "usage" in response_data, + }, + ) + + match await calculate_cost(response_data, deducted_max_cost, session): + case MaxCostData() as cost: + logger.debug( + "Using max cost data (no token adjustment)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": model, + "max_cost": cost.total_msats, + }, + ) + # Finalize by releasing reservation and charging max cost + finalize_stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) + - deducted_max_cost, + balance=col(TemporaryCredit.balance) - cost.total_msats, + ) + ) + result = await session.exec(finalize_stmt) # type: ignore[call-overload] + await session.commit() + if result.rowcount == 0: + logger.error( + "Failed to finalize max-cost payment - insufficient reserved balance", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + "current_reserved_balance": key.reserved_balance, + "total_cost": cost.total_msats, + "model": model, + }, + ) + else: + await session.refresh(key) + logger.info( + "Max cost payment finalized", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "charged_amount": cost.total_msats, + "new_balance": key.balance, + "model": model, + }, + ) + return cost.dict() + + case CostData() as cost: + # If token-based pricing is enabled and base cost is 0, use token-based cost + # Otherwise, token cost is additional to the base cost + cost_difference = cost.total_msats - deducted_max_cost + total_cost_msats: int = math.ceil(cost.total_msats) + + logger.info( + "Calculated token-based cost", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": model, + "token_cost": cost.total_msats, + "deducted_max_cost": deducted_max_cost, + "cost_difference": cost_difference, + "input_msats": cost.input_msats, + "output_msats": cost.output_msats, + }, + ) + + if cost_difference == 0: + logger.debug( + "Finalizing with exact reserved cost", + extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, + ) + finalize_stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) + - deducted_max_cost, + balance=col(TemporaryCredit.balance) - total_cost_msats, + ) + ) + await session.exec(finalize_stmt) # type: ignore[call-overload] + await session.commit() + await session.refresh(key) + 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", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "additional_charge": cost_difference, + "current_balance": key.balance, + "sufficient_balance": key.balance >= cost_difference, + "model": model, + }, + ) + + finalize_stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) + - deducted_max_cost, + balance=col(TemporaryCredit.balance) - total_cost_msats, + ) + ) + result = await session.exec(finalize_stmt) # type: ignore[call-overload] + await session.commit() + + if result.rowcount: + cost.total_msats = total_cost_msats + await session.refresh(key) + + logger.info( + "Finalized payment with additional charge", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "charged_amount": total_cost_msats, + "new_balance": key.balance, + "model": model, + }, + ) + else: + logger.warning( + "Failed to finalize additional charge (concurrent operation)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "attempted_charge": total_cost_msats, + "model": model, + }, + ) + else: + # Refund some of the base cost + refund = abs(cost_difference) + logger.info( + "Refunding excess payment", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "refund_amount": refund, + "current_balance": key.balance, + "model": model, + }, + ) + + refund_stmt = ( + update(TemporaryCredit) + .where(col(TemporaryCredit.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(TemporaryCredit.reserved_balance) + - deducted_max_cost, + balance=col(TemporaryCredit.balance) - total_cost_msats, + ) + ) + result = await session.exec(refund_stmt) # type: ignore[call-overload] + await session.commit() + + if result.rowcount == 0: + logger.error( + "Failed to finalize payment - insufficient reserved balance", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + "current_reserved_balance": key.reserved_balance, + "total_cost": total_cost_msats, + "model": model, + }, + ) + # Still return the cost data even if we couldn't properly finalize + # The reservation was already made, so the user has paid + + cost.total_msats = total_cost_msats + await session.refresh(key) + + logger.info( + "Refund processed successfully", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "refunded_amount": refund, + "new_balance": key.balance, + "final_cost": cost.total_msats, + "model": model, + }, + ) + + return cost.dict() + + case CostDataError() as error: + logger.error( + "Cost calculation error during payment adjustment", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": model, + "error_message": error.message, + "error_code": error.code, + }, + ) + + raise HTTPException( + status_code=400, + detail={ + "error": { + "message": error.message, + "type": "invalid_request_error", + "code": error.code, + } + }, + ) + # Fallback return to satisfy type checker; execution should not reach here + return { + "base_msats": deducted_max_cost, + "input_msats": 0, + "output_msats": 0, + "total_msats": deducted_max_cost, + } diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 04395311..65f88f9c 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -305,3 +305,12 @@ async def raw_send_to_lnurl( quote_id=melt_quote_resp.quote, ) return final_amount + + +async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: + from .wallet import get_wallet + + wallet = await get_wallet(mint, unit) + proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] + proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) + return await raw_send_to_lnurl(wallet, proofs, address, unit) diff --git a/routstr/wallet.py b/routstr/payment/wallet.py similarity index 87% rename from routstr/wallet.py rename to routstr/payment/wallet.py index 34569b85..edafcfd9 100644 --- a/routstr/wallet.py +++ b/routstr/payment/wallet.py @@ -7,9 +7,9 @@ from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet from sqlmodel import col, update -from .core import db, get_logger -from .core.settings import settings -from .payment.lnurl import raw_send_to_lnurl +from ..core import db, get_logger +from ..core.settings import settings +from ..payment.lnurl import raw_send_to_lnurl logger = get_logger(__name__) @@ -19,9 +19,7 @@ async def get_balance(unit: str) -> int: return wallet.available_balance.amount -async def recieve_token( - token: str, -) -> tuple[int, str, str]: # amount, unit, mint_url +async def recieve_token(token: str) -> tuple[int, str, str]: # amount, unit, mint_url token_obj = deserialize_token_from_string(token) if len(token_obj.keysets) > 1: raise ValueError("Multiple keysets per token currently not supported") @@ -104,7 +102,7 @@ async def swap_to_primary_mint( async def credit_balance( - cashu_token: str, key: db.ApiKey, session: db.AsyncSession + cashu_token: str, credit: db.TemporaryCredit, session: db.AsyncSession ) -> int: logger.info( "credit_balance: Starting token redemption", @@ -126,22 +124,22 @@ async def credit_balance( logger.info( "credit_balance: Updating balance", - extra={"old_balance": key.balance, "credit_amount": amount}, + extra={"old_balance": credit.balance, "credit_amount": amount}, ) # Use atomic SQL UPDATE to prevent race conditions during concurrent topups stmt = ( - update(db.ApiKey) - .where(col(db.ApiKey.hashed_key) == key.hashed_key) - .values(balance=(db.ApiKey.balance) + amount) + update(db.TemporaryCredit) + .where(col(db.TemporaryCredit.hashed_key) == credit.hashed_key) + .values(balance=(db.TemporaryCredit.balance) + amount) ) await session.exec(stmt) # type: ignore[call-overload] await session.commit() - await session.refresh(key) + await session.refresh(credit) logger.info( "credit_balance: Balance updated successfully", - extra={"new_balance": key.balance}, + extra={"new_balance": credit.balance}, ) logger.info( @@ -351,34 +349,3 @@ async def periodic_payout() -> None: f"Error sending payout: {type(e).__name__}", extra={"error": str(e)}, ) - - -async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: - wallet = await get_wallet(mint, unit) - proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] - proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) - return await raw_send_to_lnurl(wallet, proofs, address, unit) - - -# class Payment: -# """ -# Stores all cashu payment related data -# """ - -# def __init__(self, token: str) -> None: -# self.initial_token = token -# amount, unit, mint_url = self.parse_token(token) -# self.amount = amount -# self.unit = unit -# self.mint_url = mint_url - -# self.claimed_proofs = redeem_to_proofs(token) - -# def parse_token(self, token: str) -> tuple[int, CurrencyUnit, str]: -# raise NotImplementedError - -# def refund_full(self) -> None: -# raise NotImplementedError - -# def refund_partial(self, amount: int) -> None: -# raise NotImplementedError diff --git a/routstr/proxy.py b/routstr/proxy.py index ce558fc5..f960439f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -5,24 +5,20 @@ from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse from sqlmodel import select -from .algorithm import create_model_mappings -from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key +from .auth import get_credit from .core import get_logger from .core.db import ( - ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, create_session, get_session, ) -from .payment.helpers import ( - calculate_discounted_max_cost, - check_token_balance, - create_error_response, - get_max_cost_for_model, -) -from .payment.models import Model +from .models.algorithm import create_model_mappings +from .models.models import Model +from .payment.cashu import check_token_balance +from .payment.cost import calculate_discounted_max_cost, get_max_cost_for_model +from .payment.helpers import pay_for_request, revert_pay_for_request from .upstream import BaseUpstreamProvider from .upstream.helpers import init_upstreams @@ -130,18 +126,21 @@ async def refresh_model_maps_periodically() -> None: ) -@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) -async def proxy( - request: Request, path: str, session: AsyncSession = Depends(get_session) -) -> Response | StreamingResponse: +async def parse_request(request: Request) -> tuple[dict, dict, bytes]: headers = dict(request.headers) + body_bytes = await request.body() + body_dict = parse_request_body_json(body_bytes) + return headers, body_dict, body_bytes + +def ensure_authorization(headers: dict) -> None: if "x-cashu" not in headers and "authorization" not in headers.keys(): - return create_error_response( - "unauthorized", "Unauthorized", 401, request=request - ) + # return create_error_response("unauthorized", "Unauthorized", 401, request=request) + raise HTTPException(status_code=401, detail="Unauthorized") - logger.info( # TODO: move to middleware, async + +def log_request(request: Request, path: str) -> None: + logger.info( "Received proxy request", extra={ "method": request.method, @@ -151,58 +150,49 @@ async def proxy( }, ) - request_body = await request.body() - request_body_dict = parse_request_body_json(request_body, path) - model_id = request_body_dict.get("model", "unknown") +def get_model_and_upstream(request_body: dict) -> tuple[Model, BaseUpstreamProvider]: + model_id = request_body.get("model", "unknown") - model_obj = get_model_instance(model_id) - if not model_obj: - return create_error_response( - "invalid_model", f"Model '{model_id}' not found", 400, request=request - ) + if not (model := get_model_instance(model_id)): + raise HTTPException(status_code=400, detail=f"Model '{model_id}' not found") - upstream = get_provider_for_model(model_id) - if not upstream: - return create_error_response( - "invalid_model", - f"No provider found for model '{model_id}'", - 400, - request=request, - ) + if not (upstream := get_provider_for_model(model_id)): + raise HTTPException(status_code=400, detail=f"No provider with '{model_id}'") - _max_cost_for_model = await get_max_cost_for_model( - model=model_id, session=session, model_obj=model_obj - ) + return model, upstream + + +@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) +async def proxy( + request: Request, + path: str, + session: AsyncSession = Depends(get_session), +) -> Response | StreamingResponse: + headers, body_dict, body_bytes = await parse_request(request) + + ensure_authorization(headers) + + log_request(request, path) + + model, upstream = get_model_and_upstream(body_dict) + + _max_cost_for_model = get_max_cost_for_model(model_obj=model) max_cost_for_model = await calculate_discounted_max_cost( - _max_cost_for_model, request_body_dict, model_obj=model_obj + _max_cost_for_model, body_dict, model_obj=model ) - check_token_balance(headers, request_body_dict, max_cost_for_model) + check_token_balance(headers, body_dict, max_cost_for_model) if x_cashu := headers.get("x-cashu", None): return await upstream.handle_x_cashu( - request, x_cashu, path, max_cost_for_model, model_obj + request, x_cashu, path, max_cost_for_model, model ) elif auth := headers.get("authorization", None): - key = await get_bearer_token_key(headers, path, session, auth) - - else: - if request.method not in ["GET"]: - raise HTTPException( - status_code=401, - detail={ - "error": {"type": "invalid_request_error", "code": "unauthorized"} - }, - ) - - logger.debug("Processing unauthenticated GET request", extra={"path": path}) - # TODO: why is this needed? can we remove it? - headers = upstream.prepare_headers(dict(request.headers)) - return await upstream.forward_get_request(request, path, headers) + key = await get_credit(auth, session) # Only pay for request if we have request body data (for completions endpoints) - if request_body_dict: + if body_dict: await pay_for_request(key, max_cost_for_model, session) # Prepare headers for upstream @@ -213,11 +203,11 @@ async def proxy( request, path, headers, - request_body, + body_bytes, key, max_cost_for_model, session, - model_obj, + model, ) if response.status_code != 200: @@ -241,87 +231,7 @@ async def proxy( return response -async def get_bearer_token_key( - headers: dict, path: str, session: AsyncSession, auth: str -) -> ApiKey: - """Handle bearer token authentication proxy requests.""" - bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else "" - refund_address = headers.get("Refund-LNURL", None) - key_expiry_time = headers.get("Key-Expiry-Time", None) - - logger.debug( - "Processing bearer token", - extra={ - "path": path, - "has_refund_address": bool(refund_address), - "has_expiry_time": bool(key_expiry_time), - "bearer_key_preview": bearer_key[:20] + "..." - if len(bearer_key) > 20 - else bearer_key, - }, - ) - - # Validate key_expiry_time header - if key_expiry_time: - try: - key_expiry_time = int(key_expiry_time) # type: ignore - logger.debug( - "Key expiry time validated", - extra={"expiry_time": key_expiry_time, "path": path}, - ) - except ValueError: - logger.error( - "Invalid Key-Expiry-Time header", - extra={"key_expiry_time": key_expiry_time, "path": path}, - ) - raise HTTPException( - status_code=400, - detail="Invalid Key-Expiry-Time: must be a valid Unix timestamp", - ) - if not refund_address: - logger.error( - "Missing Refund-LNURL header with Key-Expiry-Time", - extra={"path": path, "expiry_time": key_expiry_time}, - ) - raise HTTPException( - status_code=400, - detail="Error: Refund-LNURL header required when using Key-Expiry-Time", - ) - else: - key_expiry_time = None - - try: - key = await validate_bearer_key( - bearer_key, - session, - refund_address, - key_expiry_time, # type: ignore - ) - logger.info( - "Bearer token validated successfully", - extra={ - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - }, - ) - return key - except Exception as e: - logger.error( - "Bearer token validation failed", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "path": path, - "bearer_key_preview": bearer_key[:20] + "..." - if len(bearer_key) > 20 - else bearer_key, - }, - ) - raise - - -def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]: +def parse_request_body_json(request_body: bytes) -> dict[str, Any]: request_body_dict = {} if request_body: try: @@ -341,7 +251,6 @@ def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]: logger.debug( "Request body parsed", extra={ - "path": path, "body_keys": list(request_body_dict.keys()), "model": request_body_dict.get("model", "not_specified"), }, @@ -351,7 +260,6 @@ def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]: "Invalid JSON in request body", extra={ "error": str(e), - "path": path, "body_preview": request_body[:200].decode(errors="ignore") if request_body else "empty", diff --git a/routstr/upstream/anthropic.py b/routstr/upstream/anthropic.py index 3f228e9c..3de8e454 100644 --- a/routstr/upstream/anthropic.py +++ b/routstr/upstream/anthropic.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from ..payment.models import Model, async_fetch_openrouter_models +from ..models.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 576ade2f..dca95270 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -11,28 +11,23 @@ import httpx from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from ..auth import adjust_payment_for_tokens from ..core import get_logger -from ..core.db import ApiKey, AsyncSession, create_session +from ..core.db import AsyncSession, TemporaryCredit, create_session +from ..payment.helpers import adjust_payment_for_tokens if TYPE_CHECKING: from ..core.db import UpstreamProviderRow -from ..payment.cost_calculation import ( - CostData, - CostDataError, - MaxCostData, - calculate_cost, -) -from ..payment.helpers import create_error_response -from ..payment.models import ( +from ..models.models import ( Model, Pricing, _calculate_usd_max_costs, _update_model_sats_pricing, ) +from ..payment.cost import CostData, CostDataError, MaxCostData, calculate_cost +from ..payment.helpers import create_error_response from ..payment.price import sats_usd_price -from ..wallet import recieve_token, send_token +from ..payment.wallet import recieve_token, send_token logger = get_logger(__name__) @@ -143,7 +138,6 @@ class BaseUpstreamProvider: # Explicitly define the list of supported compression encodings headers["accept-encoding"] = "gzip, deflate, br, identity" - logger.debug( "Headers prepared for upstream", extra={ @@ -338,7 +332,7 @@ class BaseUpstreamProvider: ) async def handle_streaming_chat_completion( - self, response: httpx.Response, key: ApiKey, max_cost_for_model: int + self, response: httpx.Response, key: TemporaryCredit, max_cost_for_model: int ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -512,7 +506,7 @@ class BaseUpstreamProvider: ) await finalize_without_usage() raise - + # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) response_headers.pop("content-encoding", None) @@ -521,13 +515,13 @@ class BaseUpstreamProvider: return StreamingResponse( stream_with_cost(max_cost_for_model), status_code=response.status_code, - headers=response_headers, + headers=response_headers, ) async def handle_non_streaming_chat_completion( self, response: httpx.Response, - key: ApiKey, + key: TemporaryCredit, session: AsyncSession, deducted_max_cost: int, ) -> Response: @@ -633,7 +627,7 @@ class BaseUpstreamProvider: path: str, headers: dict, request_body: bytes | None, - key: ApiKey, + key: TemporaryCredit, max_cost_for_model: int, session: AsyncSession, model_obj: Model, diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 3d91550f..712e9968 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -10,7 +10,7 @@ if TYPE_CHECKING: from ..core import get_logger from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session -from ..payment.models import Model +from ..models import Model from .base import BaseUpstreamProvider logger = get_logger(__name__) diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 24703eb2..651084ca 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -9,8 +9,8 @@ from fastapi.responses import Response, StreamingResponse from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow - from ..payment.models import Model + from ..core.db import AsyncSession, TemporaryCredit, UpstreamProviderRow + from ..models import Model from ..core.logging import get_logger @@ -73,7 +73,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): path: str, headers: dict, request_body: bytes | None, - key: ApiKey, + key: TemporaryCredit, max_cost_for_model: int, session: AsyncSession, model_obj: Model, @@ -102,7 +102,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): async def fetch_models(self) -> list[Model]: """Fetch models from Ollama API using /api/tags endpoint.""" - from ..payment.models import Architecture, Model, Pricing, TopProvider + from ..models.models import Architecture, Model, Pricing, TopProvider try: async with httpx.AsyncClient(timeout=30.0) as client: @@ -199,7 +199,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): async def refresh_models_cache(self) -> None: """Refresh the in-memory models cache from upstream API.""" try: - from ..payment.models import _update_model_sats_pricing + from ..models.models import _update_model_sats_pricing from ..payment.price import sats_usd_price models = await self.fetch_models() @@ -252,7 +252,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): Returns: Model with provider fee applied to pricing and max costs calculated """ - from ..payment.models import Model, Pricing, _calculate_usd_max_costs + from ..models.models import Model, Pricing, _calculate_usd_max_costs adjusted_pricing = Pricing.parse_obj( {k: v * self.provider_fee for k, v in model.pricing.dict().items()} diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py index 11cc4336..5de05829 100644 --- a/routstr/upstream/openai.py +++ b/routstr/upstream/openai.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from ..payment.models import Model, async_fetch_openrouter_models +from ..models.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index cc0e2908..eb599593 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from ..payment.models import Model, async_fetch_openrouter_models +from ..models.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: diff --git a/routstr/upstream/perplexity.py b/routstr/upstream/perplexity.py index b73881d8..1c791d33 100644 --- a/routstr/upstream/perplexity.py +++ b/routstr/upstream/perplexity.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from ..payment.models import Model, async_fetch_openrouter_models +from ..models.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 99e2d35a..0dad4924 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from ..payment.models import Model, async_fetch_openrouter_models +from ..models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: