diff --git a/example.py b/example.py index d62d959d..4732859d 100644 --- a/example.py +++ b/example.py @@ -13,7 +13,7 @@ client = openai.OpenAI( history: list = [] -def chat(): +def chat() -> None: while True: user_msg = {"role": "user", "content": input("\nYou: ")} history.append(user_msg) @@ -25,8 +25,10 @@ def chat(): stream=True, ): if len(chunk.choices) > 0: - ai_msg["content"] += chunk.choices[0].delta.content - print(chunk.choices[0].delta.content, end="", flush=True) + content = chunk.choices[0].delta.content + if content is not None: + ai_msg["content"] += content + print(content, end="", flush=True) print() history.append(ai_msg) diff --git a/router/account.py b/router/account.py index f9fe9fa9..a0d11db5 100644 --- a/router/account.py +++ b/router/account.py @@ -1,4 +1,4 @@ -from typing import Annotated +from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException @@ -49,8 +49,9 @@ async def topup_wallet_endpoint( cashu_token: str, key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), -): - return await credit_balance(cashu_token, key, session) +) -> dict[str, int]: + amount_msats = await credit_balance(cashu_token, key, session) + return {"msats": amount_msats} @wallet_router.post("/refund") @@ -88,9 +89,12 @@ async def refund_wallet_endpoint( @wallet_router.api_route( - "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], include_in_schema=False + "/{path:path}", + methods=["GET", "POST", "PUT", "DELETE"], + include_in_schema=False, + response_model=None, ) -async def wallet_catch_all(path: str): +async def wallet_catch_all(path: str) -> NoReturn: raise HTTPException( status_code=404, detail="Not found check /docs for available endpoints" ) diff --git a/router/admin.py b/router/admin.py index c46158dd..d44dba8a 100644 --- a/router/admin.py +++ b/router/admin.py @@ -155,7 +155,7 @@ async def dashboard(request: Request) -> str: @admin_router.get("/", response_class=HTMLResponse) -async def admin(request: Request): +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) diff --git a/router/cashu.py b/router/cashu.py index f30fea12..35ab6be3 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -25,12 +25,12 @@ async def delete_key_if_zero_balance(key: ApiKey, session: AsyncSession) -> None await session.commit() -async def init_wallet(): +async def init_wallet() -> None: global WALLET WALLET = await Wallet.create(nsec=NSEC, mint_urls=[MINT]) -async def close_wallet(): +async def close_wallet() -> None: global WALLET await WALLET.aclose() diff --git a/router/discovery.py b/router/discovery.py index 2d849b88..2a301540 100644 --- a/router/discovery.py +++ b/router/discovery.py @@ -139,7 +139,7 @@ async def fetch_onion(provider: str) -> dict: @providers_router.get("/") -async def get_providers(include_json: bool = False): +async def get_providers(include_json: bool = False) -> dict[str, list[dict | str]]: npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s" # Relays that support NIP-50 text search diff --git a/router/main.py b/router/main.py index bd0ce175..f96f6624 100644 --- a/router/main.py +++ b/router/main.py @@ -1,6 +1,7 @@ import asyncio import os from contextlib import asynccontextmanager +from typing import AsyncGenerator from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware @@ -17,7 +18,7 @@ __version__ = "0.0.1" @asynccontextmanager -async def lifespan(_: FastAPI): +async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: await init_db() await init_wallet() pricing_task = asyncio.create_task(update_sats_pricing()) @@ -51,7 +52,7 @@ app.add_middleware( @app.get("/") -async def info(): +async def info() -> dict: return { "name": app.title, "description": app.description, diff --git a/router/proxy.py b/router/proxy.py index f388a340..4ce77d55 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -1,6 +1,7 @@ import json import os import re +from typing import AsyncGenerator import httpx from fastapi import APIRouter, BackgroundTasks, Depends, Request @@ -16,12 +17,10 @@ UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "") proxy_router = APIRouter() -@proxy_router.api_route( - "/{path:path}", methods=["GET", "POST"] -) +@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) -): +) -> Response | StreamingResponse: auth = request.headers.get("Authorization", "") bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else "" refund_address = request.headers.get("Refund-LNURL", None) @@ -146,7 +145,7 @@ async def proxy( if is_streaming and response.status_code == 200: # Process streaming response and extract cost from the last chunk - async def stream_with_cost(): + async def stream_with_cost() -> AsyncGenerator[bytes, None]: # Store all chunks to analyze stored_chunks = [] diff --git a/tests/conftest.py b/tests/conftest.py index 0790e18b..cd661420 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -7,7 +7,7 @@ import pytest import pytest_asyncio from fastapi.testclient import TestClient from httpx import ASGITransport, AsyncClient -from sqlalchemy.ext.asyncio import create_async_engine +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession @@ -61,7 +61,7 @@ with patch("sixty_nuts.Wallet") as mock_wallet_class: @pytest.fixture(scope="session") -def event_loop(): +def event_loop() -> Generator[asyncio.AbstractEventLoop, None, None]: """Create an instance of the default event loop for the test session.""" loop = asyncio.get_event_loop_policy().new_event_loop() yield loop @@ -69,7 +69,7 @@ def event_loop(): @pytest_asyncio.fixture(scope="function") -async def test_engine(): +async def test_engine() -> AsyncGenerator[AsyncEngine, None]: """Create a test database engine - new for each test.""" engine = create_async_engine( "sqlite+aiosqlite:///:memory:", @@ -86,7 +86,7 @@ async def test_engine(): @pytest_asyncio.fixture -async def test_session(test_engine) -> AsyncGenerator[AsyncSession, None]: +async def test_session(test_engine: AsyncEngine) -> AsyncGenerator[AsyncSession, None]: """Create a test database session.""" from sqlmodel.ext.asyncio.session import AsyncSession as SqlModelAsyncSession @@ -123,10 +123,10 @@ def test_client() -> Generator[TestClient, None, None]: @pytest_asyncio.fixture -async def async_client(test_session) -> AsyncGenerator[AsyncClient, None]: +async def async_client(test_session: AsyncSession) -> AsyncGenerator[AsyncClient, None]: """Create an async test client with dependency overrides.""" - async def override_get_session(): + async def override_get_session() -> AsyncGenerator[AsyncSession, None]: yield test_session app.dependency_overrides[get_session] = override_get_session @@ -164,7 +164,7 @@ async def async_client(test_session) -> AsyncGenerator[AsyncClient, None]: @pytest.fixture -def mock_models(): +def mock_models() -> list[dict]: """Mock models data for testing.""" return [ { @@ -199,7 +199,7 @@ def mock_models(): # Cleanup after all tests @pytest.fixture(scope="session", autouse=True) -def cleanup(): +def cleanup() -> Generator[None, None, None]: yield # Restore original environment carefully current_keys = set(os.environ.keys()) diff --git a/tests/test_account.py b/tests/test_account.py index 28ac6efa..32f4eaeb 100644 --- a/tests/test_account.py +++ b/tests/test_account.py @@ -39,10 +39,11 @@ async def test_api_key(test_session: AsyncSession) -> ApiKey: @pytest.mark.asyncio async def test_account_info_with_valid_key( async_client: AsyncClient, test_api_key: ApiKey -): +) -> None: """Test getting account info with a valid API key.""" response = await async_client.get( - "/v1/wallet/info", headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"} + "/v1/wallet/info", + headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"}, ) assert response.status_code == 200 @@ -53,7 +54,7 @@ async def test_account_info_with_valid_key( @pytest.mark.asyncio -async def test_account_info_without_auth(async_client: AsyncClient): +async def test_account_info_without_auth(async_client: AsyncClient) -> None: """Test that account info requires authentication.""" response = await async_client.get("/v1/wallet/") @@ -61,7 +62,7 @@ async def test_account_info_without_auth(async_client: AsyncClient): @pytest.mark.asyncio -async def test_account_info_with_invalid_key(async_client: AsyncClient): +async def test_account_info_with_invalid_key(async_client: AsyncClient) -> None: """Test account info with an invalid API key.""" response = await async_client.get( "/v1/wallet/info", headers={"Authorization": "Bearer invalid-key"} @@ -73,7 +74,7 @@ async def test_account_info_with_invalid_key(async_client: AsyncClient): @pytest.mark.asyncio async def test_refund_balance_with_address( async_client: AsyncClient, test_api_key: ApiKey, test_session: AsyncSession -): +) -> None: """Test refunding balance when refund address is set.""" # Need to patch the refund_balance at the module level to intercept the call with patch("router.account.refund_balance", new_callable=AsyncMock) as mock_refund: @@ -101,7 +102,7 @@ async def test_refund_balance_with_address( @pytest.mark.asyncio async def test_refund_balance_without_address( async_client: AsyncClient, test_session: AsyncSession -): +) -> None: """Test refunding balance when no refund address is set.""" # Create key without refund address - with unique ID unique_id = str(uuid.uuid4())[:8] @@ -144,11 +145,11 @@ async def test_refund_balance_without_address( @pytest.mark.asyncio async def test_topup_balance_endpoint( async_client: AsyncClient, test_api_key: ApiKey, test_session: AsyncSession -): +) -> None: """Test topping up balance with a cashu token.""" # Mock at the router.account module level to intercept the import with patch("router.account.credit_balance", new_callable=AsyncMock) as mock_credit: - mock_credit.return_value = {"msats": 500000} + mock_credit.return_value = 500000 # Return integer msats value response = await async_client.post( "/v1/wallet/topup?cashu_token=cashuBqQSEQ...", @@ -166,7 +167,7 @@ async def test_topup_balance_endpoint( @pytest.mark.asyncio async def test_topup_balance_requires_cashu_token( async_client: AsyncClient, test_api_key: ApiKey -): +) -> None: """Test that topup endpoint requires a cashu token.""" response = await async_client.post( "/v1/wallet/topup", @@ -180,7 +181,7 @@ async def test_topup_balance_requires_cashu_token( @pytest.mark.asyncio async def test_account_with_cashu_token( async_client: AsyncClient, test_session: AsyncSession -): +) -> None: """Test authentication with a cashu token creates a new account.""" cashu_token = "cashuBqQSEQ123456" @@ -212,7 +213,7 @@ async def test_account_with_cashu_token( @pytest.mark.asyncio -async def test_account_with_invalid_cashu_token(async_client: AsyncClient): +async def test_account_with_invalid_cashu_token(async_client: AsyncClient) -> None: """Test authentication with an invalid cashu token returns 401.""" with patch("router.auth.credit_balance", new_callable=AsyncMock) as mock_credit: @@ -225,4 +226,3 @@ async def test_account_with_invalid_cashu_token(async_client: AsyncClient): assert response.status_code == 401 error = response.json() assert error["detail"]["error"]["code"] == "invalid_api_key" - diff --git a/tests/test_main.py b/tests/test_main.py index fc5b59d7..245e7303 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -1,11 +1,12 @@ from unittest.mock import patch import pytest +from fastapi.testclient import TestClient from httpx import AsyncClient @pytest.mark.asyncio -async def test_root_endpoint(async_client: AsyncClient): +async def test_root_endpoint(async_client: AsyncClient) -> None: """Test the root endpoint returns expected information.""" # Mock the environment variables for this specific test env_vars = { @@ -16,13 +17,13 @@ async def test_root_endpoint(async_client: AsyncClient): "HTTP_URL": "http://test.example.com", "ONION_URL": "http://test.onion", } - + with patch.dict("os.environ", env_vars, clear=False): response = await async_client.get("/") - + assert response.status_code == 200 data = response.json() - + # The app reads from env vars during import, so check what we actually get assert "name" in data assert "description" in data @@ -35,16 +36,16 @@ async def test_root_endpoint(async_client: AsyncClient): @pytest.mark.asyncio -async def test_cors_headers(async_client: AsyncClient): +async def test_cors_headers(async_client: AsyncClient) -> None: """Test that CORS headers are properly set.""" response = await async_client.options( "/", headers={ "Origin": "http://localhost:3000", "Access-Control-Request-Method": "GET", - } + }, ) - + assert response.status_code == 200 # Check that CORS is working (might be * or specific origin) assert "access-control-allow-origin" in response.headers @@ -52,7 +53,7 @@ async def test_cors_headers(async_client: AsyncClient): @pytest.mark.asyncio -async def test_startup_event_initializes_properly(test_client): +async def test_startup_event_initializes_properly(test_client: TestClient) -> None: """Test that the startup event runs without errors.""" # The test_client fixture already triggers the startup event # This test ensures no exceptions are raised during startup diff --git a/tests/test_models.py b/tests/test_models.py index 6476c6a9..99e14276 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1,6 +1,6 @@ import asyncio +from typing import Any from unittest.mock import AsyncMock, patch - import pytest from router.models import ( @@ -27,7 +27,7 @@ def sample_model() -> Model: input_modalities=["text"], output_modalities=["text"], tokenizer="test_tokenizer", - instruct_type="chat" + instruct_type="chat", ), pricing=Pricing( prompt=0.01, @@ -35,63 +35,71 @@ def sample_model() -> Model: request=0.001, image=0.0, web_search=0.0, - internal_reasoning=0.0 + internal_reasoning=0.0, ), top_provider=TopProvider( - context_length=4096, - max_completion_tokens=2048, - is_moderated=False - ) + context_length=4096, max_completion_tokens=2048, is_moderated=False + ), ) @pytest.mark.asyncio -async def test_update_sats_pricing_calculation(sample_model: Model): +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) as mock_price: + with patch( + "router.models.sats_usd_ask_price", new_callable=AsyncMock + ) as mock_price: mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD - + # Temporarily replace MODELS original_models = MODELS[:] MODELS.clear() MODELS.append(sample_model) - + # Run one iteration of the pricing update sleep_called = asyncio.Event() - - async def mock_sleep(duration): + + async def mock_sleep(duration: float) -> None: sleep_called.set() raise asyncio.CancelledError() - + with patch("asyncio.sleep", side_effect=mock_sleep): try: # Create and run the task task = asyncio.create_task(update_sats_pricing()) - + # Wait for the first iteration to complete await sleep_called.wait() - + # Check that sats pricing was calculated assert sample_model.sats_pricing is not None - + # Verify calculations (prices in USD / sats_to_usd) - assert sample_model.sats_pricing.prompt == pytest.approx(0.01 / 0.0001) # 100 sats - assert sample_model.sats_pricing.completion == pytest.approx(0.02 / 0.0001) # 200 sats - assert sample_model.sats_pricing.request == pytest.approx(0.001 / 0.0001) # 10 sats - + assert sample_model.sats_pricing.prompt == pytest.approx( + 0.01 / 0.0001 + ) # 100 sats + assert sample_model.sats_pricing.completion == pytest.approx( + 0.02 / 0.0001 + ) # 200 sats + assert sample_model.sats_pricing.request == pytest.approx( + 0.001 / 0.0001 + ) # 10 sats + # Verify max_cost calculation for model with top_provider expected_max_context = 4096 * sample_model.sats_pricing.prompt expected_max_completion = 2048 * sample_model.sats_pricing.completion - assert sample_model.sats_pricing.max_cost == pytest.approx(expected_max_context + expected_max_completion) - + assert sample_model.sats_pricing.max_cost == pytest.approx( + expected_max_context + expected_max_completion + ) + # Cancel and await the task task.cancel() try: await task except asyncio.CancelledError: pass - + except asyncio.CancelledError: pass finally: @@ -101,7 +109,7 @@ async def test_update_sats_pricing_calculation(sample_model: Model): @pytest.mark.asyncio -async def test_update_sats_pricing_without_top_provider(): +async def test_update_sats_pricing_without_top_provider() -> None: """Test sats pricing calculation for models without top_provider.""" model_without_top = Model( id="test-model-no-top", @@ -114,7 +122,7 @@ async def test_update_sats_pricing_without_top_provider(): input_modalities=["text"], output_modalities=["text"], tokenizer="test_tokenizer", - instruct_type=None + instruct_type=None, ), pricing=Pricing( prompt=0.01, @@ -122,31 +130,33 @@ async def test_update_sats_pricing_without_top_provider(): request=0.001, image=0.01, web_search=0.005, - internal_reasoning=0.015 + internal_reasoning=0.015, ), - top_provider=None + top_provider=None, ) - - with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price: + + with patch( + "router.models.sats_usd_ask_price", new_callable=AsyncMock + ) as mock_price: mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD - + original_models = MODELS[:] MODELS.clear() MODELS.append(model_without_top) - + sleep_called = asyncio.Event() - - async def mock_sleep(duration): + + async def mock_sleep(duration: float) -> None: sleep_called.set() raise asyncio.CancelledError() - + with patch("asyncio.sleep", side_effect=mock_sleep): try: task = asyncio.create_task(update_sats_pricing()) await sleep_called.wait() - + assert model_without_top.sats_pricing is not None - + # Verify the fallback max_cost calculation p = model_without_top.sats_pricing.prompt * 1_000_000 c = model_without_top.sats_pricing.completion * 32_000 @@ -155,16 +165,18 @@ async def test_update_sats_pricing_without_top_provider(): w = model_without_top.sats_pricing.web_search * 1000 ir = model_without_top.sats_pricing.internal_reasoning * 100 expected_max = p + c + r + i + w + ir - - assert model_without_top.sats_pricing.max_cost == pytest.approx(expected_max) - + + assert model_without_top.sats_pricing.max_cost == pytest.approx( + expected_max + ) + # Cancel and await the task task.cancel() try: await task except asyncio.CancelledError: pass - + except asyncio.CancelledError: pass finally: @@ -173,59 +185,61 @@ async def test_update_sats_pricing_without_top_provider(): @pytest.mark.asyncio -async def test_update_sats_pricing_handles_errors(): +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) as mock_price: + with patch( + "router.models.sats_usd_ask_price", new_callable=AsyncMock + ) as mock_price: mock_price.side_effect = Exception("API Error") - + error_printed = False original_print = print - - def mock_print(*args, **kwargs): + + def mock_print(*args: Any, **kwargs: Any) -> None: nonlocal error_printed message = " ".join(str(a) for a in args) if "API Error" in message and "Error updating sats pricing" in message: error_printed = True original_print(*args, **kwargs) - + with patch("builtins.print", side_effect=mock_print): sleep_called = asyncio.Event() - - async def mock_sleep(duration): + + async def mock_sleep(duration: float) -> None: sleep_called.set() raise asyncio.CancelledError() - + with patch("asyncio.sleep", side_effect=mock_sleep): try: task = asyncio.create_task(update_sats_pricing()) await sleep_called.wait() - + # Verify error was printed assert error_printed - + # Cancel and await the task task.cancel() try: await task except asyncio.CancelledError: pass - + except asyncio.CancelledError: pass -def test_model_serialization(sample_model: Model): +def test_model_serialization(sample_model: Model) -> None: """Test that models can be serialized and deserialized correctly.""" model_dict = sample_model.dict() - + # Verify all fields are present assert model_dict["id"] == "test-model" assert model_dict["name"] == "Test Model" assert model_dict["pricing"]["prompt"] == 0.01 assert model_dict["architecture"]["modality"] == "text" assert model_dict["top_provider"]["context_length"] == 4096 - + # Test deserialization new_model = Model(**model_dict) assert new_model.id == sample_model.id - assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt) + assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt) diff --git a/tests/test_proxy.py b/tests/test_proxy.py index f7970302..a07b18e1 100644 --- a/tests/test_proxy.py +++ b/tests/test_proxy.py @@ -1,6 +1,7 @@ import json import os import uuid +from typing import AsyncGenerator from unittest.mock import AsyncMock, patch import pytest @@ -28,7 +29,7 @@ async def api_key_with_balance(test_session: AsyncSession) -> ApiKey: @pytest.mark.asyncio -async def test_proxy_requires_authentication(async_client: AsyncClient): +async def test_proxy_requires_authentication(async_client: AsyncClient) -> None: """Test that proxy endpoints require authentication.""" response = await async_client.post("/v1/chat/completions") @@ -42,7 +43,7 @@ async def test_proxy_requires_authentication(async_client: AsyncClient): @pytest.mark.asyncio async def test_proxy_with_insufficient_balance( async_client: AsyncClient, test_session: AsyncSession -): +) -> None: """Test proxy request with insufficient balance.""" # Create key with minimal balance unique_id = str(uuid.uuid4())[:8] @@ -71,7 +72,7 @@ async def test_proxy_with_insufficient_balance( @pytest.mark.asyncio async def test_proxy_invalid_json_body( async_client: AsyncClient, api_key_with_balance: ApiKey -): +) -> None: """Test proxy request with invalid JSON body.""" response = await async_client.post( "/v1/chat/completions", @@ -91,7 +92,7 @@ async def test_proxy_invalid_json_body( @pytest.mark.asyncio async def test_proxy_successful_request_mock( async_client: AsyncClient, api_key_with_balance: ApiKey, test_session: AsyncSession -): +) -> None: """Test successful proxy request with mocked upstream.""" mock_response_data = { "id": "chatcmpl-123", @@ -162,7 +163,7 @@ async def test_proxy_successful_request_mock( @pytest.mark.asyncio async def test_proxy_streaming_response( async_client: AsyncClient, api_key_with_balance: ApiKey -): +) -> None: """Test proxy request with streaming response.""" # Mock SSE stream chunks stream_chunks = [ @@ -172,7 +173,7 @@ async def test_proxy_streaming_response( b"data: [DONE]\n\n", ] - async def mock_aiter_bytes(): + async def mock_aiter_bytes() -> AsyncGenerator[bytes, None]: for chunk in stream_chunks: yield chunk @@ -213,7 +214,7 @@ async def test_proxy_streaming_response( @pytest.mark.asyncio async def test_proxy_handles_upstream_errors( async_client: AsyncClient, api_key_with_balance: ApiKey -): +) -> None: """Test proxy handles upstream connection errors gracefully.""" with patch("httpx.AsyncClient") as mock_client_class: mock_client = AsyncMock() @@ -245,7 +246,7 @@ async def test_proxy_handles_upstream_errors( @pytest.mark.asyncio async def test_proxy_with_model_based_pricing( async_client: AsyncClient, test_session: AsyncSession -): +) -> None: """Test proxy with model-based pricing enabled.""" # Create API key with sufficient balance unique_id = str(uuid.uuid4())[:8] diff --git a/tests/test_shutdown.py b/tests/test_shutdown.py index 09eabb58..059c2454 100644 --- a/tests/test_shutdown.py +++ b/tests/test_shutdown.py @@ -8,11 +8,11 @@ from tests.conftest import TEST_ENV @pytest.mark.asyncio -async def test_background_tasks_cancel_on_shutdown(): +async def test_background_tasks_cancel_on_shutdown() -> None: pricing_started = asyncio.Event() pricing_cancelled = asyncio.Event() - async def fake_update(): + async def fake_update() -> None: pricing_started.set() try: await asyncio.Event().wait() @@ -23,7 +23,7 @@ async def test_background_tasks_cancel_on_shutdown(): refund_started = asyncio.Event() refund_cancelled = asyncio.Event() - async def fake_refund(): + async def fake_refund() -> None: refund_started.set() try: await asyncio.Event().wait() @@ -31,7 +31,7 @@ async def test_background_tasks_cancel_on_shutdown(): refund_cancelled.set() raise - with patch.dict('os.environ', TEST_ENV, clear=True): + with patch.dict("os.environ", TEST_ENV, clear=True): mock_wallet = AsyncMock() mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet) mock_wallet.__aexit__ = AsyncMock(return_value=None) @@ -40,13 +40,16 @@ async def test_background_tasks_cancel_on_shutdown(): mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state) mock_wallet.send_to_lnurl = AsyncMock(return_value=100) mock_wallet.redeem = AsyncMock(return_value=1) - mock_wallet.send = AsyncMock(return_value='cashu:token123') + mock_wallet.send = AsyncMock(return_value="cashu:token123") - with patch('router.cashu.Wallet.create', AsyncMock(return_value=mock_wallet)), \ - patch('router.cashu.WALLET', mock_wallet): - - with patch('router.main.update_sats_pricing', new=fake_update), \ - patch('router.main.check_for_refunds', new=fake_refund): + with ( + patch("router.cashu.Wallet.create", AsyncMock(return_value=mock_wallet)), + patch("router.cashu.WALLET", mock_wallet), + ): + with ( + patch("router.main.update_sats_pricing", new=fake_update), + patch("router.main.check_for_refunds", new=fake_refund), + ): async with lifespan(app): await pricing_started.wait() await refund_started.wait()