diff --git a/router/__init__.py b/router/__init__.py index 9c697e55..7a1bc151 100644 --- a/router/__init__.py +++ b/router/__init__.py @@ -2,6 +2,6 @@ import dotenv dotenv.load_dotenv() -from .main import app as fastapi_app # noqa +from .core.main import app as fastapi_app # noqa __all__ = ["fastapi_app"] diff --git a/router/auth.py b/router/auth.py index b51a069d..8c6cd04d 100644 --- a/router/auth.py +++ b/router/auth.py @@ -4,8 +4,8 @@ from typing import Optional from fastapi import HTTPException from sqlmodel import col, update -from .db import ApiKey, AsyncSession -from .logging import get_logger +from .core import get_logger +from .core.db import ApiKey, AsyncSession from .payment.cost_caculation import ( CostData, CostDataError, diff --git a/router/account.py b/router/balance.py similarity index 86% rename from router/account.py rename to router/balance.py index c4389a2a..0592b249 100644 --- a/router/account.py +++ b/router/balance.py @@ -3,10 +3,11 @@ from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException from .auth import validate_bearer_key -from .db import ApiKey, AsyncSession, get_session +from .core.db import ApiKey, AsyncSession, get_session from .wallet import credit_balance, send_to_lnurl, send_token -wallet_router = APIRouter(prefix="/v1/wallet") +router = APIRouter() +balance_router = APIRouter(prefix="/v1/balance") async def get_key_from_header( @@ -23,7 +24,7 @@ async def get_key_from_header( # TODO: remove this endpoint when frontend is updated -@wallet_router.get("/") +@router.get("/", include_in_schema=False) async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: return { "api_key": "sk-" + key.hashed_key, @@ -31,7 +32,7 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: } -@wallet_router.get("/info") +@router.get("/info") async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: return { "api_key": "sk-" + key.hashed_key, @@ -39,7 +40,7 @@ async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: } -@wallet_router.post("/topup") +@router.post("/topup") async def topup_wallet_endpoint( cashu_token: str, key: ApiKey = Depends(get_key_from_header), @@ -49,7 +50,7 @@ async def topup_wallet_endpoint( return {"msats": amount_msats} -@wallet_router.post("/refund") +@router.post("/refund") async def refund_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), @@ -82,7 +83,7 @@ async def refund_wallet_endpoint( return result -@wallet_router.api_route( +@router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], include_in_schema=False, @@ -92,3 +93,8 @@ async def wallet_catch_all(path: str) -> NoReturn: raise HTTPException( status_code=404, detail="Not found check /docs for available endpoints" ) + + +balance_router.include_router(router) +deprecated_wallet_router = APIRouter(prefix="/v1/wallet", include_in_schema=False) +deprecated_wallet_router.include_router(router) diff --git a/router/core/__init__.py b/router/core/__init__.py new file mode 100644 index 00000000..6affb142 --- /dev/null +++ b/router/core/__init__.py @@ -0,0 +1,3 @@ +from .logging import get_logger + +__all__ = ["get_logger"] diff --git a/router/admin.py b/router/core/admin.py similarity index 99% rename from router/admin.py rename to router/core/admin.py index 856bffa9..b511e449 100644 --- a/router/admin.py +++ b/router/core/admin.py @@ -6,10 +6,10 @@ 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 -from .wallet import get_balance, send_token -admin_router = APIRouter(prefix="/admin") +admin_router = APIRouter(prefix="/admin", include_in_schema=False) class WithdrawRequest(BaseModel): diff --git a/router/db.py b/router/core/db.py similarity index 100% rename from router/db.py rename to router/core/db.py diff --git a/router/logging.py b/router/core/logging.py similarity index 97% rename from router/logging.py rename to router/core/logging.py index 5a4b019c..57c5277a 100644 --- a/router/logging.py +++ b/router/core/logging.py @@ -252,6 +252,11 @@ def setup_logging() -> None: "handlers": ["console"] if console_enabled else [], "propagate": False, }, + "watchfiles": { + "level": "WARNING", + "handlers": ["console"] if console_enabled else [], + "propagate": False, + }, }, "root": { "level": log_level, diff --git a/router/main.py b/router/core/main.py similarity index 87% rename from router/main.py rename to router/core/main.py index 4562ee3d..592ccbfd 100644 --- a/router/main.py +++ b/router/core/main.py @@ -6,14 +6,14 @@ from typing import AsyncGenerator from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -from .account import wallet_router +from ..balance import balance_router, deprecated_wallet_router +from ..discovery import providers_router +from ..payment.models import MODELS, models_router, update_sats_pricing +from ..proxy import proxy_router +from ..wallet import check_for_refunds, periodic_payout from .admin import admin_router from .db import init_db -from .discovery import providers_router from .logging import get_logger, setup_logging -from .models import MODELS, models_router, update_sats_pricing -from .proxy import proxy_router -from .wallet import check_for_refunds, periodic_payout # Initialize logging first setup_logging() @@ -84,7 +84,8 @@ app.add_middleware( ) -@app.get("/") +@app.get("/", include_in_schema=False) +@app.get("/v1/info") async def info() -> dict: logger.info("Info endpoint accessed") return { @@ -101,7 +102,8 @@ async def info() -> dict: app.include_router(models_router) app.include_router(admin_router) -app.include_router(wallet_router) +app.include_router(balance_router) +app.include_router(deprecated_wallet_router) app.include_router(providers_router) app.include_router(proxy_router) diff --git a/router/payment/__init__.py b/router/payment/__init__.py index e69de29b..55f5a854 100644 --- a/router/payment/__init__.py +++ b/router/payment/__init__.py @@ -0,0 +1,8 @@ +from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost + +__all__ = [ + "CostData", + "CostDataError", + "MaxCostData", + "calculate_cost", +] diff --git a/router/payment/cost_caculation.py b/router/payment/cost_caculation.py index 2f2bf9da..fb664e9c 100644 --- a/router/payment/cost_caculation.py +++ b/router/payment/cost_caculation.py @@ -2,8 +2,8 @@ import os from pydantic import BaseModel -from ..logging import get_logger -from ..models import MODELS +from ..core import get_logger +from .models import MODELS logger = get_logger(__name__) diff --git a/router/payment/helpers.py b/router/payment/helpers.py index d04216d9..5c265ea5 100644 --- a/router/payment/helpers.py +++ b/router/payment/helpers.py @@ -6,9 +6,9 @@ from typing import Literal import cbor2 from fastapi import HTTPException, Response -from ..logging import get_logger -from ..models import MODELS +from ..core import get_logger from .cost_caculation import COST_PER_REQUEST, MODEL_BASED_PRICING +from .models import MODELS logger = get_logger(__name__) diff --git a/router/models.py b/router/payment/models.py similarity index 99% rename from router/models.py rename to router/payment/models.py index e9af0e85..ef545c81 100644 --- a/router/models.py +++ b/router/payment/models.py @@ -144,7 +144,6 @@ async def update_sats_pricing() -> None: break -@models_router.get("/models") @models_router.get("/v1/models") async def models() -> dict: return {"data": MODELS} diff --git a/router/price.py b/router/payment/price.py similarity index 99% rename from router/price.py rename to router/payment/price.py index a14557a6..588c8094 100644 --- a/router/price.py +++ b/router/payment/price.py @@ -3,7 +3,7 @@ import os import httpx -from .logging import get_logger +from ..core import get_logger logger = get_logger(__name__) diff --git a/router/payment/x_cashu.py b/router/payment/x_cashu.py index 30ce674c..2c25ec14 100644 --- a/router/payment/x_cashu.py +++ b/router/payment/x_cashu.py @@ -6,9 +6,14 @@ import httpx from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from ..logging import get_logger +from ..core import get_logger from ..wallet import CurrencyUnit, recieve_token, send_token -from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost +from .cost_caculation import ( + CostData, + CostDataError, + MaxCostData, + calculate_cost, +) from .helpers import ( UPSTREAM_BASE_URL, create_error_response, diff --git a/router/proxy.py b/router/proxy.py index 2b643cb9..d89afd8f 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -13,8 +13,8 @@ from .auth import ( revert_pay_for_request, validate_bearer_key, ) -from .db import ApiKey, AsyncSession, create_session, get_session -from .logging import get_logger +from .core import get_logger +from .core.db import ApiKey, AsyncSession, create_session, get_session from .payment.helpers import ( UPSTREAM_BASE_URL, check_token_balance, diff --git a/router/wallet.py b/router/wallet.py index 65fe211e..bfd3718b 100644 --- a/router/wallet.py +++ b/router/wallet.py @@ -1,12 +1,11 @@ import os from typing import Literal -from cashu.core.base import Token +from cashu.core.base import Token, Unit from cashu.wallet.helpers import deserialize_token_from_string, receive, send from cashu.wallet.wallet import Wallet -from .db import ApiKey, AsyncSession -from .logging import get_logger +from .core import db, get_logger logger = get_logger(__name__) @@ -39,10 +38,9 @@ async def recieve_token( load_all_keysets=True, unit=token_obj.unit, ) - if token_obj.mint in TRUSTED_MINTS and token_obj.mint != PRIMARY_MINT_URL: + + if token_obj.mint not in TRUSTED_MINTS: return await swap_to_primary_mint(token_obj, wallet) - elif token_obj.mint not in TRUSTED_MINTS: - raise ValueError("Mint URL is not supported by this proxy") await receive(wallet, token_obj) return token_obj.amount, token_obj.unit, token_obj.mint @@ -61,7 +59,7 @@ async def send_token( async def swap_to_primary_mint( - token_obj: Token, wallet: Wallet + token_obj: Token, token_wallet: Wallet ) -> tuple[int, CurrencyUnit, str]: print(f"swap_to_primary_mint, token_obj: {token_obj}") if token_obj.unit == "sat": @@ -72,18 +70,33 @@ async def swap_to_primary_mint( raise ValueError("Invalid unit") estimated_fee_sat = max(amount_msat // 1000 * 0.01, 2) amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000 - mint_quote = await wallet.mint_quote(amount_msat_after_fee, "sat") - melt_quote = await wallet.melt_quote(mint_quote.request, amount_msat_after_fee) - _ = await wallet.melt( + print(f"amount_msat_after_fee: {amount_msat_after_fee}") + primary_wallet = await Wallet.with_db( + PRIMARY_MINT_URL, db=".temp", load_all_keysets=True, unit="sat" + ) + await primary_wallet.load_mint_keysets() + mint_quote = await primary_wallet.mint_quote( + amount_msat_after_fee // 1000, Unit.sat + ) + print(f"mint_quote: {mint_quote}") + melt_quote = await token_wallet.melt_quote(mint_quote.request) + print(f"melt_quote: {melt_quote}") + melt_quote_resp = await token_wallet.melt( proofs=token_obj.proofs, invoice=mint_quote.request, - fee_reserve=melt_quote.fee_reserve, + fee_reserve_sat=melt_quote.fee_reserve, quote_id=melt_quote.quote, ) + print(f"melt_quote_resp: {melt_quote_resp}") + + _ = await primary_wallet.mint(amount_msat_after_fee // 1000, mint_quote.quote) + + return amount_msat_after_fee // 1000, "sat", PRIMARY_MINT_URL - -async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -> int: +async def credit_balance( + cashu_token: str, key: db.ApiKey, session: db.AsyncSession +) -> int: amount, unit, mint_url = await recieve_token(cashu_token) if unit == "sat": amount = amount * 1000 diff --git a/tests/README.md b/tests/README.md index 9e40159f..72280342 100644 --- a/tests/README.md +++ b/tests/README.md @@ -28,9 +28,8 @@ To run specific test files: ```bash pytest tests/test_main.py -pytest tests/test_account.py -pytest tests/test_proxy.py pytest tests/test_models.py +pytest tests/test_proxy.py ``` To run only async tests: diff --git a/tests/conftest.py b/tests/conftest.py index d69640c2..d5c2b6e6 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -37,8 +37,8 @@ TEST_ENV = { os.environ.update(TEST_ENV) # Now import modules that depend on environment variables -from router.db import get_session # noqa: E402 -from router.main import app # noqa: E402 +from router.core.db import get_session # noqa: E402 +from router.core.main import app # noqa: E402 @pytest.fixture(scope="session") @@ -79,7 +79,7 @@ async def test_session(test_engine: AsyncEngine) -> AsyncGenerator[AsyncSession, def test_client() -> Generator[TestClient, None, None]: """Create a test client for the FastAPI app.""" with patch.dict(os.environ, TEST_ENV, clear=True): - with patch("router.models.update_sats_pricing") as mock_update: + with patch("router.payment.models.update_sats_pricing") as mock_update: mock_update.return_value = None yield TestClient(app) @@ -95,7 +95,7 @@ async def async_client(test_session: AsyncSession) -> AsyncGenerator[AsyncClient # Mock startup tasks with patch.dict(os.environ, TEST_ENV, clear=True): - with patch("router.models.update_sats_pricing") as mock_update: + with patch("router.payment.models.update_sats_pricing") as mock_update: mock_update.return_value = None async with AsyncClient( diff --git a/tests/test_models.py b/tests/test_models.py index 06bb4eb8..f586fd93 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch import pytest -from router.models import ( +from router.payment.models import ( MODELS, Architecture, Model, @@ -49,7 +49,7 @@ async def test_update_sats_pricing_calculation(sample_model: Model) -> None: """Test that sats pricing is calculated correctly.""" # Mock the sats_usd_ask_price function with patch( - "router.models.sats_usd_ask_price", new_callable=AsyncMock + "router.payment.models.sats_usd_ask_price", new_callable=AsyncMock ) as mock_price: mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD @@ -137,7 +137,7 @@ async def test_update_sats_pricing_without_top_provider() -> None: ) with patch( - "router.models.sats_usd_ask_price", new_callable=AsyncMock + "router.payment.models.sats_usd_ask_price", new_callable=AsyncMock ) as mock_price: mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD @@ -189,7 +189,7 @@ async def test_update_sats_pricing_without_top_provider() -> None: async def test_update_sats_pricing_handles_errors() -> None: """Test that update_sats_pricing handles errors gracefully.""" with patch( - "router.models.sats_usd_ask_price", new_callable=AsyncMock + "router.payment.models.sats_usd_ask_price", new_callable=AsyncMock ) as mock_price: mock_price.side_effect = Exception("API Error")