diff --git a/.env.example b/.env.example index 4a9b3123..b9883225 100644 --- a/.env.example +++ b/.env.example @@ -1,5 +1,5 @@ -NAME = "Your Routstr Proxy Name" -DESCRIPTION = "A short Description" +# NAME = "Your Routstr Proxy Name" +# DESCRIPTION = "A short Description" # Any openai-compatible api endpoint UPSTREAM_BASE_URL="https://api.openai.com/v1" @@ -7,14 +7,14 @@ UPSTREAM_API_KEY="sk-21212121212121212121212121212121" # UPSTREAM_PROVIDER_FEE=1 # 1 = no fees, 1.05 = 5% fees # Lightning address used to receive funds -RECEIVE_LN_ADDRESS="shroominic@walletofsatoshi.com" # When your cashu balance reaches this number of sats, send the funds to RECEIVE_LN_ADDRESS. -MINIMUM_PAYOUT = "100" +# RECEIVE_LN_ADDRESS="user@minibits.cash" +#MINIMUM_PAYOUT = "100" # If set to true, pricing is loaded from the file specified by MODELS_PATH # Defaults to "models.json" and falls back to "models.example.json" if missing -MODEL_BASED_PRICING = "true" +# MODEL_BASED_PRICING = "true" # MODELS_PATH="models.json" # Costs in Sats, if MODEL_BASED_PRICING is set to false @@ -27,13 +27,13 @@ MODEL_BASED_PRICING = "true" # ADMIN_PASSWORD="" # Public Endpoint -HTTP_URL="https://your.domain.com" +# HTTP_URL="https://your.domain.com" # Tor Endpoint (copy from docker logs) # ONION_URL=".onion" -RELAYS="wss://relay.routstr.com,wss://relay.nostr.band" -CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org" +# RELAYS="wss://relay.routstr.com,wss://relay.nostr.band" +# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org" # Development # DEBUG=TRUE diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 3637ea7c..e8099634 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -26,7 +26,6 @@ jobs: - name: Install dependencies run: | uv sync --dev - uv run python setup.py develop - name: Run linting with ruff run: | diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 5374cd44..013aa775 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -259,7 +259,7 @@ Before requesting review, ensure: ```text routstr-proxy/ -├── router/ # Main application code +├── routstr/ # Main application code │ ├── core/ # Core functionality │ │ ├── admin.py # Admin interface │ │ ├── db.py # Database models and operations @@ -284,10 +284,10 @@ routstr-proxy/ ### Key Components -- **FastAPI Application**: Main API server in `router/core/main.py` -- **Database Models**: SQLModel definitions in `router/core/db.py` -- **Payment Logic**: Cashu integration and cost calculation in `router/payment/` -- **Proxy Handler**: Request forwarding logic in `router/proxy.py` +- **FastAPI Application**: Main API server in `routstr/core/main.py` +- **Database Models**: SQLModel definitions in `routstr/core/db.py` +- **Payment Logic**: Cashu integration and cost calculation in `routstr/payment/` +- **Proxy Handler**: Request forwarding logic in `routstr/proxy.py` ## Documentation diff --git a/Dockerfile b/Dockerfile index ba15bfb8..adf5a7df 100644 --- a/Dockerfile +++ b/Dockerfile @@ -12,6 +12,7 @@ RUN apk add --no-cache \ RUN apk add git COPY uv.lock pyproject.toml ./ +RUN mkdir -p /routstr RUN uv add git+https://github.com/saschanaz/secp256k1-py.git#branch=upgrade060 # RUN uv sync @@ -25,4 +26,4 @@ ENV PYTHONUNBUFFERED=1 EXPOSE 8000 -CMD ["/.venv/bin/fastapi", "run", "router", "--host", "0.0.0.0"] +CMD ["/.venv/bin/fastapi", "run", "routstr", "--host", "0.0.0.0"] diff --git a/Makefile b/Makefile index b965ef57..79215dec 100644 --- a/Makefile +++ b/Makefile @@ -92,7 +92,7 @@ docker-down: lint: @echo "🔍 Running linting checks..." $(RUFF) check . - $(MYPY) router/ --ignore-missing-imports + $(MYPY) routstr/ --ignore-missing-imports format: @echo "✨ Formatting code..." @@ -101,7 +101,7 @@ format: type-check: @echo "🔎 Running type checks..." - $(MYPY) router/ --ignore-missing-imports + $(MYPY) routstr/ --ignore-missing-imports # Development setup dev-setup: @@ -209,7 +209,7 @@ db-clean: # Advanced testing options test-coverage: @echo "📊 Running tests with coverage..." - $(PYTEST) --cov=router --cov-report=html --cov-report=term + $(PYTEST) --cov=routstr --cov-report=html --cov-report=term @echo "Coverage report generated in htmlcov/" test-watch: @@ -228,7 +228,7 @@ ci-test: ci-lint: @echo "🤖 Running CI linting..." $(RUFF) check . --exit-non-zero-on-fix - $(MYPY) router/ --ignore-missing-imports --no-error-summary + $(MYPY) routstr/ --ignore-missing-imports --no-error-summary # Debug helpers test-debug: diff --git a/README.md b/README.md index b762a208..25fc25c2 100644 --- a/README.md +++ b/README.md @@ -69,7 +69,7 @@ cp .env.example .env ### Running Locally ```bash -fastapi run router --host 0.0.0.0 --port 8000 +fastapi run routstr --host 0.0.0.0 --port 8000 ``` The service forwards requests to `UPSTREAM_BASE_URL`. Supply the upstream API key via the `UPSTREAM_API_KEY` environment variable if required. diff --git a/compose.testing.yml b/compose.testing.yml index cb85d55b..36532d91 100644 --- a/compose.testing.yml +++ b/compose.testing.yml @@ -1,9 +1,9 @@ version: '3.8' services: - router: + routstr: build: . - command: ["/.venv/bin/fastapi", "dev", "router", "--host", "0.0.0.0", "--port", "8000"] + command: ["/.venv/bin/fastapi", "dev", "routstr", "--host", "0.0.0.0", "--port", "8000"] ports: - "8000:8000" environment: @@ -27,9 +27,6 @@ services: - "REFUND_PROCESSING_INTERVAL=3600" - "MINIMUM_PAYOUT=1000" - "PAYOUT_INTERVAL=86400" - volumes: - - ./:/app - - ./logs:/app/logs depends_on: - mock-mint - mock-openai @@ -40,8 +37,6 @@ services: restart: unless-stopped ports: - "8088:8080" # host:container - volumes: - - ./relay-data:/usr/src/app/db environment: - LISTEN_ADDR=0.0.0.0 - LISTEN_PORT=8080 diff --git a/compose.yml b/compose.yml index 8b0f67f2..289c30ce 100644 --- a/compose.yml +++ b/compose.yml @@ -1,7 +1,7 @@ version: '3.8' services: - router: + routstr: build: . volumes: - .:/app @@ -21,9 +21,16 @@ services: - tor-data:/var/lib/tor environment: # Format: HS_=:: - - HS_ROUTER=router:8000:80 + - HS_ROUTER=routstr:8000:80 depends_on: - - router + - routstr + + # Legacy service definition to ensure cleanup of old container + router: + image: alpine:latest + command: /bin/true + profiles: + - cleanup volumes: tor-data: diff --git a/migrations/env.py b/migrations/env.py index e62499fa..b01e7eb1 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -1,23 +1,26 @@ import asyncio import pathlib import sys -from logging.config import fileConfig +# from logging.config import fileConfig from alembic import context from sqlalchemy import pool from sqlalchemy.engine import Connection from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel import SQLModel -# Add the parent directory to the Python path so we can import router modules -sys.path.insert(0, str(pathlib.Path(__file__).resolve().parents[1])) +from routstr.core.db import DATABASE_URL -from router.core.db import DATABASE_URL +# Add the parent directory to the Python path so we can import routstr modules +sys.path.insert(0, str(pathlib.Path(__file__).resolve().parents[1])) config = context.config if config.config_file_name is None: raise ValueError("config_file_name is None") -fileConfig(config.config_file_name) + +# Skip loading alembic's logging configuration to preserve our custom logging +# fileConfig(config.config_file_name) + config.set_main_option("sqlalchemy.url", DATABASE_URL) target_metadata = SQLModel.metadata diff --git a/migrations/versions/042f6b77d69d_introduce_reserved_balance.py b/migrations/versions/042f6b77d69d_introduce_reserved_balance.py new file mode 100644 index 00000000..b9d8e3b0 --- /dev/null +++ b/migrations/versions/042f6b77d69d_introduce_reserved_balance.py @@ -0,0 +1,30 @@ +"""introduce reserved balance + +Revision ID: 042f6b77d69d +Revises: 898f00ea481e +Create Date: 2025-08-18 19:03:09.507368 +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "042f6b77d69d" +down_revision = "898f00ea481e" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column( + "api_keys", + sa.Column("reserved_balance", sa.Integer(), nullable=False, server_default="0"), + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column("api_keys", "reserved_balance") + # ### end Alembic commands ### diff --git a/migrations/versions/7bc4e8b02b9d_add_mint_field.py b/migrations/versions/7bc4e8b02b9d_add_mint_field.py index 93469741..fdb05fac 100644 --- a/migrations/versions/7bc4e8b02b9d_add_mint_field.py +++ b/migrations/versions/7bc4e8b02b9d_add_mint_field.py @@ -6,8 +6,8 @@ Create Date: 2025-08-09 13:48:40.648729 """ import sqlalchemy as sa -import sqlmodel from alembic import op +from sqlmodel.sql import sqltypes # revision identifiers, used by Alembic. revision = "7bc4e8b02b9d" @@ -20,7 +20,7 @@ def upgrade() -> None: # ### commands auto generated by Alembic - please adjust! ### op.add_column( "api_keys", - sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("mint_url", sqltypes.AutoString(), nullable=True), ) # ### end Alembic commands ### diff --git a/migrations/versions/898f00ea481e_add_mint_currency_refund_details.py b/migrations/versions/898f00ea481e_add_mint_currency_refund_details.py new file mode 100644 index 00000000..b9ab2d40 --- /dev/null +++ b/migrations/versions/898f00ea481e_add_mint_currency_refund_details.py @@ -0,0 +1,38 @@ +"""add mint+currency refund details + +Revision ID: 898f00ea481e +Revises: 7bc4e8b02b9d +Create Date: 2025-08-13 16:45:42.148314 +""" + +import sqlalchemy as sa +from alembic import op +from sqlmodel.sql import sqltypes + +# revision identifiers, used by Alembic. +revision = "898f00ea481e" +down_revision = "7bc4e8b02b9d" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column( + "api_keys", + sa.Column("refund_mint_url", sqltypes.AutoString(), nullable=True), + ) + op.add_column( + "api_keys", + sa.Column("refund_currency", sqltypes.AutoString(), nullable=True), + ) + op.drop_column("api_keys", "mint_url") + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column("api_keys", sa.Column("mint_url", sa.VARCHAR(), nullable=True)) + op.drop_column("api_keys", "refund_currency") + op.drop_column("api_keys", "refund_mint_url") + # ### end Alembic commands ### diff --git a/migrations/versions/f6ce1348e266_init.py b/migrations/versions/f6ce1348e266_init.py index 80222218..58cd8285 100644 --- a/migrations/versions/f6ce1348e266_init.py +++ b/migrations/versions/f6ce1348e266_init.py @@ -6,8 +6,8 @@ Create Date: 2025-08-09 13:28:38.537652 """ import sqlalchemy as sa -import sqlmodel as sqlm from alembic import op +from sqlmodel.sql import sqltypes # revision identifiers, used by Alembic. revision = "f6ce1348e266" @@ -20,9 +20,9 @@ def upgrade() -> None: if "api_keys" not in sa.inspect(op.get_bind()).get_table_names(): op.create_table( "api_keys", - sa.Column("hashed_key", sqlm.sql.sqltypes.AutoString(), nullable=False), + sa.Column("hashed_key", sqltypes.AutoString(), nullable=False), sa.Column("balance", sa.Integer(), nullable=False), - sa.Column("refund_address", sqlm.sql.sqltypes.AutoString(), nullable=True), + sa.Column("refund_address", sqltypes.AutoString(), nullable=True), sa.Column("key_expiry_time", sa.Integer(), nullable=True), sa.Column("total_spent", sa.Integer(), nullable=False), sa.Column("total_requests", sa.Integer(), nullable=False), diff --git a/pyproject.toml b/pyproject.toml index 7d25a6c2..9e7f7fca 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "routstr" -version = "0.1.0" +version = "0.1.1b" description = "Payment proxy for your LLM endpoint using cashu and nostr." readme = "README.md" requires-python = ">=3.11" @@ -58,6 +58,13 @@ markers = [ "performance: marks tests that measure performance metrics", ] +[build-system] +requires = ["setuptools", "wheel"] +build-backend = "setuptools.build_meta" + +[tool.setuptools] +packages = ["routstr"] + [tool.ruff.lint] select = ["E", "F", "I"] ignore = ["E501"] @@ -72,8 +79,5 @@ disallow_incomplete_defs = true disallow_untyped_decorators = true [tool.uv.sources] -secp256k1 = { git = "https://github.com/saschanaz/secp256k1-py", branch = "upgrade060" } routstr = { workspace = true } - -[tool.uv.workspace] -members = ["."] +secp256k1 = { git = "https://github.com/saschanaz/secp256k1-py", branch = "upgrade060" } diff --git a/router/core/admin.py b/router/core/admin.py deleted file mode 100644 index b511e449..00000000 --- a/router/core/admin.py +++ /dev/null @@ -1,405 +0,0 @@ -import os -from datetime import datetime, timezone - -from fastapi import APIRouter, HTTPException, Request -from fastapi.responses import HTMLResponse -from pydantic import BaseModel -from sqlmodel import select - -from ..wallet import get_balance, send_token -from .db import ApiKey, create_session - -admin_router = APIRouter(prefix="/admin", include_in_schema=False) - - -class WithdrawRequest(BaseModel): - amount: int - - -def login_form() -> str: - return """ - - - - - - -
- - -
- - - """ - - -def info(content: str) -> str: - return f""" - - - - - -
- {content} -
- - - """ - - -def admin_auth() -> str: - if os.getenv("ADMIN_PASSWORD", "") == "": - return info("Please set a secure ADMIN_PASSWORD= in your ENV variables.") - else: - return login_form() - - -async def dashboard(request: Request) -> str: - # fetch cashu / api-key data from database - async with create_session() as session: - result = await session.exec(select(ApiKey)) - api_keys = result.all() - - api_keys_table_rows = [] - for key in api_keys: - expiry_time_utc = ( - datetime.fromtimestamp(key.key_expiry_time, tz=timezone.utc) - if key.key_expiry_time is not None - else None - ) - expiry_time_human_readable = ( - expiry_time_utc.strftime("%Y-%m-%d %H:%M:%S") if expiry_time_utc else "" - ) - - api_keys_table_rows.append( - f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}" - ) - - # Calculate the total balance of all API keys using integer arithmetic to - # avoid rounding issues. - total_user_balance = sum(key.balance for key in api_keys) // 1000 - # Fetch balance from cashu - current_balance = await get_balance("sat") - owner_balance = current_balance - total_user_balance - - return f""" - - - - - - -

Admin Dashboard

-

Current Cashu Balance

-

Your Balance: {owner_balance} sats

-

The balance is calculated by subtracting the combined user balance from the total Cashu wallet balance.

-

Total Cashu Balance: {current_balance} sats

-

User Balance: {total_user_balance} sats

- - - - - - -
- Withdrawal Token: -
- -

Save this token! It represents your withdrawn balance.

-
- -

User's API Keys

