diff --git a/example.py b/example.py index ea0a0dc2..d62d959d 100644 --- a/example.py +++ b/example.py @@ -1,4 +1,5 @@ import os + import openai client = openai.OpenAI( diff --git a/pyproject.toml b/pyproject.toml index b8a43ace..0ee6834c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -44,3 +44,17 @@ markers = [ "integration: marks tests as integration tests", "unit: marks tests as unit tests", ] + +[tool.ruff.lint] +exclude = ["gptassistant/bot/prompts.py"] +select = ["E", "F", "I"] +ignore = ["E501"] + +[tool.mypy] +python_version = "3.11" +ignore_missing_imports = true +disallow_untyped_defs = true +check_untyped_defs = true +disallow_untyped_calls = true +disallow_incomplete_defs = true +disallow_untyped_decorators = true diff --git a/router/account.py b/router/account.py index 27e391ce..f9fe9fa9 100644 --- a/router/account.py +++ b/router/account.py @@ -1,17 +1,19 @@ from typing import Annotated -from fastapi import APIRouter, Header, HTTPException, Depends + +from fastapi import APIRouter, Depends, Header, HTTPException from .auth import validate_bearer_key from .cashu import ( - refund_balance, - credit_balance, WALLET, + credit_balance, delete_key_if_zero_balance, + refund_balance, ) from .db import ApiKey, AsyncSession, get_session wallet_router = APIRouter(prefix="/v1/wallet") + async def get_key_from_header( authorization: Annotated[str, Header(...)], session: AsyncSession = Depends(get_session), @@ -24,6 +26,7 @@ async def get_key_from_header( detail="Invalid authorization. Use 'Bearer ' or 'Bearer '", ) + # TODO: remove this endpoint when frontend is updated @wallet_router.get("/") async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: @@ -32,6 +35,7 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: "balance": key.balance, } + @wallet_router.get("/info") async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: return { @@ -55,10 +59,10 @@ async def refund_wallet_endpoint( session: AsyncSession = Depends(get_session), ) -> dict: remaining_balance_msats = key.balance - + if remaining_balance_msats == 0: raise HTTPException(status_code=400, detail="No balance to refund") - + # Perform refund operation first, before modifying balance if key.refund_address: await refund_balance(remaining_balance_msats, key, session) @@ -67,17 +71,19 @@ async def refund_wallet_endpoint( # 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)") - + raise HTTPException( + status_code=400, detail="Balance too small to refund (less than 1 sat)" + ) + token = await WALLET.send(remaining_balance_sats) result = {"msats": remaining_balance_msats, "recipient": None, "token": token} - + # Only after successful refund, zero out the balance key.balance = 0 session.add(key) await session.commit() await delete_key_if_zero_balance(key, session) - + return result @@ -85,4 +91,6 @@ async def refund_wallet_endpoint( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], include_in_schema=False ) async def wallet_catch_all(path: str): - raise HTTPException(status_code=404, detail="Not found check /docs for available endpoints") + raise HTTPException( + status_code=404, detail="Not found check /docs for available endpoints" + ) diff --git a/router/admin.py b/router/admin.py index a4734de6..c46158dd 100644 --- a/router/admin.py +++ b/router/admin.py @@ -5,8 +5,8 @@ from fastapi import APIRouter, Request from fastapi.responses import HTMLResponse from sqlmodel import select -from .db import ApiKey, create_session from .cashu import WALLET +from .db import ApiKey, create_session admin_router = APIRouter(prefix="/admin") diff --git a/router/auth.py b/router/auth.py index 91b4381e..67ee1041 100644 --- a/router/auth.py +++ b/router/auth.py @@ -1,12 +1,11 @@ import asyncio import hashlib -import os import json +import os from typing import Optional - from fastapi import HTTPException, Request -from sqlmodel import update, col +from sqlmodel import col, update from .cashu import credit_balance, pay_out from .db import ApiKey, AsyncSession diff --git a/router/cashu.py b/router/cashu.py index 2e59a333..f30fea12 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -1,11 +1,11 @@ -import os import asyncio +import os import time from sixty_nuts import Wallet -from sqlmodel import select, func, col, update -from .db import ApiKey, AsyncSession, get_session +from sqlmodel import col, func, select, update +from .db import ApiKey, AsyncSession, get_session RECEIVE_LN_ADDRESS = os.environ["RECEIVE_LN_ADDRESS"] MINT = os.environ.get("MINT", "https://mint.minibits.cash/Bitcoin") diff --git a/router/db.py b/router/db.py index 10e45d8b..e959e6b4 100644 --- a/router/db.py +++ b/router/db.py @@ -1,10 +1,10 @@ -from contextlib import asynccontextmanager import os +from contextlib import asynccontextmanager from typing import AsyncGenerator -from sqlmodel import Field, SQLModel -from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlmodel.ext.asyncio.session import AsyncSession +from sqlalchemy.ext.asyncio.engine import create_async_engine +from sqlmodel import Field, SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") diff --git a/router/discovery.py b/router/discovery.py index 34461144..2d849b88 100644 --- a/router/discovery.py +++ b/router/discovery.py @@ -1,12 +1,13 @@ -from fastapi import APIRouter import asyncio import json -import websockets -import random -import string -import re -import httpx import os +import random +import re +import string + +import httpx +import websockets +from fastapi import APIRouter providers_router = APIRouter(prefix="/v1/providers") diff --git a/router/main.py b/router/main.py index a29fdb4e..bd0ce175 100644 --- a/router/main.py +++ b/router/main.py @@ -1,19 +1,21 @@ import asyncio -from contextlib import asynccontextmanager import os +from contextlib import asynccontextmanager + from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -from .db import init_db -from .admin import admin_router -from .proxy import proxy_router from .account import wallet_router -from .models import MODELS, update_sats_pricing -from .cashu import check_for_refunds, init_wallet, close_wallet +from .admin import admin_router +from .cashu import check_for_refunds, close_wallet, init_wallet +from .db import init_db from .discovery import providers_router +from .models import MODELS, update_sats_pricing +from .proxy import proxy_router __version__ = "0.0.1" + @asynccontextmanager async def lifespan(_: FastAPI): await init_db() diff --git a/router/models.py b/router/models.py index 3bad6925..60b2b93e 100644 --- a/router/models.py +++ b/router/models.py @@ -2,6 +2,7 @@ import asyncio import json import os from pathlib import Path + from pydantic.v1 import BaseModel from .price import sats_usd_ask_price diff --git a/router/price.py b/router/price.py index 885c9759..ab584122 100644 --- a/router/price.py +++ b/router/price.py @@ -1,7 +1,8 @@ -import os -import httpx import asyncio import logging +import os + +import httpx # artifical spread to cover conversion fees EXCHANGE_FEE = float(os.environ.get("EXCHANGE_FEE", "1.005")) # 0.5% default diff --git a/router/proxy.py b/router/proxy.py index 93b268df..f388a340 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -1,13 +1,13 @@ -import os import json -from fastapi import APIRouter, Request, BackgroundTasks, Depends -from fastapi.responses import Response, StreamingResponse -import httpx +import os import re -from .cashu import pay_out +import httpx +from fastapi import APIRouter, BackgroundTasks, Depends, Request +from fastapi.responses import Response, StreamingResponse -from .auth import validate_bearer_key, pay_for_request, adjust_payment_for_tokens +from .auth import adjust_payment_for_tokens, pay_for_request, validate_bearer_key +from .cashu import pay_out from .db import AsyncSession, get_session UPSTREAM_BASE_URL = os.environ["UPSTREAM_BASE_URL"] diff --git a/scripts/models_meta.py b/scripts/models_meta.py index 585becd2..f2871ef3 100644 --- a/scripts/models_meta.py +++ b/scripts/models_meta.py @@ -1,8 +1,9 @@ -import httpx -import json import asyncio +import json from typing import TypedDict +import httpx + class ModelArchitecture(TypedDict): modality: str diff --git a/tests/conftest.py b/tests/conftest.py index 09c39c62..0790e18b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,14 +1,15 @@ import asyncio import os +from typing import AsyncGenerator, Generator +from unittest.mock import AsyncMock, MagicMock, patch + import pytest import pytest_asyncio -from typing import AsyncGenerator, Generator from fastapi.testclient import TestClient -from httpx import AsyncClient, ASGITransport -from sqlmodel import SQLModel +from httpx import ASGITransport, AsyncClient from sqlalchemy.ext.asyncio import create_async_engine +from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession -from unittest.mock import patch, MagicMock, AsyncMock # Save original environment variables ORIGINAL_ENV = os.environ.copy() @@ -55,8 +56,8 @@ with patch("sixty_nuts.Wallet") as mock_wallet_class: # Make the Wallet class return our mock when instantiated mock_wallet_class.return_value = mock_wallet - from router.main import app from router.db import get_session + from router.main import app @pytest.fixture(scope="session") diff --git a/tests/test_account.py b/tests/test_account.py index 95b0713a..28ac6efa 100644 --- a/tests/test_account.py +++ b/tests/test_account.py @@ -1,9 +1,11 @@ -import pytest -import pytest_asyncio import hashlib import uuid -from unittest.mock import patch, AsyncMock +from unittest.mock import AsyncMock, patch + +import pytest +import pytest_asyncio from httpx import AsyncClient + from router.db import ApiKey, AsyncSession diff --git a/tests/test_main.py b/tests/test_main.py index d4d026ca..fc5b59d7 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -1,6 +1,7 @@ +from unittest.mock import patch + import pytest from httpx import AsyncClient -from unittest.mock import patch @pytest.mark.asyncio diff --git a/tests/test_models.py b/tests/test_models.py index 1fce1710..6476c6a9 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1,7 +1,16 @@ -import pytest import asyncio -from unittest.mock import patch, AsyncMock -from router.models import Model, Architecture, Pricing, TopProvider, update_sats_pricing, MODELS +from unittest.mock import AsyncMock, patch + +import pytest + +from router.models import ( + MODELS, + Architecture, + Model, + Pricing, + TopProvider, + update_sats_pricing, +) @pytest.fixture diff --git a/tests/test_proxy.py b/tests/test_proxy.py index 06f83d30..f7970302 100644 --- a/tests/test_proxy.py +++ b/tests/test_proxy.py @@ -1,10 +1,12 @@ -import pytest -import pytest_asyncio import json import os import uuid from unittest.mock import AsyncMock, patch + +import pytest +import pytest_asyncio from httpx import AsyncClient + from router.db import ApiKey, AsyncSession @@ -260,7 +262,7 @@ async def test_proxy_with_model_based_pricing( with patch.dict(os.environ, {"MODEL_BASED_PRICING": "true"}): with patch("os.path.exists", return_value=True): # Mock a model with pricing - from router.models import MODELS, Model, Pricing, Architecture, TopProvider + from router.models import MODELS, Architecture, Model, Pricing, TopProvider test_model = Model( id="gpt-4", diff --git a/tests/test_shutdown.py b/tests/test_shutdown.py index 81857686..09eabb58 100644 --- a/tests/test_shutdown.py +++ b/tests/test_shutdown.py @@ -2,8 +2,9 @@ import asyncio from unittest.mock import AsyncMock, MagicMock, patch import pytest -from tests.conftest import TEST_ENV + from router.main import app, lifespan +from tests.conftest import TEST_ENV @pytest.mark.asyncio