From eb5d832719c83f739eb6a3b1d64a0eecb4592d3e Mon Sep 17 00:00:00 2001 From: Shroominic Date: Wed, 6 Aug 2025 23:25:06 -0300 Subject: [PATCH] fix mypy + ruff linting --- router/core/main.py | 15 +++- tests/integration/conftest.py | 35 ++++----- tests/integration/test_background_tasks.py | 7 +- tests/integration/test_real_mint.py | 2 +- .../integration/test_wallet_authentication.py | 7 +- tests/integration/test_wallet_refund.py | 1 - tests/unit/test_payment_helpers.py | 12 +-- tests/unit/test_wallet.py | 77 +++++++++++-------- 8 files changed, 81 insertions(+), 75 deletions(-) diff --git a/router/core/main.py b/router/core/main.py index 2928f3f8..e7b77e91 100644 --- a/router/core/main.py +++ b/router/core/main.py @@ -46,11 +46,20 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: finally: logger.info("Application shutdown initiated") - pricing_task.cancel() - payout_task.cancel() + if pricing_task is not None: + pricing_task.cancel() + if payout_task is not None: + payout_task.cancel() try: - await asyncio.gather(pricing_task, payout_task, return_exceptions=True) + tasks_to_wait = [] + if pricing_task is not None: + tasks_to_wait.append(pricing_task) + if payout_task is not None: + tasks_to_wait.append(payout_task) + + if tasks_to_wait: + await asyncio.gather(*tasks_to_wait, return_exceptions=True) logger.info("Background tasks stopped successfully") except Exception as e: logger.error( diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 17d18cff..75378357 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -2,7 +2,7 @@ import asyncio import json import os from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock, patch import pytest import pytest_asyncio @@ -61,8 +61,8 @@ else: # Set test environment variables before importing the app os.environ.update(test_env) -from router.core.db import ApiKey, get_session -from router.core.main import app, lifespan +from router.core.db import ApiKey, get_session # noqa: E402 +from router.core.main import app, lifespan # noqa: E402 class TestmintWallet: @@ -113,9 +113,10 @@ class TestmintWallet: async def _create_real_token(self, amount: int) -> str: """Create real tokens using the testmint""" - from cashu.wallet.wallet import Wallet import tempfile + from cashu.wallet.wallet import Wallet + logger.info( f"Creating real token for {amount} sats from testmint {self.connection_url}" ) @@ -155,10 +156,10 @@ class TestmintWallet: async def _create_fallback_token(self, amount: int) -> str: """Fallback method to create a basic test token""" - import json import base64 - import time + import json import random + import time unique_id = int(time.time() * 1000000) + random.randint(1000, 9999) token_data = { @@ -198,7 +199,7 @@ class TestmintWallet: token_base64 = token[6:] # Remove "cashuA" prefix # Add padding if necessary padding = (4 - len(token_base64) % 4) % 4 - token_base64 += '=' * padding + token_base64 += "=" * padding token_json = base64.urlsafe_b64decode(token_base64).decode() token_data = json.loads(token_json) @@ -247,7 +248,9 @@ class TestmintWallet: # For testing, return a simulated balance return 100000 # 100k sats - async def credit_balance(self, cashu_token: str, key: ApiKey, session) -> int: + async def credit_balance( + self, cashu_token: str, key: ApiKey, session: AsyncSession + ) -> int: """Credit balance to API key - test implementation""" try: print(f"DEBUG: credit_balance called with token: {cashu_token[:20]}...") @@ -449,7 +452,6 @@ async def integration_app( if use_real_mint: # Use real mint with sixty_nuts wallet - from .real_testmint import create_real_mint_wallet # Use real mint - no wallet patches needed with patch("router.core.db.engine", integration_engine): @@ -660,15 +662,11 @@ async def background_tasks_controller() -> AsyncGenerator[Any, None]: # Patch background task functions to respect controller original_update_pricing: Optional[Callable] = None - original_check_refunds: Optional[Callable] = None original_periodic_payout: Optional[Callable] = None try: - from router.core.main import ( - check_for_refunds, - periodic_payout, - update_sats_pricing, - ) + from router.payment.models import update_sats_pricing + from router.wallet import periodic_payout async def controlled_update_pricing() -> None: while not controller.cancelled: @@ -676,12 +674,6 @@ async def background_tasks_controller() -> AsyncGenerator[Any, None]: await original_update_pricing() await asyncio.sleep(1) - async def controlled_check_refunds() -> None: - while not controller.cancelled: - if not controller.paused and original_check_refunds: - await original_check_refunds() - await asyncio.sleep(1) - async def controlled_periodic_payout() -> None: while not controller.cancelled: if not controller.paused and original_periodic_payout: @@ -690,7 +682,6 @@ async def background_tasks_controller() -> AsyncGenerator[Any, None]: # Store originals and patch original_update_pricing = update_sats_pricing - original_check_refunds = check_for_refunds original_periodic_payout = periodic_payout except ImportError: diff --git a/tests/integration/test_background_tasks.py b/tests/integration/test_background_tasks.py index cea1f67c..dcb67fba 100644 --- a/tests/integration/test_background_tasks.py +++ b/tests/integration/test_background_tasks.py @@ -10,8 +10,8 @@ from unittest.mock import AsyncMock, patch import pytest from router.core.db import ApiKey -from router.wallet import periodic_payout from router.payment.models import MODELS, Model, Pricing, update_sats_pricing +from router.wallet import periodic_payout @pytest.mark.asyncio @@ -456,7 +456,6 @@ class TestPeriodicPayoutTask: # Mock wallet balance higher than user balances (indicating revenue) wallet_balance = 200000 # 200 sats total - expected_revenue = wallet_balance - total_user_balance # 50 sats revenue with ( patch("router.wallet.get_balance", AsyncMock(return_value=wallet_balance)), @@ -704,9 +703,7 @@ class TestTaskInteractions: "router.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("refund")), patch("router.wallet.periodic_payout", lambda: task_with_cleanup("payout")), ): # Start all tasks diff --git a/tests/integration/test_real_mint.py b/tests/integration/test_real_mint.py index de217c80..9e48bf84 100644 --- a/tests/integration/test_real_mint.py +++ b/tests/integration/test_real_mint.py @@ -10,7 +10,7 @@ try: from .real_testmint import create_real_mint_wallet except ImportError: # sixty_nuts not available, tests will be skipped - create_real_mint_wallet = None + create_real_mint_wallet = None # type: ignore async def test_real_wallet() -> None: diff --git a/tests/integration/test_wallet_authentication.py b/tests/integration/test_wallet_authentication.py index c83b8670..504308dc 100644 --- a/tests/integration/test_wallet_authentication.py +++ b/tests/integration/test_wallet_authentication.py @@ -224,9 +224,10 @@ async def test_malformed_authorization_header(integration_client: AsyncClient) - response = await integration_client.get("/v1/wallet/") # Should return 401 for invalid auth (not 400 in this implementation) - assert response.status_code in [400, 401], ( - f"Malformed header '{auth_value[:20]}...' should fail" - ) + assert response.status_code in [ + 400, + 401, + ], f"Malformed header '{auth_value[:20]}...' should fail" @pytest.mark.integration diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 94d382fd..8731d5ee 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -404,7 +404,6 @@ async def test_mint_unavailability_handling( # The global mock in conftest.py is already in place, # so we need to temporarily modify it - import router.wallet from unittest.mock import patch # Make the send_token method raise an exception diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index b00cfbcb..c167486b 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -5,30 +5,30 @@ 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 +from router.payment.helpers import get_max_cost_for_model # noqa: E402 -def test_get_max_cost_for_model_known(): +def test_get_max_cost_for_model_known() -> None: mock_model = Mock() mock_model.id = "gpt-4" 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") assert cost == 500000 # 500 sats * 1000 = msats -def test_get_max_cost_for_model_unknown(): +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") assert cost == 100 -def test_get_max_cost_for_model_disabled(): +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") - assert cost == 200 \ No newline at end of file + assert cost == 200 diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 216f2cc9..fb5425e7 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -8,32 +8,36 @@ from router.wallet import credit_balance, get_balance, recieve_token, send_token @pytest.mark.asyncio -async def test_get_balance(): +async def test_get_balance() -> None: mock_wallet = Mock() mock_wallet.available_balance = Mock(amount=50000) mock_wallet.load_proofs = AsyncMock() - + with patch("router.wallet.Wallet.with_db", return_value=mock_wallet): balance = await get_balance("sat") assert balance == 50000 @pytest.mark.asyncio -async def test_recieve_token_valid(): +async def test_recieve_token_valid() -> None: token_data = { - "token": [{ - "mint": "http://mint:3338", - "proofs": [{"amount": 1000, "id": "test", "secret": "secret", "C": "curve"}] - }], - "unit": "sat" + "token": [ + { + "mint": "http://mint:3338", + "proofs": [ + {"amount": 1000, "id": "test", "secret": "secret", "C": "curve"} + ], + } + ], + "unit": "sat", } token_json = json.dumps(token_data) token_b64 = base64.urlsafe_b64encode(token_json.encode()).decode() token_str = f"cashuA{token_b64}" - + mock_wallet = Mock() mock_wallet.redeem = AsyncMock() - + with patch("router.wallet.TRUSTED_MINTS", ["http://mint:3338"]): with patch("router.wallet.deserialize_token_from_string") as mock_deserialize: mock_token = Mock() @@ -43,10 +47,10 @@ async def test_recieve_token_valid(): mock_token.amount = 1000 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() - + amount, unit, mint = await recieve_token(token_str) assert amount == 1000 assert unit == "sat" @@ -54,9 +58,9 @@ async def test_recieve_token_valid(): @pytest.mark.asyncio -async def test_send_token(): +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")): token = await send_token(1000, "sat", "http://mint:3338") @@ -64,24 +68,24 @@ async def test_send_token(): @pytest.mark.asyncio -async def test_credit_balance(): +async def test_credit_balance() -> None: token_data = { - "token": [{ - "mint": "http://mint:3338", - "proofs": [{"amount": 1000}] - }], - "unit": "sat" + "token": [{"mint": "http://mint:3338", "proofs": [{"amount": 1000}]}], + "unit": "sat", } token_json = json.dumps(token_data) token_b64 = base64.urlsafe_b64encode(token_json.encode()).decode() token_str = f"cashuA{token_b64}" - + mock_key = Mock() mock_key.balance = 5000000 mock_session = AsyncMock() - + with patch("router.wallet.PRIMARY_MINT_URL", "http://mint:3338"): - with patch("router.wallet.recieve_token", return_value=(1000, "sat", "http://mint:3338")): + with patch( + "router.wallet.recieve_token", + return_value=(1000, "sat", "http://mint:3338"), + ): amount = await credit_balance(token_str, mock_key, mock_session) assert amount == 1000000 # converted to msat assert mock_key.balance == 6000000 @@ -89,32 +93,37 @@ async def test_credit_balance(): mock_session.commit.assert_called_once() -@pytest.mark.asyncio -async def test_credit_balance_invalid_mint(): +@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 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(): +async def test_recieve_token_untrusted_mint() -> None: mock_wallet = Mock() - + with patch("router.wallet.deserialize_token_from_string") as mock_deserialize: mock_token = Mock() - mock_token.keysets = ["keyset1"] + mock_token.keysets = ["keyset1"] mock_token.mint = "http://untrusted:3338" mock_token.unit = "sat" 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() - with patch("router.wallet.swap_to_primary_mint", return_value=(900, "sat", "http://mint:3338")): + with patch( + "router.wallet.swap_to_primary_mint", + return_value=(900, "sat", "http://mint:3338"), + ): amount, unit, mint = await recieve_token("test_token") assert amount == 900 - assert unit == "sat" - assert mint == "http://mint:3338" \ No newline at end of file + assert unit == "sat" + assert mint == "http://mint:3338"