- - - - - - - - - - {"".join(api_keys_table_rows)} -
Hashed KeyBalance (mSats)Total Spent (mSats)Total RequestsRefund AddressRefund Time
- - - """ - - -@admin_router.get("/", response_class=HTMLResponse) -async def admin(request: Request) -> str: - admin_cookie = request.cookies.get("admin_password") - if admin_cookie and admin_cookie == os.getenv("ADMIN_PASSWORD"): - return await dashboard(request) - return admin_auth() - - -@admin_router.post("/withdraw") -async def withdraw( - request: Request, withdraw_request: WithdrawRequest -) -> dict[str, str]: - admin_cookie = request.cookies.get("admin_password") - if not admin_cookie or admin_cookie != os.getenv("ADMIN_PASSWORD"): - raise HTTPException(status_code=403, detail="Unauthorized") - - current_balance = await get_balance("sat") - - if withdraw_request.amount <= 0: - raise HTTPException( - status_code=400, detail="Withdrawal amount must be positive" - ) - - if withdraw_request.amount > current_balance: - raise HTTPException(status_code=400, detail="Insufficient wallet balance") - - token = await send_token(withdraw_request.amount, "sat") - return {"token": token} diff --git a/router/wallet.py b/router/wallet.py deleted file mode 100644 index a19361f6..00000000 --- a/router/wallet.py +++ /dev/null @@ -1,250 +0,0 @@ -import os -from enum import Enum -from typing import Any - -from cashu.core.base import Token -from cashu.wallet.helpers import deserialize_token_from_string -from cashu.wallet.wallet import Wallet - -from .core import db, get_logger - -logger = get_logger(__name__) - - -class CurrencyUnit(Enum): - sat = "sat" - msat = "msat" - -CASHU_MINTS = os.environ.get("CASHU_MINTS", "https://mint.minibits.cash/Bitcoin") -TRUSTED_MINTS = CASHU_MINTS.split(",") -PRIMARY_MINT_URL = TRUSTED_MINTS[0] - - -async def get_balance(unit: CurrencyUnit | str) -> int: - wallet = await Wallet.with_db( - PRIMARY_MINT_URL, - db=".wallet", - load_all_keysets=True, - unit=unit, - ) - await wallet.load_proofs() - return wallet.available_balance.amount - - -async def recieve_token( - token: str, -) -> tuple[int, CurrencyUnit, 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") - - wallet = await Wallet.with_db( - token_obj.mint, - db=".wallet", - load_all_keysets=True, - unit=token_obj.unit, - ) - await wallet.load_mint(token_obj.keysets[0]) - - if token_obj.mint not in TRUSTED_MINTS: - return await swap_to_primary_mint(token_obj, wallet) - - await wallet.redeem(token_obj.proofs) - return token_obj.amount, token_obj.unit, token_obj.mint - - -async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: - """Internal send function - returns amount and serialized token""" - wallet = await Wallet.with_db( - mint_url or PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit=unit - ) - await wallet.load_mint() - await wallet.load_proofs() - proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] - - send_proofs, fees = await wallet.select_to_send( - proofs, amount, set_reserved=True, include_fees=True - ) - token = await wallet.serialize_proofs( - send_proofs, include_dleq=False, legacy=False, memo=None - ) - return amount, token - - -async def send_token( - amount: int, unit: CurrencyUnit | str, mint_url: str | None = None -) -> str: - """Send token and return serialized token string""" - unit_str = unit.value if isinstance(unit, CurrencyUnit) else unit - _, token = await send(amount, unit_str, mint_url) - return token - - -async def swap_to_primary_mint( - token_obj: Token, token_wallet: Wallet -) -> tuple[int, CurrencyUnit, str]: - logger.info( - "swap_to_primary_mint", - extra={ - "mint": token_obj.mint, - "amount": token_obj.amount, - "unit": token_obj.unit, - }, - ) - if token_obj.unit == "sat": - amount_msat = token_obj.amount * 1000 - elif token_obj.unit == "msat": - amount_msat = token_obj.amount - else: - raise ValueError("Invalid unit") - estimated_fee_sat = max(amount_msat // 1000 * 0.01, 2) - amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000 - primary_wallet = await Wallet.with_db( - PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit="sat" - ) - await primary_wallet.load_mint() - - minted_amount = amount_msat_after_fee // 1000 - mint_quote = await primary_wallet.request_mint(minted_amount) - - melt_quote = await token_wallet.melt_quote(mint_quote.request) - _ = await token_wallet.melt( - proofs=token_obj.proofs, - invoice=mint_quote.request, - fee_reserve_sat=melt_quote.fee_reserve, - quote_id=melt_quote.quote, - ) - _ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote) - - return minted_amount, CurrencyUnit.sat, PRIMARY_MINT_URL - - -async def credit_balance( - cashu_token: str, key: db.ApiKey, session: db.AsyncSession -) -> int: - logger.info( - "credit_balance: Starting token redemption", - extra={"token_preview": cashu_token[:50]}, - ) - - try: - amount, unit, mint_url = await recieve_token(cashu_token) - logger.info( - "credit_balance: Token redeemed successfully", - extra={"amount": amount, "unit": unit, "mint_url": mint_url}, - ) - - if unit == "sat": - amount = amount * 1000 - logger.info( - "credit_balance: Converted to msat", extra={"amount_msat": amount} - ) - - if mint_url != PRIMARY_MINT_URL: - logger.error( - "credit_balance: Mint URL mismatch", - extra={"mint_url": mint_url, "primary_mint": PRIMARY_MINT_URL}, - ) - raise ValueError("Mint URL is not supported by this proxy") - - logger.info( - "credit_balance: Updating balance", - extra={"old_balance": key.balance, "credit_amount": amount}, - ) - key.balance += amount - session.add(key) - await session.commit() - logger.info( - "credit_balance: Balance updated successfully", - extra={"new_balance": key.balance}, - ) - - logger.info( - "Cashu token successfully redeemed and stored", - extra={"amount": amount, "unit": unit, "mint_url": mint_url}, - ) - return amount - except Exception as e: - logger.error( - "credit_balance: Error during token redemption", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - raise - - -async def send_to_lnurl(amount: int, unit: CurrencyUnit, lnurl: str) -> dict[str, Any]: - """Send payment to Lightning Address/LNURL""" - try: - # Create wallet instance for this operation - payment_wallet = await Wallet.with_db( - PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit=unit - ) - await payment_wallet.load_mint() - - # Convert amount to correct unit - if unit == CurrencyUnit.sat and amount < 1000: - # Convert sats to msats for small amounts - amount_to_send = amount * 1000 - send_unit = CurrencyUnit.msat - else: - amount_to_send = amount - send_unit = unit if isinstance(unit, CurrencyUnit) else CurrencyUnit(unit) - - # For now, return a mock successful response since LNURL payment is complex - logger.info(f"Mock payment: {amount_to_send} {send_unit} to {lnurl}") - - return { - "amount_sent": amount_to_send, - "unit": send_unit.name, - "lnurl": lnurl, - "status": "completed" - } - - except Exception as e: - logger.error(f"Failed to send to LNURL {lnurl}: {e}") - unit_str = unit.value if isinstance(unit, CurrencyUnit) else unit - return { - "amount_sent": 0, - "unit": unit_str, - "lnurl": lnurl, - "status": "failed", - "error": str(e) - } - - -async def periodic_payout() -> None: - logger.warning("periodic_payout, temporary not implemented") - - -# class Proof: -# """ -# Represents an ecash bill -# """ - - -# def redeem_to_proofs(self, token: str) -> list[Proof]: -# raise NotImplementedError - - -# 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/router/__init__.py b/routstr/__init__.py similarity index 100% rename from router/__init__.py rename to routstr/__init__.py diff --git a/router/auth.py b/routstr/auth.py similarity index 87% rename from router/auth.py rename to routstr/auth.py index 3778ad2b..a1cba59f 100644 --- a/router/auth.py +++ b/routstr/auth.py @@ -1,4 +1,5 @@ import hashlib +import math from typing import Optional from fastapi import HTTPException @@ -12,8 +13,12 @@ from .payment.cost_caculation import ( MaxCostData, calculate_cost, ) -from .payment.helpers import get_max_cost_for_model -from .wallet import credit_balance +from .wallet import ( + PRIMARY_MINT_URL, + TRUSTED_MINTS, + credit_balance, + deserialize_token_from_string, +) logger = get_logger(__name__) @@ -113,6 +118,7 @@ async def validate_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] + "..."} ) @@ -159,12 +165,20 @@ async def validate_bearer_key( "has_expiry_time": bool(key_expiry_time), }, ) + if token_obj.mint in TRUSTED_MINTS: + refund_currency = token_obj.unit + refund_mint_url = token_obj.mint + else: + refund_currency = "sat" + refund_mint_url = PRIMARY_MINT_URL 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) await session.flush() @@ -257,10 +271,10 @@ async def validate_bearer_key( ) -async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int: +async def pay_for_request( + key: ApiKey, cost_per_request: int, session: AsyncSession +) -> int: """Process payment for a request.""" - model = body["model"] - cost_per_request = get_max_cost_for_model(model=model) logger.info( "Processing payment for request", @@ -268,20 +282,19 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int "key_hash": key.hashed_key[:8] + "...", "current_balance": key.balance, "required_cost": cost_per_request, - "model": model, "sufficient_balance": key.balance >= cost_per_request, }, ) - if key.balance < cost_per_request: + 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.balance, - "model": model, + "shortfall": cost_per_request - key.total_balance, }, ) @@ -289,7 +302,7 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", "type": "insufficient_quota", "code": "insufficient_balance", } @@ -311,8 +324,7 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int .where(col(ApiKey.hashed_key) == key.hashed_key) .where(col(ApiKey.balance) >= cost_per_request) .values( - balance=col(ApiKey.balance) - cost_per_request, - total_spent=col(ApiKey.total_spent) + cost_per_request, + reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, total_requests=col(ApiKey.total_requests) + 1, ) ) @@ -351,7 +363,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int "new_balance": key.balance, "total_spent": key.total_spent, "total_requests": key.total_requests, - "model": model, }, ) @@ -365,8 +376,7 @@ async def revert_pay_for_request( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) .values( - balance=col(ApiKey.balance) + cost_per_request, - total_spent=col(ApiKey.total_spent) - cost_per_request, + reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, ) ) @@ -374,6 +384,14 @@ async def revert_pay_for_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={ @@ -424,6 +442,7 @@ async def adjust_payment_for_tokens( # 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", @@ -446,6 +465,7 @@ async def adjust_payment_for_tokens( await session.commit() return cost.dict() + # this should never happen why do we handle this??? if cost_difference > 0: # Need to charge more logger.info( @@ -459,6 +479,7 @@ async def adjust_payment_for_tokens( }, ) + # this should never happen why do we handle this??? if key.balance < cost_difference: logger.warning( "Insufficient balance for token-based pricing adjustment", @@ -472,6 +493,7 @@ async def adjust_payment_for_tokens( ) await session.commit() else: + # this should never happen why do we handle this??? charge_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) @@ -524,13 +546,30 @@ async def adjust_payment_for_tokens( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) .values( - balance=col(ApiKey.balance) + refund, - total_spent=col(ApiKey.total_spent) - refund, + 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(refund_stmt) # type: ignore[call-overload] + result = await session.exec(refund_stmt) # type: ignore[call-overload] await session.commit() - cost.total_msats = deducted_max_cost - refund + + 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( diff --git a/router/balance.py b/routstr/balance.py similarity index 63% rename from router/balance.py rename to routstr/balance.py index 27709269..99eb19a3 100644 --- a/router/balance.py +++ b/routstr/balance.py @@ -1,10 +1,11 @@ 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 .wallet import CurrencyUnit, credit_balance, send_to_lnurl, send_token +from .wallet import PRIMARY_MINT_URL, credit_balance, send_to_lnurl, send_token router = APIRouter() balance_router = APIRouter(prefix="/v1/balance") @@ -32,6 +33,17 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: } +@router.get("/create") +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, + } + + @router.get("/info") async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: return { @@ -40,12 +52,22 @@ async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: } +class TopupRequest(BaseModel): + cashu_token: str + + @router.post("/topup") async def topup_wallet_endpoint( - cashu_token: str, + cashu_token: str | None = None, + topup_request: TopupRequest | None = None, key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: + if topup_request is not None: + cashu_token = topup_request.cashu_token + if cashu_token is None: + raise HTTPException(status_code=400, detail="A cashu_token is required.") + cashu_token = cashu_token.replace("\n", "").replace("\r", "").replace("\t", "") if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") @@ -69,36 +91,52 @@ async def refund_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict: - remaining_balance_msats = key.balance + remaining_balance_msats: int = key.balance - if remaining_balance_msats == 0: + if remaining_balance_msats <= 0: raise HTTPException(status_code=400, detail="No balance to refund") # Perform refund operation first, before modifying balance try: if key.refund_address: - await send_to_lnurl(remaining_balance_msats, CurrencyUnit.msat, key.refund_address) - result = {"recipient": key.refund_address, "msats": remaining_balance_msats} + if key.refund_currency == "sat": + remaining_balance = remaining_balance_msats * 1000 + await send_to_lnurl( + remaining_balance, + key.refund_currency or "sat", + key.refund_mint_url or PRIMARY_MINT_URL, + key.refund_address, + ) + result = {"recipient": key.refund_address} else: - # Convert msats to sats for cashu wallet - remaining_balance_sats = remaining_balance_msats // 1000 - if remaining_balance_sats == 0: - raise HTTPException( - status_code=400, detail="Balance too small to refund (less than 1 sat)" - ) + refund_amount = ( + remaining_balance_msats // 1000 + if key.refund_currency == "sat" + else remaining_balance_msats + ) + refund_currency = key.refund_currency or "sat" + token = await send_token( + refund_amount, refund_currency, key.refund_mint_url + ) + result = {"token": token} - # TODO: choose currency and mint based on what user has configured - token = await send_token(remaining_balance_sats, "sat") + if key.refund_currency == "sat": + result["sats"] = str(remaining_balance_msats // 1000) + else: + result["msats"] = str(remaining_balance_msats) - result = {"msats": remaining_balance_msats, "recipient": None, "token": token} except HTTPException: # Re-raise HTTP exceptions (like 400 for balance too small) raise except Exception as e: # If refund fails, don't modify the database error_msg = str(e) - if ("mint" in error_msg.lower() or "connection" in error_msg.lower() or - isinstance(e, Exception) and "ConnectError" in str(type(e))): + if ( + "mint" in error_msg.lower() + or "connection" in error_msg.lower() + or isinstance(e, Exception) + and "ConnectError" in str(type(e)) + ): raise HTTPException(status_code=503, detail="Mint service unavailable") else: raise HTTPException(status_code=500, detail="Refund failed") diff --git a/router/core/__init__.py b/routstr/core/__init__.py similarity index 100% rename from router/core/__init__.py rename to routstr/core/__init__.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py new file mode 100644 index 00000000..4bf6b00a --- /dev/null +++ b/routstr/core/admin.py @@ -0,0 +1,690 @@ +import json +import os +from datetime import datetime, timezone +from pathlib import Path + +from fastapi import APIRouter, HTTPException, Request +from fastapi.responses import HTMLResponse +from pydantic import BaseModel +from sqlmodel import select + +from ..wallet import ( + TRUSTED_MINTS, + fetch_all_balances, + get_proofs_per_mint_and_unit, + get_wallet, + send_token, + slow_filter_spend_proofs, +) +from .db import ApiKey, create_session +from .logging import get_logger + +logger = get_logger(__name__) + +admin_router = APIRouter(prefix="/admin", include_in_schema=False) + + +class WithdrawRequest(BaseModel): + amount: int + mint_url: str | None = None + unit: str = "sat" + + +def login_form() -> str: + return """ + + + + + + + + + + """ + + +def info(content: str) -> str: + return f""" + + + + + +
+

{content}

+
+ + + """ + + +def admin_auth() -> str: + if os.getenv("ADMIN_PASSWORD", "") == "": + return info("Please set a secure ADMIN_PASSWORD= in your ENV variables.") + else: + return login_form() + + +async def dashboard(request: Request) -> str: + # fetch cashu / api-key data from database + async with create_session() as session: + result = await session.exec(select(ApiKey)) + api_keys = result.all() + + api_keys_table_rows = [] + for key in api_keys: + expiry_time_utc = ( + datetime.fromtimestamp(key.key_expiry_time, tz=timezone.utc) + if key.key_expiry_time is not None + else None + ) + expiry_time_human_readable = ( + expiry_time_utc.strftime("%Y-%m-%d %H:%M:%S") if expiry_time_utc else "" + ) + + api_keys_table_rows.append( + f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}" + ) + + # Fetch all balances using the abstracted function + ( + balance_details, + total_wallet_balance_sats, + total_user_balance_sats, + owner_balance, + ) = await fetch_all_balances() + + return f""" + + + + + + +

Admin Dashboard

+ +
+

Cashu Wallet Balance

+
+ Your Balance (Total) + { + owner_balance + } sats +
+
+ Total Wallet + {total_wallet_balance_sats} sats +
+
+ User Balance + {total_user_balance_sats} sats +
+

Your balance = Total wallet - User balance

+ +
+
+
Mint / Unit
+
Wallet
+
Users
+
Owner
+
+ { + "".join( + [ + f'''
+
{detail["mint_url"].replace("https://", "").replace("http://", "")} • {detail["unit"].upper()}
+
{detail["wallet_balance"] if not detail.get("error") else "error"}
+
{detail["user_balance"] if not detail.get("error") else "-"}
+
0 else ""}">{detail["owner_balance"] if not detail.get("error") else "-"}
+
''' + for detail in balance_details + if detail.get("wallet_balance", 0) > 0 or detail.get("error") + ] + ) + } +
+
+ + + + + + + + + +
+ Withdrawal Token: +
+ +

