refactor names and folders

This commit is contained in:
Shroominic
2025-08-02 14:04:37 -03:00
parent d94a27d16e
commit 2a56998450
19 changed files with 92 additions and 52 deletions
+1 -1
View File
@@ -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"]
+2 -2
View File
@@ -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,
+13 -7
View File
@@ -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)
+3
View File
@@ -0,0 +1,3 @@
from .logging import get_logger
__all__ = ["get_logger"]
+2 -2
View File
@@ -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):
View File
@@ -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,
+9 -7
View File
@@ -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)
+8
View File
@@ -0,0 +1,8 @@
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
__all__ = [
"CostData",
"CostDataError",
"MaxCostData",
"calculate_cost",
]
+2 -2
View File
@@ -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__)
+2 -2
View File
@@ -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__)
@@ -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}
+1 -1
View File
@@ -3,7 +3,7 @@ import os
import httpx
from .logging import get_logger
from ..core import get_logger
logger = get_logger(__name__)
+7 -2
View File
@@ -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,
+2 -2
View File
@@ -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,
+26 -13
View File
@@ -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
+1 -2
View File
@@ -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:
+4 -4
View File
@@ -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(
+4 -4
View File
@@ -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")