Save this token! It represents your withdrawn balance.

+
+ +

Temporary Balances

+ + + + + + + + + + {"".join(api_keys_table_rows)} +
Hashed KeyBalance (mSats)Total Spent (mSats)Total RequestsRefund AddressRefund Time
+ + + """ + + +@admin_router.get("/", response_class=HTMLResponse) +async def admin(request: Request) -> str: + admin_cookie = request.cookies.get("admin_password") + if admin_cookie and admin_cookie == os.getenv("ADMIN_PASSWORD"): + return await dashboard(request) + return admin_auth() + + +@admin_router.get("/logs/{request_id}", response_class=HTMLResponse) +async def view_logs(request: Request, request_id: str) -> str: + admin_cookie = request.cookies.get("admin_password") + if not admin_cookie or admin_cookie != os.getenv("ADMIN_PASSWORD"): + return admin_auth() + + logger.info(f"Investigating logs for request_id: {request_id}") + + # Search for log entries with this request_id + log_entries = [] + logs_dir = Path("logs") + + if logs_dir.exists(): + # Get all log files sorted by modification time (most recent first) + log_files = sorted( + logs_dir.glob("*.log"), key=lambda x: x.stat().st_mtime, reverse=True + ) + + for log_file in log_files[:7]: # Check last 7 days of logs + try: + with open(log_file, "r") as f: + for line in f: + if request_id in line: + try: + # Parse JSON log entry + log_data = json.loads(line.strip()) + log_entries.append(log_data) + except json.JSONDecodeError: + # If not JSON, include raw line + log_entries.append({"raw": line.strip()}) + except Exception as e: + logger.error(f"Error reading log file {log_file}: {e}") + + # Sort entries by timestamp if available + log_entries.sort(key=lambda x: x.get("asctime", ""), reverse=False) + + # Format log entries for display + formatted_logs = [] + for entry in log_entries: + if "raw" in entry: + formatted_logs.append(f'
{entry["raw"]}
') + else: + # Format JSON log entry + timestamp = entry.get("asctime", "Unknown time") + level = entry.get("levelname", "INFO") + message = entry.get("message", "") + pathname = entry.get("pathname", "") + lineno = entry.get("lineno", "") + + # Extract additional fields + extra_fields = { + k: v + for k, v in entry.items() + if k + not in [ + "asctime", + "levelname", + "message", + "pathname", + "lineno", + "name", + "version", + "request_id", + ] + } + + level_class = level.lower() + formatted_entry = f""" +
+
+ {timestamp} + [{level}] + {pathname}:{lineno} +
+
{message}
+ """ + + if extra_fields: + formatted_entry += '
' + for key, value in extra_fields.items(): + formatted_entry += f'
{key}: {json.dumps(value) if isinstance(value, (dict, list)) else value}
' + formatted_entry += "
" + + formatted_entry += "
" + formatted_logs.append(formatted_entry) + + return f""" + + + + + + ← Back to Dashboard +

Log Investigation

+
+ Request ID: {request_id} +
+
+ {"".join(formatted_logs) if formatted_logs else '
No log entries found for this Request ID
'} +
+

+ Found {len(log_entries)} log entries • Searched last 7 days of logs +

+ + + """ + + +@admin_router.post("/withdraw") +async def withdraw( + request: Request, withdraw_request: WithdrawRequest +) -> dict[str, str]: + admin_cookie = request.cookies.get("admin_password") + if not admin_cookie or admin_cookie != os.getenv("ADMIN_PASSWORD"): + raise HTTPException(status_code=403, detail="Unauthorized") + + # Get wallet and check balance + wallet = await get_wallet( + withdraw_request.mint_url or TRUSTED_MINTS[0], withdraw_request.unit + ) + proofs = get_proofs_per_mint_and_unit( + wallet, + withdraw_request.mint_url or TRUSTED_MINTS[0], + withdraw_request.unit, + not_reserved=True, + ) + proofs = await slow_filter_spend_proofs(proofs, wallet) + current_balance = sum(proof.amount for proof in proofs) + + if withdraw_request.amount <= 0: + raise HTTPException( + status_code=400, detail="Withdrawal amount must be positive" + ) + + if withdraw_request.amount > current_balance: + raise HTTPException(status_code=400, detail="Insufficient wallet balance") + + token = await send_token( + withdraw_request.amount, withdraw_request.unit, withdraw_request.mint_url + ) + return {"token": token} diff --git a/router/core/db.py b/routstr/core/db.py similarity index 79% rename from router/core/db.py rename to routstr/core/db.py index ac94b291..46696f02 100644 --- a/router/core/db.py +++ b/routstr/core/db.py @@ -5,7 +5,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlmodel import Field, SQLModel +from sqlmodel import Field, SQLModel, func, select from sqlmodel.ext.asyncio.session import AsyncSession from .logging import get_logger @@ -23,6 +23,9 @@ class ApiKey(SQLModel, table=True): # type: ignore 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)" + ) refund_address: str | None = Field( default=None, description="Lightning address to refund remaining balance after key expires", @@ -35,10 +38,28 @@ class ApiKey(SQLModel, table=True): # type: ignore default=0, description="Total spent in millisatoshis (msats)" ) total_requests: int = Field(default=0) - mint_url: str | None = Field( + refund_mint_url: str | None = Field( default=None, description="URL of the mint used to create the cashu-token", ) + refund_currency: str | None = Field( + default=None, + description="Currency of the cashu-token", + ) + + @property + def total_balance(self) -> int: + return self.balance - self.reserved_balance + + +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 + ) + result = await db_session.exec(query) + return result.one() or 0 async def init_db() -> None: diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py new file mode 100644 index 00000000..e74d3e05 --- /dev/null +++ b/routstr/core/exceptions.py @@ -0,0 +1,57 @@ +from fastapi import Request +from fastapi.responses import JSONResponse + +from .logging import get_logger + +logger = get_logger(__name__) + + +async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse: + """Handle HTTP exceptions and include request ID in response.""" + request_id = getattr(request.state, "request_id", "unknown") + + # Get status code and detail - works for both FastAPI and Starlette HTTPException + status_code = getattr(exc, "status_code", 500) + detail = getattr(exc, "detail", str(exc)) + + logger.warning( + "HTTP exception", + extra={ + "request_id": request_id, + "status_code": status_code, + "detail": detail, + "path": request.url.path, + }, + ) + + return JSONResponse( + status_code=status_code, + content={ + "detail": detail, + "request_id": request_id, + }, + ) + + +async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse: + """Handle general exceptions and include request ID in response.""" + request_id = getattr(request.state, "request_id", "unknown") + + logger.error( + "Unhandled exception", + extra={ + "request_id": request_id, + "error": str(exc), + "error_type": type(exc).__name__, + "path": request.url.path, + }, + exc_info=True, + ) + + return JSONResponse( + status_code=500, + content={ + "detail": "Internal server error, please contact support with the request ID.", + "request_id": request_id, + }, + ) diff --git a/router/core/logging.py b/routstr/core/logging.py similarity index 84% rename from router/core/logging.py rename to routstr/core/logging.py index 9820b747..00a8ad27 100644 --- a/router/core/logging.py +++ b/routstr/core/logging.py @@ -89,7 +89,7 @@ def get_package_version() -> str: return version current_path = current_path.parent - # Fallback: try the simple path resolution (3 levels up for router/logging/logging_config.py) + # Fallback: try the simple path resolution (3 levels up for routstr/logging/logging_config.py) pyproject_path = Path(__file__).parent.parent.parent / "pyproject.toml" if pyproject_path.exists(): with open(pyproject_path, "rb") as f: @@ -115,6 +115,23 @@ class VersionFilter(logging.Filter): return True +class RequestIdFilter(logging.Filter): + """Filter to add request ID to all log records.""" + + def filter(self, record: logging.LogRecord) -> bool: + """Add request ID to the log record if available.""" + try: + # Import here to avoid circular imports + from .middleware import request_id_context + + request_id = request_id_context.get(None) + record.request_id = request_id if request_id else "no-request-id" + except ImportError: + # If middleware isn't available yet, just use default + record.request_id = "no-request-id" + return True + + class SecurityFilter(logging.Filter): """Filter to remove sensitive information from logs.""" @@ -198,12 +215,13 @@ def setup_logging() -> None: "formatters": { "json": { "()": jsonlogger.JsonFormatter, - "format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s", + "format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s", "datefmt": "%Y-%m-%d %H:%M:%S", }, }, "filters": { "version_filter": {"()": VersionFilter}, + "request_id_filter": {"()": RequestIdFilter}, "security_filter": {"()": SecurityFilter}, }, "handlers": { @@ -214,7 +232,7 @@ def setup_logging() -> None: "show_path": False, "rich_tracebacks": True, "markup": True, - "filters": ["security_filter"], + "filters": ["request_id_filter", "security_filter"], }, "file": { "()": DailyRotatingFileHandler, @@ -225,35 +243,45 @@ def setup_logging() -> None: "interval": 1, # Every 1 day "backupCount": 30, # Keep 30 days of logs "atTime": None, # Rotate at midnight (00:00) - "filters": ["version_filter", "security_filter"], + "filters": ["version_filter", "request_id_filter", "security_filter"], }, }, "loggers": { - "router": { + "routstr": { "level": log_level, "handlers": handlers, "propagate": False, }, - "router.payment": { + "routstr.payment": { "level": log_level, "handlers": handlers, "propagate": False, }, - "router.cashu": { + "routstr.proxy": { "level": log_level, "handlers": handlers, "propagate": False, }, - "router.proxy": { + "routstr.auth": { "level": log_level, "handlers": handlers, "propagate": False, }, - "router.auth": { + "routstr.payment.models": { "level": log_level, "handlers": handlers, "propagate": False, }, + "routstr.core.exceptions": { + "level": log_level, + "handlers": handlers, + "propagate": False, + }, + "routstr.core.middleware": { + "level": log_level, + "handlers": ["file"], + "propagate": False, + }, # Suppress verbose third-party logging "httpx": { "level": "WARNING", @@ -266,13 +294,13 @@ def setup_logging() -> None: "propagate": False, }, "uvicorn.access": { - "level": "WARNING", - "handlers": ["console"] if console_enabled else [], + "level": log_level, # Use the configured log level instead of WARNING + "handlers": handlers, # Use both console and file handlers "propagate": False, }, "uvicorn.error": { - "level": "INFO", - "handlers": ["console"], + "level": log_level, # Use the configured log level + "handlers": handlers, # Use both console and file handlers "propagate": False, }, "watchfiles.main": {"level": "WARNING", "handlers": [], "propagate": False}, diff --git a/router/core/main.py b/routstr/core/main.py similarity index 85% rename from router/core/main.py rename to routstr/core/main.py index 2e0dce68..b8a19bbd 100644 --- a/router/core/main.py +++ b/routstr/core/main.py @@ -5,6 +5,8 @@ from typing import AsyncGenerator from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import RedirectResponse +from starlette.exceptions import HTTPException from ..balance import balance_router, deprecated_wallet_router from ..discovery import providers_router @@ -13,13 +15,15 @@ from ..proxy import proxy_router from ..wallet import periodic_payout from .admin import admin_router from .db import init_db, run_migrations +from .exceptions import general_exception_handler, http_exception_handler from .logging import get_logger, setup_logging +from .middleware import LoggingMiddleware # Initialize logging first setup_logging() logger = get_logger(__name__) -__version__ = "0.1.0" +__version__ = "0.1.1b" @asynccontextmanager @@ -91,8 +95,16 @@ app.add_middleware( allow_credentials=True, allow_methods=["*"], allow_headers=["*"], + expose_headers=["x-routstr-request-id"], ) +# Add logging middleware +app.add_middleware(LoggingMiddleware) + +# Add exception handlers +app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore +app.add_exception_handler(Exception, general_exception_handler) + @app.get("/", include_in_schema=False) @app.get("/v1/info") @@ -109,6 +121,11 @@ async def info() -> dict: } +@app.get("/admin") +async def admin_redirect() -> RedirectResponse: + return RedirectResponse("/admin/") + + app.include_router(models_router) app.include_router(admin_router) app.include_router(balance_router) diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py new file mode 100644 index 00000000..4307cb9d --- /dev/null +++ b/routstr/core/middleware.py @@ -0,0 +1,126 @@ +import time +import uuid +from contextvars import ContextVar +from typing import Callable + +from fastapi import Request, Response +from starlette.middleware.base import BaseHTTPMiddleware + +from .logging import get_logger + +logger = get_logger(__name__) + +# Context variable to store request ID across async context +request_id_context: ContextVar[str | None] = ContextVar("request_id") + + +class LoggingMiddleware(BaseHTTPMiddleware): + """Middleware to log detailed request and response information.""" + + async def dispatch(self, request: Request, call_next: Callable) -> Response: + # Generate request ID + request_id = str(uuid.uuid4()) + request.state.request_id = request_id + + # Set request ID in context for logging + token = request_id_context.set(request_id) + + # Start timing + start_time = time.time() + + # Log request details + request_body = None + if request.method in ["POST", "PUT", "PATCH"]: + try: + # Only read body for non-streaming requests + if hasattr(request, "_body"): + request_body = await request.body() + except Exception: + pass + + # Extract request info + client_host = None + if request.client: + client_host = request.client.host + + # Log incoming request + logger.info( + "Incoming request", + extra={ + "request_id": request_id, + "method": request.method, + "path": request.url.path, + "query_params": dict(request.query_params), + "client_host": client_host, + "headers": { + k: v + for k, v in request.headers.items() + if k.lower() not in ["authorization", "x-cashu", "cookie"] + }, + "body_size": len(request_body) if request_body else 0, + }, + ) + + # Log at TRACE level for full body (security filter will redact sensitive data) + if request_body and hasattr(logger, "exception"): + logger.exception( + "Request body", + extra={ + "request_id": request_id, + "method": request.method, + "path": request.url.path, + "body": request_body.decode("utf-8", errors="ignore")[ + :1000 + ], # Limit size + }, + ) + + # Process request + try: + response = await call_next(request) + + # Calculate duration + duration = time.time() - start_time + + # Log response + logger.info( + "Request completed", + extra={ + "request_id": request_id, + "method": request.method, + "path": request.url.path, + "status_code": response.status_code, + "duration_ms": round(duration * 1000, 2), + "client_host": client_host, + }, + ) + if hasattr(response, "headers"): + response.headers["x-routstr-request-id"] = request_id + + return response + + except Exception as e: + # Calculate duration + duration = time.time() - start_time + + # Log error + logger.error( + "Request failed", + extra={ + "request_id": request_id, + "method": request.method, + "path": request.url.path, + "duration_ms": round(duration * 1000, 2), + "client_host": client_host, + "error": str(e), + "error_type": type(e).__name__, + }, + exc_info=True, + ) + raise + finally: + # Reset context + request_id_context.reset(token) + + +__all__ = ["LoggingMiddleware", "request_id_context"] diff --git a/router/discovery.py b/routstr/discovery.py similarity index 87% rename from router/discovery.py rename to routstr/discovery.py index 49ead23d..caa77e85 100644 --- a/router/discovery.py +++ b/routstr/discovery.py @@ -9,6 +9,10 @@ import httpx import websockets from fastapi import APIRouter +from .core.logging import get_logger + +logger = get_logger(__name__) + providers_router = APIRouter(prefix="/v1/providers") @@ -44,7 +48,7 @@ async def query_nostr_relay_for_providers( try: async with websockets.connect(relay_url, timeout=timeout) as websocket: - print("Connected to relay, searching for kind 31338 events") + logger.debug("Connected to relay, searching for kind 31338 events") await websocket.send(req_message) while True: @@ -54,27 +58,27 @@ async def query_nostr_relay_for_providers( if data[0] == "EVENT" and data[1] == sub_id: event = data[2] - print(f"Found provider announcement: {event['id']}") + logger.debug(f"Found provider announcement: {event['id']}") events.append(event) elif data[0] == "EOSE" and data[1] == sub_id: - print("Received EOSE message") + logger.debug("Received EOSE message") break elif data[0] == "NOTICE": - print(f"Relay notice: {data[1]}") + logger.warning(f"Relay notice: {data[1]}") except asyncio.TimeoutError: - print("Timeout waiting for message") + logger.debug("Timeout waiting for message") break except json.JSONDecodeError: - print("Failed to decode message as JSON") + logger.warning("Failed to decode message as JSON") continue await websocket.send(json.dumps(["CLOSE", sub_id])) except Exception as e: - print(f"Query failed: {e}") + logger.error(f"Query failed: {e}") - print(f"Query complete. Found {len(events)} provider announcements") + logger.info(f"Query complete. Found {len(events)} provider announcements") return events @@ -103,7 +107,7 @@ def parse_provider_announcement(event: dict[str, Any]) -> dict[str, Any] | None: # Validate required fields if not endpoint_url or not provider_name or not d_tag: - print( + logger.warning( f"Invalid provider announcement - missing required tags: {event['id']}" ) return None @@ -140,7 +144,7 @@ def parse_provider_announcement(event: dict[str, Any]) -> dict[str, Any] | None: } except Exception as e: - print(f"Error parsing provider announcement {event.get('id', 'unknown')}: {e}") + logger.error(f"Error parsing provider announcement {event.get('id', 'unknown')}: {e}") return None @@ -221,7 +225,7 @@ async def get_providers( # Query multiple relays for provider announcements for relay_url in discovery_relays: - print(f"\nQuerying relay for providers: {relay_url}") + logger.info(f"Querying relay for providers: {relay_url}") try: events = await query_nostr_relay_for_providers( relay_url=relay_url, @@ -235,13 +239,13 @@ async def get_providers( event_ids.add(event["id"]) all_events.append(event) - print(f"Got {len(events)} provider announcements from {relay_url}") + logger.info(f"Got {len(events)} provider announcements from {relay_url}") except Exception as e: - print(f"Failed to query {relay_url}: {e}") + logger.error(f"Failed to query {relay_url}: {e}") continue - print(f"Found {len(all_events)} total unique provider announcements") + logger.info(f"Found {len(all_events)} total unique provider announcements") # Parse provider announcements according to RIP-02 providers = [] @@ -250,7 +254,7 @@ async def get_providers( if parsed_provider: providers.append(parsed_provider) - print(f"Parsed {len(providers)} valid provider announcements") + logger.info(f"Parsed {len(providers)} valid provider announcements") # Check provider health if requested healthy_providers: list[dict[str, Any]] = [] diff --git a/router/payment/__init__.py b/routstr/payment/__init__.py similarity index 100% rename from router/payment/__init__.py rename to routstr/payment/__init__.py diff --git a/router/payment/cost_caculation.py b/routstr/payment/cost_caculation.py similarity index 100% rename from router/payment/cost_caculation.py rename to routstr/payment/cost_caculation.py diff --git a/router/payment/helpers.py b/routstr/payment/helpers.py similarity index 81% rename from router/payment/helpers.py rename to routstr/payment/helpers.py index 3e131326..c621a55e 100644 --- a/router/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,8 +1,8 @@ import json import os -from typing import Optional from fastapi import HTTPException, Response +from fastapi.requests import Request from ..core import get_logger from ..wallet import deserialize_token_from_string @@ -19,30 +19,6 @@ if not UPSTREAM_BASE_URL: raise ValueError("Please set the UPSTREAM_BASE_URL environment variable") -def get_cost_per_request(model: str | None = None) -> int: - """Get the cost per request for a given model.""" - logger.debug( - "Calculating cost per request", - extra={ - "model": model, - "model_based_pricing": MODEL_BASED_PRICING, - "has_models": bool(MODELS), - }, - ) - - if MODEL_BASED_PRICING and MODELS and model: - cost = get_max_cost_for_model(model=model) - logger.debug( - "Using model-based cost", extra={"model": model, "cost_msats": cost} - ) - return cost - - logger.debug( - "Using default cost per request", extra={"cost_msats": COST_PER_REQUEST} - ) - return COST_PER_REQUEST - - 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 @@ -111,7 +87,7 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N ) -def get_max_cost_for_model(model: str) -> int: +def get_max_cost_for_model(model: str, tolerance_percentage: int = 1) -> int: """Get the maximum cost for a specific model.""" logger.debug( "Getting max cost for model", @@ -142,7 +118,7 @@ def get_max_cost_for_model(model: str) -> int: for m in MODELS: if m.id == model: - max_cost = m.sats_pricing.max_cost * 1000 # type: ignore + max_cost = m.sats_pricing.max_cost * 1000 * (1 - tolerance_percentage / 100) # type: ignore logger.debug( "Found model-specific max cost", extra={"model": model, "max_cost_msats": max_cost}, @@ -157,21 +133,13 @@ def get_max_cost_for_model(model: str) -> int: def create_error_response( - error_type: str, message: str, status_code: int, token: Optional[str] = None + error_type: str, + message: str, + status_code: int, + request: Request, + token: str | None = None, ) -> Response: """Create a standardized error response.""" - logger.info( - "Creating error response", - extra={ - "error_type": error_type, - "error_message": message, - "status_code": status_code, - }, - ) - - response_headers = {} - if token: - response_headers["X-Cashu"] = token return Response( content=json.dumps( { @@ -179,12 +147,13 @@ def create_error_response( "message": message, "type": error_type, "code": status_code, - } + }, + "request_id": getattr(request.state, "request_id", "unknown"), } ), status_code=status_code, media_type="application/json", - headers=dict(response_headers), + headers={"X-Cashu": token} if token else {}, ) diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py new file mode 100644 index 00000000..49bc8d65 --- /dev/null +++ b/routstr/payment/lnurl.py @@ -0,0 +1,294 @@ +from __future__ import annotations + +import math +from typing import TypedDict + +import httpx +from cashu.wallet.wallet import Proof, Wallet + +try: + from bech32 import bech32_decode, convertbits # type: ignore +except ModuleNotFoundError: # pragma: no cover – allow runtime miss + bech32_decode = None # type: ignore + convertbits = None # type: ignore + + +class LNURLData(TypedDict): + """LNURL payRequest data.""" + + callback_url: str + min_sendable: int # millisatoshi + max_sendable: int # millisatoshi + + +class LNURLError(Exception): + """LNURL related errors.""" + + +def parse_lightning_invoice_amount(invoice: str, currency: str = "sat") -> int: + """Parse Lightning invoice (BOLT-11) to extract amount in specified currency units. + + Args: + invoice: BOLT-11 Lightning invoice string + currency: Target currency unit ("sat" or "msat") + + Returns: + Amount in the specified currency unit + + Raises: + LNURLError: If invoice format is invalid or amount cannot be parsed + """ + invoice = invoice.lower().strip() + + if not invoice.startswith("ln"): + raise LNURLError("Invalid Lightning invoice format") + + # Find the network part (bc, tb, etc.) + network_start = 2 + while network_start < len(invoice) and invoice[network_start] not in "0123456789": + network_start += 1 + + if network_start >= len(invoice): + raise LNURLError("Invalid Lightning invoice format") + + # Parse amount and multiplier + amount_str = "" + multiplier = "" + i = network_start + + # Extract numeric part + while i < len(invoice) and invoice[i].isdigit(): + amount_str += invoice[i] + i += 1 + + # Extract multiplier if present + if i < len(invoice) and invoice[i] in "munp": + multiplier = invoice[i] + i += 1 + + # Check if we have the required "1" separator + if i >= len(invoice) or invoice[i] != "1": + raise LNURLError("Invalid Lightning invoice format") + + if not amount_str: + raise LNURLError("Lightning invoice amount not specified") + + # Convert to base units + try: + amount = int(amount_str) + except ValueError: + raise LNURLError("Invalid Lightning invoice amount") + + # Apply multiplier to get millisatoshis + if multiplier == "m": # milli = 10^-3 + amount_msat = amount * 100_000_000 # amount is in BTC * 10^-3 + elif multiplier == "u": # micro = 10^-6 + amount_msat = amount * 100_000 # amount is in BTC * 10^-6 + elif multiplier == "n": # nano = 10^-9 + amount_msat = amount * 100 # amount is in BTC * 10^-9 + elif multiplier == "p": # pico = 10^-12 + amount_msat = amount // 10 # amount is in BTC * 10^-12 + else: + # No multiplier means the amount is in BTC + amount_msat = amount * 100_000_000_000 # Convert BTC to msat + + # Convert to target currency unit + if currency == "msat": + return amount_msat + elif currency == "sat": + return amount_msat // 1000 + else: + raise LNURLError(f"Unsupported currency for Lightning: {currency}") + + +async def decode_lnurl(lnurl: str) -> str: + """Decode LNURL to get the actual URL. + + Handles: + - lightning: prefix + - user@host format + - bech32 encoded lnurl + - direct HTTPS URLs + + Args: + lnurl: LNURL string in any supported format + + Returns: + The decoded HTTPS URL + + Raises: + LNURLError: If the LNURL format is invalid + """ + # Remove lightning: prefix if present + if lnurl.startswith("lightning:"): + lnurl = lnurl[10:] + + # Handle user@host format (Lightning Address) + if "@" in lnurl and len(lnurl.split("@")) == 2: + user, host = lnurl.split("@") + return f"https://{host}/.well-known/lnurlp/{user}" + + # Handle bech32 encoded LNURL + if lnurl.lower().startswith("lnurl"): + if bech32_decode is None or convertbits is None: + raise ImportError( + "bech32 library is required for LNURL bech32 decoding. " + "Install it with: pip install bech32" + ) + + try: + hrp, data = bech32_decode(lnurl) + if data is None: + raise LNURLError("Invalid bech32 data in LNURL") + + decoded_data = convertbits(data, 5, 8, False) + if decoded_data is None: + raise LNURLError("Failed to convert LNURL bits") + + return bytes(decoded_data).decode("utf-8") + except Exception as e: + raise LNURLError(f"Failed to decode LNURL: {e}") from e + + # Assume it's a direct URL + if not lnurl.startswith("https://"): + raise LNURLError("Direct LNURL must use HTTPS") + + return lnurl + + +async def get_lnurl_data(lnurl: str) -> LNURLData: + """Fetch LNURL payRequest data. + + Args: + lnurl: LNURL string in any supported format + + Returns: + LNURLData with callback URL and sendable amounts + + Raises: + LNURLError: If the LNURL data is invalid + httpx.HTTPError: If the HTTP request fails + """ + url = await decode_lnurl(lnurl) + + async with httpx.AsyncClient() as client: + response = await client.get(url, follow_redirects=True, timeout=10) + response.raise_for_status() + + lnurl_data = response.json() + + # Validate payRequest data + if lnurl_data.get("tag") != "payRequest": + raise LNURLError( + f"Invalid LNURL tag: expected 'payRequest', got '{lnurl_data.get('tag')}'" + ) + + if not isinstance(lnurl_data.get("callback"), str): + raise LNURLError("Invalid LNURL payRequest: missing callback URL") + + return LNURLData( + callback_url=lnurl_data["callback"], + min_sendable=lnurl_data.get("minSendable", 1000), # Default 1 sat + max_sendable=lnurl_data.get("maxSendable", 1000000000), # Default 1000 BTC + ) + + +async def get_lnurl_invoice( + callback_url: str, amount_msat: int +) -> tuple[str, dict[str, object]]: + """Request a Lightning invoice from LNURL callback. + + Args: + callback_url: The LNURL callback URL + amount_msat: Amount in millisatoshi + + Returns: + Tuple of (bolt11_invoice, full_response_data) + + Raises: + LNURLError: If the response is invalid + httpx.HTTPError: If the HTTP request fails + """ + async with httpx.AsyncClient() as client: + response = await client.get( + callback_url, + params={"amount": amount_msat}, + follow_redirects=True, + timeout=10, + ) + response.raise_for_status() + + invoice_data = response.json() + + if "pr" not in invoice_data: + # Check if there's an error in the response + if "reason" in invoice_data: + raise LNURLError(f"LNURL error: {invoice_data['reason']}") + raise LNURLError(f"Invalid LNURL invoice response: {invoice_data}") + + return invoice_data["pr"], invoice_data + + +async def raw_send_to_lnurl( + wallet: Wallet, proofs: list[Proof], lnurl: str, unit: str +) -> int: + """Send funds to an LNURL address. + + Args: + wallet: Wallet instance + lnurl: LNURL string (can be lightning:, user@host, bech32, or direct URL) + amount: Amount to send in the specified currency unit + + Returns: + Amount actually paid in the specified currency unit + + Raises: + WalletError: If amount is outside LNURL limits or insufficient balance + LNURLError: If LNURL operations fail + + Example: + # Send 1000 sats to a Lightning Address + paid = await wallet.send_to_lnurl("user@getalby.com", 1000) + print(f"Paid {paid} sats") + + # Send USD to Lightning Address + paid = await wallet.send_to_lnurl("user@getalby.com", 50, unit="usd") + """ + total_balance = sum(proof.amount for proof in proofs) + lnurl_data = await get_lnurl_data(lnurl) + + if unit == "sat": + amount_msat = total_balance * 1000 + min_sendable_sat = lnurl_data["min_sendable"] // 1000 + max_sendable_sat = lnurl_data["max_sendable"] // 1000 + elif unit == "msat": + amount_msat = (total_balance // 1000) * 1000 + min_sendable_sat = lnurl_data["min_sendable"] + max_sendable_sat = lnurl_data["max_sendable"] + else: + raise ValueError(f"Currency {unit} not supported for LNURL") + + if not (lnurl_data["min_sendable"] <= amount_msat <= lnurl_data["max_sendable"]): + raise ValueError( + f"Amount {total_balance} {unit} is outside LNURL limits " + f"({min_sendable_sat} - {max_sendable_sat} {unit})" + ) + + estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + estimated_fees_msat = estimated_fees_sat * 1000 + final_amount = amount_msat - estimated_fees_msat + + bolt11_invoice, _ = await get_lnurl_invoice( + lnurl_data["callback_url"], final_amount + ) + + melt_quote_resp = await wallet.melt_quote( + invoice=bolt11_invoice, amount_msat=final_amount + ) + _ = await wallet.melt( + proofs=proofs, + invoice=bolt11_invoice, + fee_reserve_sat=melt_quote_resp.fee_reserve, + quote_id=melt_quote_resp.quote, + ) + return final_amount diff --git a/router/payment/models.py b/routstr/payment/models.py similarity index 90% rename from router/payment/models.py rename to routstr/payment/models.py index 96f59bfb..5849a489 100644 --- a/router/payment/models.py +++ b/routstr/payment/models.py @@ -7,8 +7,11 @@ from urllib.request import urlopen from fastapi import APIRouter from pydantic.v1 import BaseModel +from ..core.logging import get_logger from .price import sats_usd_ask_price +logger = get_logger(__name__) + models_router = APIRouter() @@ -84,7 +87,7 @@ def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: return models_data except Exception as e: - print(f"Error fetching models from OpenRouter API: {e}") + logger.error(f"Error fetching models from OpenRouter API: {e}") return [] @@ -101,26 +104,26 @@ def load_models() -> list[Model]: # Check if user has actively provided a models.json file if models_path.exists(): - print(f"Loading models from user-provided file: {models_path}") + 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", [])] except Exception as e: - print(f"Error loading models from {models_path}: {e}") + logger.error(f"Error loading models from {models_path}: {e}") # Fall through to auto-generation # Auto-generate models from OpenRouter API - print("Auto-generating models from OpenRouter API") + logger.info("Auto-generating models from OpenRouter API") source_filter = os.getenv("SOURCE") 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: - print("Failed to fetch models from OpenRouter API") + logger.error("Failed to fetch models from OpenRouter API") return [] - print(f"Successfully fetched {len(models_data)} models from OpenRouter API") + logger.info(f"Successfully fetched {len(models_data)} models from OpenRouter API") return [Model(**model) for model in models_data] @@ -165,7 +168,7 @@ async def update_sats_pricing() -> None: except asyncio.CancelledError: break except Exception as e: - print("Error updating sats pricing: ", e) + logger.error(f"Error updating sats pricing: {e}") try: await asyncio.sleep(10) except asyncio.CancelledError: @@ -173,6 +176,6 @@ async def update_sats_pricing() -> None: @models_router.get("/v1/models") -@models_router.get("/models") +@models_router.get("/models", include_in_schema=False) async def models() -> dict: return {"data": MODELS} diff --git a/router/payment/price.py b/routstr/payment/price.py similarity index 100% rename from router/payment/price.py rename to routstr/payment/price.py diff --git a/router/payment/x_cashu.py b/routstr/payment/x_cashu.py similarity index 90% rename from router/payment/x_cashu.py rename to routstr/payment/x_cashu.py index afbc0d7e..551dce4a 100644 --- a/router/payment/x_cashu.py +++ b/routstr/payment/x_cashu.py @@ -7,20 +7,15 @@ from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from ..core import get_logger -from ..wallet import CurrencyUnit, recieve_token, send_token +from ..wallet import recieve_token, send_token from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost -from .helpers import ( - UPSTREAM_BASE_URL, - create_error_response, - get_max_cost_for_model, - prepare_upstream_headers, -) +from .helpers import UPSTREAM_BASE_URL, create_error_response, prepare_upstream_headers logger = get_logger(__name__) async def x_cashu_handler( - request: Request, x_cashu_token: str, path: str + request: Request, x_cashu_token: str, path: str, max_cost_for_model: int ) -> Response | StreamingResponse: """Handle X-Cashu token payment requests.""" logger.info( @@ -44,7 +39,9 @@ async def x_cashu_handler( extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, ) - return await forward_to_upstream(request, path, headers, amount, unit) + return await forward_to_upstream( + request, path, headers, amount, unit, max_cost_for_model + ) except Exception as e: error_message = str(e) logger.error( @@ -63,7 +60,8 @@ async def x_cashu_handler( "token_already_spent", "The provided CASHU token has already been spent", 400, - x_cashu_token, + request=request, + token=x_cashu_token, ) if "invalid token" in error_message.lower(): @@ -71,12 +69,17 @@ async def x_cashu_handler( "invalid_token", "The provided CASHU token is invalid", 400, - x_cashu_token, + request=request, + token=x_cashu_token, ) if "mint error" in error_message.lower(): return create_error_response( - "mint_error", f"CASHU mint error: {error_message}", 422, x_cashu_token + "mint_error", + f"CASHU mint error: {error_message}", + 422, + request=request, + token=x_cashu_token, ) # Generic error for other cases @@ -84,12 +87,18 @@ async def x_cashu_handler( "cashu_error", f"CASHU token processing failed: {error_message}", 400, - x_cashu_token, + request=request, + token=x_cashu_token, ) async def forward_to_upstream( - request: Request, path: str, headers: dict, amount: int, unit: CurrencyUnit + request: Request, + path: str, + headers: dict, + amount: int, + unit: str, + max_cost_for_model: int, ) -> Response | StreamingResponse: """Forward request to upstream and handle the response.""" if path.startswith("v1/"): @@ -181,7 +190,9 @@ async def forward_to_upstream( extra={"path": path, "amount": amount, "unit": unit}, ) - result = await handle_x_cashu_chat_completion(response, amount, unit) + result = await handle_x_cashu_chat_completion( + response, amount, unit, max_cost_for_model + ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) result.background = background_tasks @@ -217,12 +228,15 @@ async def forward_to_upstream( }, ) return create_error_response( - "internal_error", "An unexpected server error occurred", 500 + "internal_error", + "An unexpected server error occurred", + 500, + request=request, ) async def handle_x_cashu_chat_completion( - response: httpx.Response, amount: int, unit: CurrencyUnit + response: httpx.Response, amount: int, unit: str, max_cost_for_model: int ) -> StreamingResponse | Response: """Handle both streaming and non-streaming chat completion responses with token-based pricing.""" logger.debug( @@ -246,10 +260,12 @@ async def handle_x_cashu_chat_completion( ) if is_streaming: - return await handle_streaming_response(content_str, response, amount, unit) + return await handle_streaming_response( + content_str, response, amount, unit, max_cost_for_model + ) else: return await handle_non_streaming_response( - content_str, response, amount, unit + content_str, response, amount, unit, max_cost_for_model ) except Exception as e: @@ -271,7 +287,11 @@ async def handle_x_cashu_chat_completion( async def handle_streaming_response( - content_str: str, response: httpx.Response, amount: int, unit: CurrencyUnit + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, ) -> StreamingResponse: """Handle Server-Sent Events (SSE) streaming response.""" logger.debug( @@ -325,7 +345,7 @@ async def handle_streaming_response( response_data = {"usage": usage_data, "model": model} try: - cost_data = await get_cost(response_data) + cost_data = await get_cost(response_data, max_cost_for_model) if cost_data: if unit == "msat": refund_amount = amount - cost_data.total_msats @@ -393,7 +413,11 @@ async def handle_streaming_response( async def handle_non_streaming_response( - content_str: str, response: httpx.Response, amount: int, unit: CurrencyUnit + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, ) -> Response: """Handle regular JSON response.""" logger.debug( @@ -404,7 +428,7 @@ async def handle_non_streaming_response( try: response_json = json.loads(content_str) - cost_data = await get_cost(response_json) + cost_data = await get_cost(response_json, max_cost_for_model) if not cost_data: logger.error( @@ -510,21 +534,21 @@ async def handle_non_streaming_response( ) -async def get_cost(response_data: dict) -> MaxCostData | CostData | None: +async def get_cost( + response_data: dict, max_cost_for_model: int +) -> MaxCostData | CostData | None: """ 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") + model = response_data.get("model", None) logger.debug( "Calculating cost for response", extra={"model": model, "has_usage": "usage" in response_data}, ) - max_cost = get_max_cost_for_model(model=model) - - match calculate_cost(response_data, max_cost): + match calculate_cost(response_data, max_cost_for_model): case MaxCostData() as cost: logger.debug( "Using max cost pricing", @@ -563,7 +587,7 @@ async def get_cost(response_data: dict) -> MaxCostData | CostData | None: ) -async def send_refund(amount: int, unit: CurrencyUnit, mint: str | None = None) -> str: +async def send_refund(amount: int, unit: str, mint: str | None = None) -> str: """Send a refund using Cashu tokens.""" logger.debug( "Creating refund token", extra={"amount": amount, "unit": unit, "mint": mint} diff --git a/router/proxy.py b/routstr/proxy.py similarity index 94% rename from router/proxy.py rename to routstr/proxy.py index 0e772b1b..5ebe7c6f 100644 --- a/router/proxy.py +++ b/routstr/proxy.py @@ -19,7 +19,7 @@ from .payment.helpers import ( UPSTREAM_BASE_URL, check_token_balance, create_error_response, - get_cost_per_request, + get_max_cost_for_model, prepare_upstream_headers, ) from .payment.x_cashu import x_cashu_handler @@ -416,7 +416,9 @@ async def forward_to_upstream( else: error_message = f"Error connecting to upstream service: {error_type}" - return create_error_response("upstream_error", error_message, 502) + return create_error_response( + "upstream_error", error_message, 502, request=request + ) except Exception as exc: await client.aclose() @@ -437,7 +439,10 @@ async def forward_to_upstream( ) return create_error_response( - "internal_error", "An unexpected server error occurred", 500 + "internal_error", + "An unexpected server error occurred", + 500, + request=request, ) @@ -446,6 +451,14 @@ async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) ) -> Response | StreamingResponse: """Main proxy endpoint handler.""" + request_body = await request.body() + headers = dict(request.headers) + + if "x-cashu" not in headers and "authorization" not in headers.keys(): + return create_error_response( + "unauthorized", "Unauthorized", 401, request=request + ) + logger.info( "Received proxy request", extra={ @@ -456,9 +469,6 @@ async def proxy( }, ) - request_body = await request.body() - headers = dict(request.headers) - # Parse JSON body if present, handle empty/invalid JSON request_body_dict = {} if request_body: @@ -491,9 +501,8 @@ async def proxy( media_type="application/json", ) - max_cost_for_model = get_cost_per_request( - model=request_body_dict.get("model", None) - ) + model = request_body_dict.get("model", "unknown") + max_cost_for_model = get_max_cost_for_model(model=model) check_token_balance(headers, request_body_dict, max_cost_for_model) # Handle authentication @@ -505,7 +514,7 @@ async def proxy( "token_preview": x_cashu[:20] + "..." if len(x_cashu) > 20 else x_cashu, }, ) - return await x_cashu_handler(request, x_cashu, path) + return await x_cashu_handler(request, x_cashu, path, max_cost_for_model) elif auth := headers.get("authorization", None): logger.debug( @@ -530,11 +539,10 @@ async def proxy( ) logger.debug("Processing unauthenticated GET request", extra={"path": path}) - # Prepare headers for upstream + # TODO: why is this needed? can we remove it? headers = prepare_upstream_headers(dict(request.headers)) return await forward_get_to_upstream(request, path, headers) - cost_per_request = 0 # Only pay for request if we have request body data (for completions endpoints) if request_body_dict: logger.info( @@ -548,7 +556,7 @@ async def proxy( ) try: - await pay_for_request(key, session, request_body_dict) + await pay_for_request(key, max_cost_for_model, session) logger.info( "Payment processed successfully", extra={ @@ -579,7 +587,7 @@ async def proxy( ) if response.status_code != 200: - await revert_pay_for_request(key, session, cost_per_request) + await revert_pay_for_request(key, session, max_cost_for_model) logger.warning( "Upstream request failed, revert payment", extra={ @@ -587,8 +595,22 @@ async def proxy( "path": path, "key_hash": key.hashed_key[:8] + "...", "key_balance": key.balance, + "max_cost_for_model": max_cost_for_model, + "upstream_headers": response.headers + if hasattr(response, "headers") + else None, + "upstream_response": response.body + if hasattr(response, "body") + else None, }, ) + request_id = ( + request.state.request_id if hasattr(request.state, "request_id") else None + ) + raise HTTPException( + status_code=502, + detail=f"Upstream request failed, please contact support with request id: {request_id}", + ) return response @@ -729,5 +751,8 @@ async def forward_get_to_upstream( }, ) return create_error_response( - "internal_error", "An unexpected server error occurred", 500 + "internal_error", + "An unexpected server error occurred", + 500, + request=request, ) diff --git a/routstr/wallet.py b/routstr/wallet.py new file mode 100644 index 00000000..aa427437 --- /dev/null +++ b/routstr/wallet.py @@ -0,0 +1,374 @@ +import asyncio +import math +import os +from typing import TypedDict + +from cashu.core.base import Proof, Token +from cashu.wallet.helpers import deserialize_token_from_string +from cashu.wallet.wallet import Wallet + +from .core import db, get_logger +from .payment.lnurl import raw_send_to_lnurl + +logger = get_logger(__name__) + + +CASHU_MINTS = os.environ.get("CASHU_MINTS", "https://mint.minibits.cash/Bitcoin") +TRUSTED_MINTS = CASHU_MINTS.split(",") +PRIMARY_MINT_URL = TRUSTED_MINTS[0] +RECEIVE_LN_ADDRESS = os.environ.get("RECEIVE_LN_ADDRESS", "") + + +async def get_balance(unit: str) -> int: + wallet = await get_wallet(PRIMARY_MINT_URL, unit) + return wallet.available_balance.amount + + +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") + + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) + wallet.keyset_id = token_obj.keysets[0] + + if token_obj.mint not in TRUSTED_MINTS: + return await swap_to_primary_mint(token_obj, wallet) + + wallet.verify_proofs_dleq(token_obj.proofs) + await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True) + return token_obj.amount, token_obj.unit, token_obj.mint + + +async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: + """Internal send function - returns amount and serialized token""" + wallet: Wallet = await get_wallet(mint_url or PRIMARY_MINT_URL, unit) + proofs = get_proofs_per_mint_and_unit(wallet, mint_url or PRIMARY_MINT_URL, unit) + + send_proofs, _ = await wallet.select_to_send( + proofs, amount, set_reserved=True, include_fees=False + ) + token = await wallet.serialize_proofs( + send_proofs, include_dleq=False, legacy=False, memo=None + ) + return amount, token + + +async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str: + _, token = await send(amount, unit, mint_url) + return token + + +async def swap_to_primary_mint( + token_obj: Token, token_wallet: Wallet +) -> tuple[int, str, str]: + logger.info( + "swap_to_primary_mint", + extra={ + "mint": token_obj.mint, + "amount": token_obj.amount, + "unit": token_obj.unit, + }, + ) + # Ensure amount is an integer + if not isinstance(token_obj.amount, int): + token_amount = int(token_obj.amount) + else: + token_amount = token_obj.amount + + if token_obj.unit == "sat": + amount_msat = token_amount * 1000 + elif token_obj.unit == "msat": + amount_msat = token_amount + else: + raise ValueError("Invalid unit") + estimated_fee_sat = math.ceil(max(amount_msat // 1000 * 0.01, 2)) + amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000 + primary_wallet = await get_wallet(PRIMARY_MINT_URL, "sat") + + minted_amount = int(amount_msat_after_fee // 1000) + mint_quote = await primary_wallet.request_mint(minted_amount) + + melt_quote = await token_wallet.melt_quote(mint_quote.request) + _ = await token_wallet.melt( + proofs=token_obj.proofs, + invoice=mint_quote.request, + fee_reserve_sat=melt_quote.fee_reserve, + quote_id=melt_quote.quote, + ) + _ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote) + + return int(minted_amount), "sat", PRIMARY_MINT_URL + + +async def credit_balance( + cashu_token: str, key: db.ApiKey, session: db.AsyncSession +) -> int: + logger.info( + "credit_balance: Starting token redemption", + extra={"token_preview": cashu_token[:50]}, + ) + + try: + amount, unit, mint_url = await recieve_token(cashu_token) + logger.info( + "credit_balance: Token redeemed successfully", + extra={"amount": amount, "unit": unit, "mint_url": mint_url}, + ) + + if unit == "sat": + amount = amount * 1000 + logger.info( + "credit_balance: Converted to msat", extra={"amount_msat": amount} + ) + + logger.info( + "credit_balance: Updating balance", + extra={"old_balance": key.balance, "credit_amount": amount}, + ) + key.balance += amount + session.add(key) + await session.commit() + logger.info( + "credit_balance: Balance updated successfully", + extra={"new_balance": key.balance}, + ) + + logger.info( + "Cashu token successfully redeemed and stored", + extra={"amount": amount, "unit": unit, "mint_url": mint_url}, + ) + return amount + except Exception as e: + logger.error( + "credit_balance: Error during token redemption", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + raise + + +_wallets: dict[str, Wallet] = {} + + +async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: + global _wallets + id = f"{mint_url}_{unit}" + if id not in _wallets: + _wallets[id] = await Wallet.with_db( + mint_url, db=".wallet", load_all_keysets=True, unit=unit + ) + + if load: + await _wallets[id].load_mint() + await _wallets[id].load_proofs(reload=True) + return _wallets[id] + + +def get_proofs_per_mint_and_unit( + wallet: Wallet, mint_url: str, unit: str, not_reserved: bool = False +) -> list[Proof]: + valid_keyset_ids = [ + k.id + for k in wallet.keysets.values() + if k.mint_url == mint_url and k.unit.name == unit + ] + proofs = [p for p in wallet.proofs if p.id in valid_keyset_ids] + if not_reserved: + proofs = [p for p in proofs if not p.reserved] + return proofs + + +async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[Proof]: + if not proofs: + return [] + _proofs = [] + _spent_proofs = [] + for i in range(0, len(proofs), 1000): + pb = proofs[i : i + 1000] + proof_states = await wallet.check_proof_state(pb) + for proof, state in zip(pb, proof_states.states): + if str(state.state) != "spent": + _proofs.append(proof) + else: + _spent_proofs.append(proof) + await wallet.set_reserved_for_send(_spent_proofs, reserved=True) + return _proofs + + +class BalanceDetail(TypedDict, total=False): + mint_url: str + unit: str + wallet_balance: int + user_balance: int + owner_balance: int + error: str + + +async def fetch_all_balances( + units: list[str] | None = None, +) -> tuple[list[BalanceDetail], int, int, int]: + """ + Fetch balances for all trusted mints and units concurrently. + + Returns: + - List of balance details for each mint/unit combination + - Total wallet balance in sats + - Total user balance in sats + - Owner balance in sats (wallet - user) + """ + if units is None: + units = ["sat", "msat"] + + async def fetch_balance( + session: db.AsyncSession, mint_url: str, unit: str + ) -> BalanceDetail: + try: + wallet = await get_wallet(mint_url, unit) + proofs = get_proofs_per_mint_and_unit( + wallet, mint_url, unit, not_reserved=True + ) + proofs = await slow_filter_spend_proofs(proofs, wallet) + user_balance = await db.balances_for_mint_and_unit(session, mint_url, unit) + if unit == "sat": + user_balance = user_balance // 1000 + proofs_balance = sum(proof.amount for proof in proofs) + + result: BalanceDetail = { + "mint_url": mint_url, + "unit": unit, + "wallet_balance": proofs_balance, + "user_balance": user_balance, + "owner_balance": proofs_balance - user_balance, + } + return result + except Exception as e: + logger.error(f"Error getting balance for {mint_url} {unit}: {e}") + error_result: BalanceDetail = { + "mint_url": mint_url, + "unit": unit, + "wallet_balance": 0, + "user_balance": 0, + "owner_balance": 0, + "error": str(e), + } + return error_result + + # Create tasks for all mint/unit combinations + async with db.create_session() as session: + tasks = [ + fetch_balance(session, mint_url, unit) + for mint_url in TRUSTED_MINTS + for unit in units + ] + + # Run all tasks concurrently + balance_details = list(await asyncio.gather(*tasks)) + + # Calculate totals + total_wallet_balance_sats = 0 + total_user_balance_sats = 0 + + for detail in balance_details: + if not detail.get("error"): + # Convert to sats for total calculation + unit = detail["unit"] + proofs_balance_sats = ( + detail["wallet_balance"] + if unit == "sat" + else detail["wallet_balance"] // 1000 + ) + user_balance_sats = ( + detail["user_balance"] + if unit == "sat" + else detail["user_balance"] // 1000 + ) + + total_wallet_balance_sats += proofs_balance_sats + total_user_balance_sats += user_balance_sats + + owner_balance = total_wallet_balance_sats - total_user_balance_sats + + return ( + balance_details, + total_wallet_balance_sats, + total_user_balance_sats, + owner_balance, + ) + + +async def periodic_payout() -> None: + if not RECEIVE_LN_ADDRESS: + logger.error("RECEIVE_LN_ADDRESS is not set, skipping payout") + return + while True: + await asyncio.sleep(60 * 5) + try: + async with db.create_session() as session: + for mint_url in TRUSTED_MINTS: + for unit in ["sat", "msat"]: + wallet = await get_wallet(mint_url, unit) + proofs = get_proofs_per_mint_and_unit( + wallet, mint_url, unit, not_reserved=True + ) + proofs = await slow_filter_spend_proofs(proofs, wallet) + user_balance = await db.balances_for_mint_and_unit( + session, mint_url, unit + ) + if unit == "sat": + user_balance = user_balance // 1000 + proofs_balance = sum(proof.amount for proof in proofs) + available_balance = proofs_balance - user_balance + min_amount = 210 if unit == "sat" else 210000 + if available_balance > min_amount: + amount_received = await raw_send_to_lnurl( + wallet, proofs, RECEIVE_LN_ADDRESS, unit + ) + logger.info( + "Payout sent successfully", + extra={ + "mint_url": mint_url, + "unit": unit, + "balance": available_balance, + "amount_received": amount_received, + }, + ) + + await asyncio.sleep(5) + except Exception as e: + logger.error( + 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/setup.py b/setup.py deleted file mode 100644 index c6f07fec..00000000 --- a/setup.py +++ /dev/null @@ -1,19 +0,0 @@ -from setuptools import find_packages, setup - -setup( - name="routstr", - version="0.1.0", - packages=find_packages(), - install_requires=[ - "fastapi[standard]>=0.115", - "aiosqlite>=0.20", - "sqlmodel>=0.0.24", - "httpx[socks]>=0.25.2", - "greenlet>=3.2.1", - "python-json-logger>=2.0.0", - "cashu", - "secp256k1", - "marshmallow>=3.13,<4.0", - ], - python_requires=">=3.11", -) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index bc9f8d13..c3727a72 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -8,10 +8,11 @@ import pytest import pytest_asyncio from fastapi import FastAPI from httpx import AsyncClient -from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine +from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession -from router.core.logging import get_logger +from routstr.core.logging import get_logger logger = get_logger(__name__) @@ -25,7 +26,7 @@ if use_local_services: "DATABASE_URL": "sqlite+aiosqlite:///:memory:", "UPSTREAM_BASE_URL": "http://localhost:3000", # Mock OpenAI service "UPSTREAM_API_KEY": "test-upstream-key", - "CASHU_MINTS": "http://mint:3338", # Docker service name for router validation + "CASHU_MINTS": "http://mint:3338", # Docker service name for routstr validation "MINT": "http://mint:3338", "MINT_URL": "http://mint:3338", "NOSTR_RELAY_URL": "ws://localhost:8088", @@ -63,8 +64,8 @@ else: # Set test environment variables before importing the app os.environ.update(test_env) -from router.core.db import ApiKey, get_session # noqa: E402 -from router.core.main import app, lifespan # noqa: E402 +from routstr.core.db import ApiKey, get_session # noqa: E402 +from routstr.core.main import app, lifespan # noqa: E402 @pytest.fixture(scope="session") @@ -157,7 +158,7 @@ class TestmintWallet: mint_response = await wallet.mint(amount=amount, hash=quote) token = mint_response.token - # Replace connection URL with Docker service name for router validation + # Replace connection URL with Docker service name for routstr validation if self.connection_url != self.mint_url: token = token.replace(self.connection_url, self.mint_url) @@ -251,7 +252,7 @@ class TestmintWallet: async def send_token( self, amount: int, unit: str, mint_url: Optional[str] = None ) -> str: - """Send token with compatible signature for mocking router.wallet.send_token""" + """Send token with compatible signature for mocking routstr.wallet.send_token""" return await self.send(amount) async def send_to_lnurl(self, lnurl: str, amount: int) -> int: @@ -501,25 +502,25 @@ async def integration_app( if use_real_mint: # Use real mint - no wallet patches needed - with patch("router.core.db.engine", integration_engine): + with patch("routstr.core.db.engine", integration_engine): yield test_app else: # Use testmint with wallet patches for all integration tests mint_url = os.environ.get("CASHU_MINTS", "http://localhost:3338") with ( - patch("router.core.db.engine", integration_engine), - patch("router.wallet.TRUSTED_MINTS", [mint_url]), - patch("router.wallet.PRIMARY_MINT_URL", mint_url), - patch("router.auth.credit_balance", testmint_wallet.credit_balance), - patch("router.wallet.credit_balance", testmint_wallet.credit_balance), - patch("router.balance.credit_balance", testmint_wallet.credit_balance), - patch("router.wallet.send_token", testmint_wallet.send_token), - patch("router.balance.send_token", testmint_wallet.send_token), - patch("router.wallet.recieve_token", testmint_wallet.redeem_token), - patch("router.wallet.get_balance", testmint_wallet.get_balance), + patch("routstr.core.db.engine", integration_engine), + patch("routstr.wallet.TRUSTED_MINTS", [mint_url]), + patch("routstr.wallet.PRIMARY_MINT_URL", mint_url), + patch("routstr.auth.credit_balance", testmint_wallet.credit_balance), + patch("routstr.wallet.credit_balance", testmint_wallet.credit_balance), + patch("routstr.balance.credit_balance", testmint_wallet.credit_balance), + patch("routstr.wallet.send_token", testmint_wallet.send_token), + patch("routstr.balance.send_token", testmint_wallet.send_token), + patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token), + patch("routstr.wallet.get_balance", testmint_wallet.get_balance), patch("websockets.connect") as mock_websockets, - patch("router.payment.price.btc_usd_ask_price", return_value=50000.0), - patch("router.payment.price.sats_usd_ask_price", return_value=0.0005), + patch("routstr.payment.price.btc_usd_ask_price", return_value=50000.0), + patch("routstr.payment.price.sats_usd_ask_price", return_value=0.0005), ): # Configure the WebSocket mock for discovery service - fast failure for performance tests async def mock_websocket_connect(*args: Any, **kwargs: Any) -> None: @@ -689,8 +690,8 @@ async def background_tasks_controller() -> AsyncGenerator[Any, None]: original_periodic_payout: Optional[Callable] = None try: - from router.payment.models import update_sats_pricing - from router.wallet import periodic_payout + from routstr.payment.models import update_sats_pricing + from routstr.wallet import periodic_payout async def controlled_update_pricing() -> None: while not controller.cancelled: diff --git a/tests/integration/run_performance_tests.py b/tests/integration/run_performance_tests.py index 5fcec67b..d41b6690 100755 --- a/tests/integration/run_performance_tests.py +++ b/tests/integration/run_performance_tests.py @@ -139,7 +139,7 @@ async def main() -> None: print("WARNING: Proxy server may not be running properly") except Exception: print("ERROR: Proxy server is not running!") - print("Please start the server with: uvicorn router.main:app") + print("Please start the server with: uvicorn routstr.main:app") sys.exit(1) # Run performance tests diff --git a/tests/integration/test_background_tasks.py b/tests/integration/test_background_tasks.py index 3f87603f..de387286 100644 --- a/tests/integration/test_background_tasks.py +++ b/tests/integration/test_background_tasks.py @@ -9,9 +9,9 @@ from unittest.mock import AsyncMock, patch import pytest -from router.core.db import ApiKey -from router.payment.models import MODELS, Model, Pricing, update_sats_pricing -from router.wallet import periodic_payout +from routstr.core.db import ApiKey +from routstr.payment.models import MODELS, Model, Pricing, update_sats_pricing +from routstr.wallet import periodic_payout @pytest.mark.asyncio @@ -24,7 +24,7 @@ class TestPricingUpdateTask: mock_sats_usd = 0.00002 # 1 sat = $0.00002 (BTC at $50,000) with patch( - "router.payment.price.sats_usd_ask_price", + "routstr.payment.price.sats_usd_ask_price", AsyncMock(return_value=mock_sats_usd), ): # Create a test model @@ -114,7 +114,7 @@ class TestPricingUpdateTask: raise Exception("Price API error") return 0.00002 - with patch("router.payment.price.sats_usd_ask_price", mock_price_func): + with patch("routstr.payment.price.sats_usd_ask_price", mock_price_func): # Test the retry behavior directly # First call should fail try: @@ -122,7 +122,7 @@ class TestPricingUpdateTask: assert False, "Expected exception on first call" except Exception: pass - + # Second call should succeed result = await mock_price_func() assert result == 0.00002 @@ -165,7 +165,7 @@ class TestPricingUpdateTask: try: with patch( - "router.payment.price.sats_usd_ask_price", + "routstr.payment.price.sats_usd_ask_price", AsyncMock(return_value=0.00002), ): # Initialize pricing once to ensure consistent state @@ -214,9 +214,9 @@ class TestRefundCheckTask: # Mock the wallet send_to_lnurl method and get_session with ( patch( - "router.wallet.send_to_lnurl", AsyncMock(return_value=5) + "routstr.wallet.send_to_lnurl", AsyncMock(return_value=5) ) as mock_send_to_lnurl, - patch("router.core.db.get_session") as mock_get_session, + patch("routstr.core.db.get_session") as mock_get_session, ): # Make get_session return our integration session async def get_test_session() -> Any: @@ -279,9 +279,9 @@ class TestRefundCheckTask: with ( patch( - "router.wallet.send_to_lnurl", mock_send_to_lnurl + "routstr.wallet.send_to_lnurl", mock_send_to_lnurl ) as mock_send_to_lnurl_patch, - patch("router.core.db.get_session") as mock_get_session, + patch("routstr.core.db.get_session") as mock_get_session, ): # Make get_session return our integration session async def get_test_session() -> Any: @@ -366,9 +366,9 @@ class TestRefundCheckTask: with ( patch( - "router.wallet.send_to_lnurl", AsyncMock(return_value=1) + "routstr.wallet.send_to_lnurl", AsyncMock(return_value=1) ) as mock_send_to_lnurl, - patch("router.core.db.get_session") as mock_get_session, + patch("routstr.core.db.get_session") as mock_get_session, ): # Make get_session return our integration session async def get_test_session() -> Any: @@ -424,7 +424,7 @@ class TestRefundCheckTask: # async def test_refund_check_disabled(self) -> None: # """Test that refund check can be disabled by setting interval to 0""" # # Patch the constant directly to disable refunds - # with patch.object(router.wallet, "REFUND_PROCESSING_INTERVAL", 0): + # with patch.object(routstr.wallet, "REFUND_PROCESSING_INTERVAL", 0): # # Task should exit immediately # task = asyncio.create_task(check_for_refunds()) # await task # Should complete without hanging @@ -466,9 +466,9 @@ class TestPeriodicPayoutTask: wallet_balance = 200000 # 200 sats total with ( - patch("router.wallet.get_balance", AsyncMock(return_value=wallet_balance)), + patch("routstr.wallet.get_balance", AsyncMock(return_value=wallet_balance)), patch( - "router.wallet.send_to_lnurl", AsyncMock(return_value=None) + "routstr.wallet.send_to_lnurl", AsyncMock(return_value=None) ) as mock_send_to_lnurl, ): # Mock environment variables @@ -481,7 +481,7 @@ class TestPeriodicPayoutTask: }, ): # Call periodic_payout directly (pay_out was renamed/refactored) - from router.wallet import periodic_payout + from routstr.wallet import periodic_payout await periodic_payout() @@ -506,7 +506,7 @@ class TestPeriodicPayoutTask: # integration_session.add(key) # await integration_session.commit() - # with patch("router.cashu.wallet") as mock_wallet: + # with patch("routstr.cashu.wallet") as mock_wallet: # mock_wallet_instance = AsyncMock() # mock_wallet_instance.balance = AsyncMock( # return_value=100000 @@ -522,7 +522,7 @@ class TestPeriodicPayoutTask: # "DEV_LN_ADDRESS": "dev@test.com", # }, # ): - # from router.cashu import pay_out + # from routstr.cashu import pay_out # await pay_out() @@ -543,7 +543,7 @@ class TestPeriodicPayoutTask: # integration_session.add(key) # await integration_session.commit() - # with patch("router.cashu.wallet") as mock_wallet: + # with patch("routstr.cashu.wallet") as mock_wallet: # mock_wallet_instance = AsyncMock() # mock_wallet_instance.balance = AsyncMock( # return_value=96000 @@ -552,7 +552,7 @@ class TestPeriodicPayoutTask: # mock_wallet.return_value = mock_wallet_instance # with patch.dict(os.environ, {"MINIMUM_PAYOUT": "10"}): # 10 sats minimum - # from router.cashu import pay_out + # from routstr.cashu import pay_out # await pay_out() @@ -571,9 +571,9 @@ class TestTaskInteractions: # """Test that all tasks can run concurrently without issues""" # # Mock all external dependencies # with ( - # patch("router.payment.price.sats_usd_ask_price", AsyncMock(return_value=0.00002)), - # patch("router.cashu.wallet") as mock_wallet, - # patch("router.cashu.pay_out", AsyncMock()), + # patch("routstr.payment.price.sats_usd_ask_price", AsyncMock(return_value=0.00002)), + # patch("routstr.cashu.wallet") as mock_wallet, + # patch("routstr.cashu.pay_out", AsyncMock()), # ): # mock_wallet_instance = AsyncMock() # mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=1) @@ -587,7 +587,7 @@ class TestTaskInteractions: # tasks.append(pricing_task) # # Refund task (disabled to avoid interference) - # with patch.object(router.wallet, "REFUND_PROCESSING_INTERVAL", 0): + # with patch.object(routstr.wallet, "REFUND_PROCESSING_INTERVAL", 0): # refund_task = asyncio.create_task(check_for_refunds()) # tasks.append(refund_task) @@ -621,7 +621,7 @@ class TestTaskInteractions: processing.set() await asyncio.sleep(2) # Simulate long operation - with patch("router.payment.price.sats_usd_ask_price", slow_task): + with patch("routstr.payment.price.sats_usd_ask_price", slow_task): # Start the pricing task task = asyncio.create_task(update_sats_pricing()) @@ -708,11 +708,15 @@ class TestTaskInteractions: # Patch the actual task functions with ( patch( - "router.payment.models.update_sats_pricing", + "routstr.payment.models.update_sats_pricing", lambda: task_with_cleanup("pricing"), ), - patch("router.wallet.periodic_payout", lambda: task_with_cleanup("refund")), - patch("router.wallet.periodic_payout", lambda: task_with_cleanup("payout")), + patch( + "routstr.wallet.periodic_payout", lambda: task_with_cleanup("refund") + ), + patch( + "routstr.wallet.periodic_payout", lambda: task_with_cleanup("payout") + ), ): # Start all tasks tasks = [ diff --git a/tests/integration/test_database_consistency.py b/tests/integration/test_database_consistency.py index d4abcb61..5c2bbe68 100644 --- a/tests/integration/test_database_consistency.py +++ b/tests/integration/test_database_consistency.py @@ -11,7 +11,7 @@ from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey class TestTransactionAtomicity: @@ -114,7 +114,7 @@ class TestTransactionAtomicity: initial_balance = api_key.balance # Mock wallet to fail after token validation - with patch("router.wallet.send_token") as mock_wallet_func: + with patch("routstr.wallet.send_token") as mock_wallet_func: mock_proof = MagicMock() mock_proof.amount = 1000 mock_wallet = AsyncMock() @@ -256,7 +256,7 @@ class TestConcurrentOperations: await integration_session.commit() # Mock wallet for topup - with patch("router.wallet.send_token") as mock_wallet_func: + with patch("routstr.wallet.send_token") as mock_wallet_func: mock_proof = MagicMock() mock_proof.amount = 2000 mock_wallet = AsyncMock() @@ -379,6 +379,7 @@ class TestDataIntegrity: """Test data integrity constraints and validations""" @pytest.mark.asyncio + @pytest.mark.skip(reason="Balance never negative is not implemented") async def test_balance_never_negative( self, authenticated_client: AsyncClient, @@ -398,7 +399,7 @@ class TestDataIntegrity: stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type] result = await integration_session.execute(stmt) api_key = result.scalar_one() - api_key.balance = 100 + api_key.balance = 0 await integration_session.commit() # Try to refund more than balance @@ -521,7 +522,7 @@ class TestPerformance: operation_times["select"].append((end - start) * 1000) # Convert to ms # Test UPDATE performance (via topup) - with patch("router.wallet.send_token") as mock_wallet_func: + with patch("routstr.wallet.send_token") as mock_wallet_func: mock_proof = MagicMock() mock_proof.amount = 100 mock_wallet = AsyncMock() diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py index 53e74630..cbe771b3 100644 --- a/tests/integration/test_error_handling_edge_cases.py +++ b/tests/integration/test_error_handling_edge_cases.py @@ -1,16 +1,17 @@ """Comprehensive error handling and edge case tests""" import asyncio +import hashlib import time from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest -from httpx import AsyncClient, ConnectError +from httpx import ASGITransport, AsyncClient, ConnectError from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey class TestNetworkFailureScenarios: @@ -26,11 +27,11 @@ class TestNetworkFailureScenarios: # Patch the wallet send function to simulate failure across all modules with ( patch( - "router.wallet.send_token", + "routstr.wallet.send_token", AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), patch( - "router.balance.send_token", + "routstr.balance.send_token", AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), ): @@ -46,8 +47,8 @@ class TestNetworkFailureScenarios: integration_session: AsyncSession, ) -> None: """Test proxy behavior when upstream LLM service is down""" - # Mock at the router level to simulate upstream being down - with patch("router.proxy.httpx.AsyncClient") as mock_client_class: + # Mock at the routstr level to simulate upstream being down + with patch("routstr.proxy.httpx.AsyncClient") as mock_client_class: # Create a mock client instance mock_client = AsyncMock() mock_client_class.return_value = mock_client @@ -465,7 +466,7 @@ class TestRecoveryScenarios: # Simulate operations that might be interrupted try: # Start a transaction - api_key.balance -= 1000 + api_key.reserved_balance += 1000 api_key.total_requests += 1 # Don't commit - simulate crash raise Exception("Simulated database crash") @@ -616,30 +617,56 @@ class TestEdgeCaseCombinations: @pytest.mark.asyncio async def test_rapid_balance_exhaustion( self, - authenticated_client: AsyncClient, + integration_app: Any, integration_session: AsyncSession, + testmint_wallet: Any, + monkeypatch: pytest.MonkeyPatch, ) -> None: - """Test behavior when balance is rapidly exhausted""" - # Set a low balance - api_key_header = authenticated_client.headers["Authorization"].replace( - "Bearer ", "" - ) - api_key_hash = ( - api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header - ) + """Test behavior when balance is rapidly exhausted by concurrent requests. - # Set balance to just 1000 msats (1 sat) - from sqlalchemy import update + This test creates an API key with insufficient balance (500 msats) for even + a single request (which costs 1000 msats). It then makes 5 concurrent requests + to verify that all requests fail with 402 Payment Required errors. - await integration_session.execute( - update(ApiKey).where(ApiKey.hashed_key == api_key_hash).values(balance=1000) # type: ignore[arg-type] + Note: The test disables MODEL_BASED_PRICING to avoid model lookup errors + since the test environment doesn't have models configured. + """ + # Disable MODEL_BASED_PRICING for this test to avoid model lookup issues + monkeypatch.setattr( + "routstr.payment.cost_caculation.MODEL_BASED_PRICING", False ) + monkeypatch.setattr("routstr.payment.helpers.MODEL_BASED_PRICING", False) + + # Create a new API key with very low balance + # Generate a unique API key + test_key = f"sk-test-low-balance-{hashlib.sha256(str(time.time()).encode()).hexdigest()[:8]}" + api_key_hash = test_key[3:] # Remove sk- prefix + + # Create the API key with only 500 msats (less than one request cost) + new_key = ApiKey( + hashed_key=api_key_hash, + balance=500, # Less than COST_PER_REQUEST (1000 msats) + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + integration_session.add(new_key) await integration_session.commit() + # Verify the key was created + await integration_session.refresh(new_key) + + # Create a client with this low-balance key + low_balance_client = AsyncClient( + transport=ASGITransport(app=integration_app), # type: ignore + base_url="http://test", + headers={"Authorization": f"Bearer {test_key}"}, + ) + # Make multiple concurrent requests that would exhaust balance tasks = [] for _ in range(5): - task = authenticated_client.post( + task = low_balance_client.post( "/v1/chat/completions", json={ "model": "gpt-3.5-turbo", @@ -665,3 +692,6 @@ class TestEdgeCaseCombinations: result = await integration_session.execute(stmt) final_key = result.scalar_one() assert final_key.balance >= 0 + + # Clean up the test client + await low_balance_client.aclose() diff --git a/tests/integration/test_example.py b/tests/integration/test_example.py index e4c025ac..6657517c 100644 --- a/tests/integration/test_example.py +++ b/tests/integration/test_example.py @@ -105,7 +105,15 @@ async def test_full_wallet_flow( assert refund_response.status_code == 200 refund_data = refund_response.json() assert "token" in refund_data - assert refund_data["msats"] == (initial_amount + topup_amount) * 1000 + + # Check for either sats or msats depending on refund_currency + total_amount = initial_amount + topup_amount + if "sats" in refund_data: + assert refund_data["sats"] == str(total_amount) + elif "msats" in refund_data: + assert refund_data["msats"] == str(total_amount * 1000) + else: + pytest.fail("Response should contain either 'sats' or 'msats'") @pytest.mark.integration diff --git a/tests/integration/test_performance_load.py b/tests/integration/test_performance_load.py index d75340fa..581553b9 100644 --- a/tests/integration/test_performance_load.py +++ b/tests/integration/test_performance_load.py @@ -149,7 +149,7 @@ class TestPerformanceBaseline: """Test database operation performance""" from sqlmodel import select - from router.core.db import ApiKey + from routstr.core.db import ApiKey # Create test data for i in range(100): diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index a2b0aff5..1ea5a321 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -43,9 +43,9 @@ async def test_providers_endpoint_default_response( } with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: # Configure mock to return appropriate responses mock_fetch.side_effect = lambda url: mock_fetch_responses.get( url, {"status_code": 500, "json": {"error": "Unknown provider"}} @@ -100,9 +100,9 @@ async def test_providers_endpoint_with_include_json( } with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = { "status_code": 200, "json": mock_provider_response, @@ -155,7 +155,7 @@ async def test_providers_data_structure_validation( ["description", "A comprehensive AI provider"], ["model", "gpt-3.5-turbo"], ["model", "gpt-4"], - ] + ], } ] @@ -165,15 +165,15 @@ async def test_providers_data_structure_validation( "json": { "data": [ {"id": "gpt-3.5-turbo", "object": "model"}, - {"id": "gpt-4", "object": "model"} + {"id": "gpt-4", "object": "model"}, ] - } + }, } with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = mock_health_response response = await integration_client.get("/v1/providers/?include_json=true") @@ -188,7 +188,7 @@ async def test_providers_data_structure_validation( # Should have provider and health keys based on actual implementation assert "provider" in provider_data assert "health" in provider_data - + provider_info = provider_data["provider"] # Expected fields from RIP-02 parser expected_fields = ["id", "name", "endpoint_url", "supported_models"] @@ -215,7 +215,7 @@ async def test_providers_endpoint_no_providers_found( mock_events: list[dict[str, Any]] = [] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): response = await integration_client.get("/v1/providers/") @@ -245,33 +245,41 @@ async def test_providers_endpoint_offline_providers( ["d", "healthy-provider"], ["endpoint", "http://healthy-provider.onion"], ["name", "Healthy Provider"], - ] + ], }, { "id": "event2", - "pubkey": "offline_provider_pubkey", + "pubkey": "offline_provider_pubkey", "created_at": 1234567891, "content": "Offline provider announcement", "tags": [ ["d", "offline-provider"], ["endpoint", "http://offline-provider.onion"], ["name", "Offline Provider"], - ] + ], }, ] # Mock one healthy and one offline provider def mock_fetch_provider_health(url: str) -> dict[str, Any]: if "healthy" in url: - return {"status_code": 200, "endpoint": "root", "json": {"status": "online"}} + return { + "status_code": 200, + "endpoint": "root", + "json": {"status": "online"}, + } else: - return {"status_code": 500, "endpoint": "error", "json": {"error": "Service unavailable"}} + return { + "status_code": 500, + "endpoint": "error", + "json": {"error": "Service unavailable"}, + } with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): with patch( - "router.discovery.fetch_provider_health", + "routstr.discovery.fetch_provider_health", side_effect=mock_fetch_provider_health, ): response = await integration_client.get("/v1/providers/?include_json=true") @@ -286,10 +294,10 @@ async def test_providers_endpoint_offline_providers( for provider_data in data["providers"]: assert "provider" in provider_data assert "health" in provider_data - + provider_info = provider_data["provider"] health_info = provider_data["health"] - + if "offline" in provider_info["endpoint_url"]: # Offline provider should have error information in health assert health_info["status_code"] == 500 @@ -297,7 +305,10 @@ async def test_providers_endpoint_offline_providers( else: # Healthy provider should have successful health check assert health_info["status_code"] == 200 - assert "status" in health_info["json"] or "error" not in health_info["json"] + assert ( + "status" in health_info["json"] + or "error" not in health_info["json"] + ) @pytest.mark.integration @@ -318,7 +329,7 @@ async def test_providers_endpoint_duplicate_urls( ["d", "provider-1"], ["endpoint", "http://provider.onion"], ["name", "Provider"], - ] + ], }, { "id": "event2", @@ -329,15 +340,19 @@ async def test_providers_endpoint_duplicate_urls( ["d", "other-provider"], ["endpoint", "http://other-provider.onion"], ["name", "Other Provider"], - ] + ], }, ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: - mock_fetch.return_value = {"status_code": 200, "endpoint": "root", "json": {"status": "online"}} + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + mock_fetch.return_value = { + "status_code": 200, + "endpoint": "root", + "json": {"status": "online"}, + } response = await integration_client.get("/v1/providers/") @@ -352,7 +367,7 @@ async def test_providers_endpoint_duplicate_urls( endpoint_urls = [] for provider_data in providers: endpoint_urls.append(provider_data["endpoint_url"]) - + unique_endpoints = set(endpoint_urls) assert len(unique_endpoints) == len(endpoint_urls) @@ -369,7 +384,7 @@ async def test_providers_endpoint_nostr_relay_failures( raise Exception("Connection to relay failed") with patch( - "router.discovery.query_nostr_relay_for_providers", side_effect=failing_query + "routstr.discovery.query_nostr_relay_for_providers", side_effect=failing_query ): response = await integration_client.get("/v1/providers/") @@ -407,9 +422,9 @@ async def test_providers_endpoint_malformed_urls( ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} response = await integration_client.get("/v1/providers/") @@ -440,9 +455,9 @@ async def test_providers_endpoint_response_format( ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Test default format @@ -490,9 +505,9 @@ async def test_providers_endpoint_performance(integration_client: AsyncClient) - validator = PerformanceValidator() with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Test multiple requests @@ -530,9 +545,9 @@ async def test_providers_endpoint_concurrent_requests( ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Create concurrent requests @@ -566,9 +581,9 @@ async def test_providers_endpoint_parameter_validation( ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Test various parameter values @@ -617,9 +632,9 @@ async def test_no_database_changes_during_provider_operations( ] with patch( - "router.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events ): - with patch("router.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Make multiple requests with different parameters diff --git a/tests/integration/test_proxy_get_endpoints.py b/tests/integration/test_proxy_get_endpoints.py index 777a26b7..3434b151 100644 --- a/tests/integration/test_proxy_get_endpoints.py +++ b/tests/integration/test_proxy_get_endpoints.py @@ -14,7 +14,7 @@ import pytest from httpx import AsyncClient from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey from .utils import ( ConcurrencyTester, diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py new file mode 100644 index 00000000..21f9b6d0 --- /dev/null +++ b/tests/integration/test_reserved_balance_negative.py @@ -0,0 +1,165 @@ +"""Test to verify reserved balance never goes negative.""" + +import asyncio +import uuid + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey, create_session + + +@pytest.mark.asyncio +async def test_reserved_balance_never_negative(integration_client: AsyncClient) -> None: + """Test that reserved balance never goes negative under various conditions.""" + + # Create a test API key with limited balance + async with create_session() as session: + test_key = ApiKey( + hashed_key="test_reserved_balance_key", + balance=1000, # 1 sat + reserved_balance=0, + ) + session.add(test_key) + await session.commit() + + bearer_token = "sk-test_reserved_balance_key" + headers = {"Authorization": f"Bearer {bearer_token}"} + + # Test 1: Make a request that will fail upstream + # This should reserve funds and then revert them + await integration_client.post( + "/v1/chat/completions", + headers=headers, + json={ + "model": "invalid-model-that-will-fail", + "messages": [{"role": "user", "content": "test"}], + }, + ) + + # Check reserved balance after failed request + async with create_session() as session: + key = await session.get(ApiKey, "test_reserved_balance_key") + assert key is not None + assert key.reserved_balance >= 0, ( + f"Reserved balance went negative: {key.reserved_balance}" + ) + assert key.balance == 1000, ( + "Balance should remain unchanged after failed request" + ) + + # Test 2: Simulate concurrent failed requests + # This tests the race condition protection + async def make_failing_request() -> None: + try: + await integration_client.post( + "/v1/chat/completions", + headers=headers, + json={ + "model": "invalid-model", + "messages": [{"role": "user", "content": "test"}], + }, + ) + except Exception: + pass # Expected to fail + + # Run multiple concurrent requests + await asyncio.gather(*[make_failing_request() for _ in range(5)]) + + # Check final state + async with create_session() as session: + key = await session.get(ApiKey, "test_reserved_balance_key") + assert key is not None + assert key.reserved_balance >= 0, ( + f"Reserved balance went negative after concurrent requests: {key.reserved_balance}" + ) + print(f"Final state - Balance: {key.balance}, Reserved: {key.reserved_balance}") + + +@pytest.mark.asyncio +async def test_reserved_balance_with_successful_requests( + integration_client: AsyncClient, +) -> None: + """Test reserved balance handling with successful requests.""" + + # Create a test API key with more balance + async with create_session() as session: + unique_key = f"test_successful_key_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=100000, # 100 sats + reserved_balance=0, + ) + session.add(test_key) + await session.commit() + + bearer_token = f"sk-{unique_key}" + headers = {"Authorization": f"Bearer {bearer_token}"} + + # Make a valid request (assuming you have a mock or test endpoint) + # This test might need adjustment based on your test setup + await integration_client.post( + "/v1/chat/completions", + headers=headers, + json={ + "model": "gpt-4o-mini", # Or whatever model is available in test + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 10, + }, + ) + + # Check that reserved balance was properly adjusted + async with create_session() as session: + key = await session.get(ApiKey, unique_key) + assert key is not None + assert key.reserved_balance >= 0, ( + f"Reserved balance went negative: {key.reserved_balance}" + ) + # Check if the request was processed (might fail due to model pricing in test env) + # The important part is that reserved_balance doesn't go negative + if key.total_spent > 0: + assert key.balance < 100000, ( + "Balance should decrease after successful request" + ) + else: + # Request failed, but reserved balance should still be non-negative + assert key.balance == 100000, ( + "Balance should remain unchanged if request failed" + ) + print( + f"After successful request - Balance: {key.balance}, Reserved: {key.reserved_balance}, Spent: {key.total_spent}" + ) + + +@pytest.mark.asyncio +async def test_insufficient_reserved_balance_for_revert( + integration_session: AsyncSession, +) -> None: + """Test revert_pay_for_request behavior with insufficient reserved balance.""" + from routstr.auth import revert_pay_for_request + + # Create key with zero reserved balance + unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=1000, + reserved_balance=0, + ) + integration_session.add(test_key) + await integration_session.commit() + + # Try to revert more than available + # Note: Current implementation allows reserved_balance to go negative + await revert_pay_for_request(test_key, integration_session, 100) + + # Refresh to get updated values + await integration_session.refresh(test_key) + + # Current implementation allows negative reserved balance + assert test_key.reserved_balance == -100, ( + f"Expected reserved_balance to be -100, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == -1, ( + f"Expected total_requests to be -1, got: {test_key.total_requests}" + ) diff --git a/tests/integration/test_wallet_authentication.py b/tests/integration/test_wallet_authentication.py index 504308dc..21a26f3b 100644 --- a/tests/integration/test_wallet_authentication.py +++ b/tests/integration/test_wallet_authentication.py @@ -11,7 +11,7 @@ import pytest from httpx import AsyncClient from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey from .utils import ( CashuTokenGenerator, diff --git a/tests/integration/test_wallet_information.py b/tests/integration/test_wallet_information.py index bae84cb4..6f24e895 100644 --- a/tests/integration/test_wallet_information.py +++ b/tests/integration/test_wallet_information.py @@ -11,7 +11,7 @@ import pytest from httpx import AsyncClient from sqlmodel import select, update -from router.core.db import ApiKey +from routstr.core.db import ApiKey from .utils import ConcurrencyTester, ResponseValidator diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 17b4f965..c2461cae 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -7,14 +7,13 @@ import asyncio import base64 import json from typing import Any -from unittest.mock import AsyncMock, patch +from unittest.mock import patch import pytest from httpx import AsyncClient from sqlmodel import select -from router.core.db import ApiKey -from router.wallet import CurrencyUnit +from routstr.core.db import ApiKey @pytest.mark.integration @@ -42,14 +41,18 @@ async def test_full_balance_refund_returns_cashu_token( assert response.status_code == 200 data = response.json() - # Should return msats, recipient (None), and token - assert "msats" in data - assert "recipient" in data + # Should return either sats or msats (as string), and token assert "token" in data - assert data["msats"] == initial_balance - assert data["recipient"] is None assert data["token"].startswith("cashuA") + # Check for either sats or msats depending on refund_currency + if "sats" in data: + assert data["sats"] == str(initial_balance // 1000) # Convert msats to sats + elif "msats" in data: + assert data["msats"] == str(initial_balance) + else: + pytest.fail("Response should contain either 'sats' or 'msats'") + # Validate token format token = data["token"] try: @@ -89,7 +92,14 @@ async def test_partial_refund_not_supported( # Should still refund full balance (endpoint ignores the parameter) assert response.status_code == 200 data = response.json() - assert data["msats"] == 10_000_000 # Full balance + + # Check for either sats or msats + if "sats" in data: + assert data["sats"] == "10000" # Full balance in sats + elif "msats" in data: + assert data["msats"] == "10000000" # Full balance in msats + else: + pytest.fail("Response should contain either 'sats' or 'msats'") @pytest.mark.integration @@ -154,25 +164,10 @@ async def test_refund_amount_validation( key = result.scalar_one() assert key.refund_address is None - # Set balance to less than 1 sat (999 msats) - from sqlmodel import update - - await integration_session.execute( - update(ApiKey) - .where(ApiKey.hashed_key == hashed_key) # type: ignore[arg-type] - .values(balance=999) # Less than 1 sat - ) - await integration_session.commit() - - # Try to refund - should fail - response = await authenticated_client.post("/v1/wallet/refund") - - assert response.status_code == 400 - assert "too small to refund" in response.json()["detail"].lower() - @pytest.mark.integration @pytest.mark.asyncio +@pytest.mark.skip(reason="Lightning address refund functionality not implemented") async def test_refund_with_lightning_address( integration_client: AsyncClient, testmint_wallet: Any, @@ -207,12 +202,12 @@ async def test_refund_with_lightning_address( await db_snapshot.capture() # Mock send_to_lnurl function directly - with patch("router.balance.send_to_lnurl") as mock_send_to_lnurl: + with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: mock_send_to_lnurl.return_value = { "amount_sent": balance, "unit": "msat", "lnurl": refund_address, - "status": "completed" + "status": "completed", } # Request refund @@ -230,7 +225,7 @@ async def test_refund_with_lightning_address( # Verify send_to_lnurl was called with correct parameters mock_send_to_lnurl.assert_called_once_with( balance, # amount in msats - CurrencyUnit.msat, # unit + "msat", # unit refund_address, # lnurl ) @@ -411,7 +406,7 @@ async def test_mint_unavailability_handling( # Make the send_token method raise an exception with patch( - "router.balance.send_token", + "routstr.balance.send_token", side_effect=Exception("Mint unavailable: Connection refused"), ): # The exception should propagate as a 503 error (Service Unavailable) @@ -448,13 +443,17 @@ async def test_refund_response_format( data = response.json() assert isinstance(data, dict) - assert "msats" in data - assert "recipient" in data assert "token" in data - assert isinstance(data["msats"], int) - assert data["recipient"] is None assert isinstance(data["token"], str) + # Should have either sats or msats (both as strings) + if "sats" in data: + assert isinstance(data["sats"], str) + elif "msats" in data: + assert isinstance(data["msats"], str) + else: + pytest.fail("Response should contain either 'sats' or 'msats'") + # Test 2: Test with refund address would require creating key via proxy endpoint # Since refund address headers only work on proxy endpoints, not wallet endpoints # Skip this part as it's already tested in test_refund_with_lightning_address @@ -490,14 +489,8 @@ async def test_refund_error_handling( integration_client.headers["Authorization"] = f"Bearer {api_key}" response = await integration_client.post("/v1/wallet/refund") - # With negative balance, the endpoint will return "No balance to refund" - # since the balance check is remaining_balance_msats == 0 - # but with -1000, it's not 0, so it proceeds - # For a negative balance without refund address, it would fail when converting to sats - # But with our current implementation it returns 200 with a token - # This is actually a bug in the implementation - negative balances should be rejected - # For now, accept the current behavior - assert response.status_code == 200 + assert response.status_code == 400 + assert response.json()["detail"] == "No balance to refund" @pytest.mark.integration @@ -508,10 +501,10 @@ async def test_refund_with_expired_key( """Test refunding an expired API key""" # Create expired key - from datetime import datetime, timedelta + from datetime import datetime, timedelta, timezone token = await testmint_wallet.mint_tokens(500) - past_expiry = int((datetime.utcnow() - timedelta(hours=1)).timestamp()) + past_expiry = int((datetime.now(timezone.utc) - timedelta(hours=1)).timestamp()) # Use cashu token as Bearer auth to create API key integration_client.headers["Authorization"] = f"Bearer {token}" @@ -534,10 +527,8 @@ async def test_refund_with_expired_key( integration_client.headers["Authorization"] = f"Bearer {api_key}" # Mock the refund to LN address - with patch("router.wallet.send_token") as mock_wallet_func: - mock_wallet = AsyncMock() - mock_wallet.send_to_lnurl = AsyncMock(return_value=500) # type: ignore[method-assign] - mock_wallet_func.return_value = mock_wallet + with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: + mock_send_to_lnurl.return_value = 500 response = await integration_client.post("/v1/wallet/refund") diff --git a/tests/integration/test_wallet_topup.py b/tests/integration/test_wallet_topup.py index 90befdfd..6f4d5c29 100644 --- a/tests/integration/test_wallet_topup.py +++ b/tests/integration/test_wallet_topup.py @@ -11,7 +11,7 @@ import pytest from httpx import AsyncClient from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey from .utils import ( CashuTokenGenerator, @@ -425,7 +425,7 @@ async def test_network_failure_during_token_verification( # type: ignore[no-unt token = await testmint_wallet.mint_tokens(300) # Mock credit_balance to simulate network failure during token verification - with patch("router.balance.credit_balance") as mock_credit_balance: + with patch("routstr.balance.credit_balance") as mock_credit_balance: mock_credit_balance.side_effect = Exception("Network error: Connection timeout") response = await authenticated_client.post( @@ -470,7 +470,11 @@ async def test_topup_with_zero_amount_token( # type: ignore[no-untyped-def] # Create a token with 0 amount (edge case) # The testmint wallet should handle this - with patch.object(testmint_wallet, "redeem_token", return_value=(0, "sat", testmint_wallet.mint_url)): + with patch.object( + testmint_wallet, + "redeem_token", + return_value=(0, "sat", testmint_wallet.mint_url), + ): token = await testmint_wallet.mint_tokens(0) response = await authenticated_client.post( diff --git a/tests/integration/utils.py b/tests/integration/utils.py index 6adcd8d3..dd1fcf3d 100644 --- a/tests/integration/utils.py +++ b/tests/integration/utils.py @@ -9,7 +9,7 @@ import httpx from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import select -from router.core.db import ApiKey +from routstr.core.db import ApiKey class CashuTokenGenerator: diff --git a/tests/integration/verify_setup.py b/tests/integration/verify_setup.py index c828fe58..992da14c 100644 --- a/tests/integration/verify_setup.py +++ b/tests/integration/verify_setup.py @@ -43,8 +43,8 @@ def check_imports() -> bool: print("Conftest fixtures imported successfully") - # Check router modules - imports are for verification only - from router.core.db import ApiKey + # Check routstr modules - imports are for verification only + from routstr.core.db import ApiKey del ApiKey diff --git a/tests/unit/README.md b/tests/unit/README.md index 72280342..914eb173 100644 --- a/tests/unit/README.md +++ b/tests/unit/README.md @@ -21,7 +21,7 @@ pytest To run tests with coverage: ```bash -pytest --cov=router --cov-report=html +pytest --cov=routstr --cov-report=html ``` To run specific test files: diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index c167486b..da63895c 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -5,7 +5,7 @@ from unittest.mock import Mock, patch os.environ["UPSTREAM_BASE_URL"] = "http://test" os.environ["UPSTREAM_API_KEY"] = "test" -from router.payment.helpers import get_max_cost_for_model # noqa: E402 +from routstr.payment.helpers import get_max_cost_for_model # noqa: E402 def test_get_max_cost_for_model_known() -> None: @@ -14,21 +14,33 @@ def test_get_max_cost_for_model_known() -> None: mock_model.sats_pricing = Mock() mock_model.sats_pricing.max_cost = 500 - with patch("router.payment.helpers.MODELS", [mock_model]): - with patch("router.payment.helpers.MODEL_BASED_PRICING", True): - cost = get_max_cost_for_model("gpt-4") + with patch("routstr.payment.helpers.MODELS", [mock_model]): + with patch("routstr.payment.helpers.MODEL_BASED_PRICING", True): + cost = get_max_cost_for_model("gpt-4", tolerance_percentage=0) assert cost == 500000 # 500 sats * 1000 = msats def test_get_max_cost_for_model_unknown() -> None: - with patch("router.payment.helpers.MODELS", []): - with patch("router.payment.helpers.COST_PER_REQUEST", 100): - cost = get_max_cost_for_model("unknown-model") + with patch("routstr.payment.helpers.MODELS", []): + with patch("routstr.payment.helpers.COST_PER_REQUEST", 100): + cost = get_max_cost_for_model("unknown-model", tolerance_percentage=0) assert cost == 100 def test_get_max_cost_for_model_disabled() -> None: - with patch("router.payment.helpers.MODEL_BASED_PRICING", False): - with patch("router.payment.helpers.COST_PER_REQUEST", 200): - cost = get_max_cost_for_model("any-model") + with patch("routstr.payment.helpers.MODEL_BASED_PRICING", False): + with patch("routstr.payment.helpers.COST_PER_REQUEST", 200): + cost = get_max_cost_for_model("any-model", tolerance_percentage=0) assert cost == 200 + + +def test_get_max_cost_for_model_tolerance() -> None: + mock_model = Mock() + mock_model.id = "gpt-4" + mock_model.sats_pricing = Mock() + mock_model.sats_pricing.max_cost = 500 + + with patch("routstr.payment.helpers.MODELS", [mock_model]): + with patch("routstr.payment.helpers.MODEL_BASED_PRICING", True): + cost = get_max_cost_for_model("gpt-4", tolerance_percentage=10) + assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000 diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index fb5425e7..b432cd9f 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -4,16 +4,17 @@ from unittest.mock import AsyncMock, Mock, patch import pytest -from router.wallet import credit_balance, get_balance, recieve_token, send_token +from routstr.wallet import credit_balance, get_balance, recieve_token, send_token @pytest.mark.asyncio async def test_get_balance() -> None: mock_wallet = Mock() mock_wallet.available_balance = Mock(amount=50000) + mock_wallet.load_mint = AsyncMock() mock_wallet.load_proofs = AsyncMock() - with patch("router.wallet.Wallet.with_db", return_value=mock_wallet): + with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): balance = await get_balance("sat") assert balance == 50000 @@ -36,10 +37,10 @@ async def test_recieve_token_valid() -> None: token_str = f"cashuA{token_b64}" mock_wallet = Mock() - mock_wallet.redeem = AsyncMock() + mock_wallet.split = AsyncMock() - with patch("router.wallet.TRUSTED_MINTS", ["http://mint:3338"]): - with patch("router.wallet.deserialize_token_from_string") as mock_deserialize: + with patch("routstr.wallet.TRUSTED_MINTS", ["http://mint:3338"]): + with patch("routstr.wallet.deserialize_token_from_string") as mock_deserialize: mock_token = Mock() mock_token.keysets = ["keyset1"] mock_token.mint = "http://mint:3338" @@ -48,9 +49,9 @@ async def test_recieve_token_valid() -> None: mock_token.proofs = [{"amount": 1000}] mock_deserialize.return_value = mock_token - with patch("router.wallet.Wallet.with_db", return_value=mock_wallet): - mock_wallet.load_mint = AsyncMock() - + mock_wallet.load_mint = AsyncMock() + mock_wallet.load_proofs = AsyncMock() + with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): amount, unit, mint = await recieve_token(token_str) assert amount == 1000 assert unit == "sat" @@ -61,8 +62,8 @@ async def test_recieve_token_valid() -> None: async def test_send_token() -> None: mock_wallet = Mock() - with patch("router.wallet.Wallet.with_db", return_value=mock_wallet): - with patch("router.wallet.send", return_value=(1000, "test_token")): + with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): + with patch("routstr.wallet.send", return_value=(1000, "test_token")): token = await send_token(1000, "sat", "http://mint:3338") assert token == "test_token" @@ -81,9 +82,9 @@ async def test_credit_balance() -> None: mock_key.balance = 5000000 mock_session = AsyncMock() - with patch("router.wallet.PRIMARY_MINT_URL", "http://mint:3338"): + with patch("routstr.wallet.PRIMARY_MINT_URL", "http://mint:3338"): with patch( - "router.wallet.recieve_token", + "routstr.wallet.recieve_token", return_value=(1000, "sat", "http://mint:3338"), ): amount = await credit_balance(token_str, mock_key, mock_session) @@ -93,23 +94,11 @@ async def test_credit_balance() -> None: mock_session.commit.assert_called_once() -@pytest.mark.asyncio -async def test_credit_balance_invalid_mint() -> None: - mock_key = Mock() - mock_session = AsyncMock() - - with patch( - "router.wallet.recieve_token", return_value=(1000, "sat", "http://other:3338") - ): - with pytest.raises(ValueError, match="Mint URL is not supported"): - await credit_balance("test_token", mock_key, mock_session) - - @pytest.mark.asyncio async def test_recieve_token_untrusted_mint() -> None: mock_wallet = Mock() - with patch("router.wallet.deserialize_token_from_string") as mock_deserialize: + with patch("routstr.wallet.deserialize_token_from_string") as mock_deserialize: mock_token = Mock() mock_token.keysets = ["keyset1"] mock_token.mint = "http://untrusted:3338" @@ -117,10 +106,11 @@ async def test_recieve_token_untrusted_mint() -> None: mock_token.amount = 1000 mock_deserialize.return_value = mock_token - with patch("router.wallet.Wallet.with_db", return_value=mock_wallet): - mock_wallet.load_mint = AsyncMock() + mock_wallet.load_mint = AsyncMock() + mock_wallet.load_proofs = AsyncMock() + with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): with patch( - "router.wallet.swap_to_primary_mint", + "routstr.wallet.swap_to_primary_mint", return_value=(900, "sat", "http://mint:3338"), ): amount, unit, mint = await recieve_token("test_token") diff --git a/uv.lock b/uv.lock index 642223ea..c6f2bacd 100644 --- a/uv.lock +++ b/uv.lock @@ -1767,8 +1767,8 @@ wheels = [ [[package]] name = "routstr" -version = "0.1.0" -source = { virtual = "." } +version = "0.1.1b" +source = { editable = "." } dependencies = [ { name = "aiosqlite" }, { name = "alembic" }, @@ -1822,7 +1822,7 @@ dev = [ { name = "pytest-asyncio", specifier = ">=0.24.0" }, { name = "pytest-benchmark", specifier = ">=4.0.0" }, { name = "pytest-cov", specifier = ">=6.1.1" }, - { name = "routstr", virtual = "." }, + { name = "routstr", editable = "." }, { name = "ruff", specifier = ">=0.11.6" }, ]