diff --git a/router/balance.py b/router/balance.py index 8f9d6836..27709269 100644 --- a/router/balance.py +++ b/router/balance.py @@ -75,21 +75,33 @@ async def refund_wallet_endpoint( raise HTTPException(status_code=400, detail="No balance to refund") # Perform refund operation first, before modifying balance - if key.refund_address: - await send_to_lnurl(remaining_balance_msats, CurrencyUnit.msat, key.refund_address) - result = {"recipient": key.refund_address, "msat": remaining_balance_msats} - else: - # 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)" - ) + try: + if key.refund_address: + await send_to_lnurl(remaining_balance_msats, CurrencyUnit.msat, key.refund_address) + result = {"recipient": key.refund_address, "msats": remaining_balance_msats} + else: + # 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)" + ) - # TODO: choose currency and mint based on what user has configured - token = await send_token(remaining_balance_sats, "sat") + # TODO: choose currency and mint based on what user has configured + token = await send_token(remaining_balance_sats, "sat") - result = {"msats": remaining_balance_msats, "recipient": None, "token": token} + result = {"msats": remaining_balance_msats, "recipient": None, "token": token} + except HTTPException: + # Re-raise HTTP exceptions (like 400 for balance too small) + raise + except Exception as e: + # If refund fails, don't modify the database + error_msg = str(e) + if ("mint" in error_msg.lower() or "connection" in error_msg.lower() or + isinstance(e, Exception) and "ConnectError" in str(type(e))): + raise HTTPException(status_code=503, detail="Mint service unavailable") + else: + raise HTTPException(status_code=500, detail="Refund failed") await session.delete(key) await session.commit() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index b373ad06..58e2678e 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -276,10 +276,21 @@ class TestmintWallet: amount_msat = amount * 1000 # Assume tokens are in sats logger.info(f"TestmintWallet.credit_balance amount in msat: {amount_msat}") - # Credit the balance - key.balance += amount_msat - session.add(key) + # Credit the balance using atomic database update to prevent race conditions + from sqlmodel import update + + # Use atomic update to avoid lost update problem in concurrent scenarios + stmt = ( + update(ApiKey) + .where(ApiKey.hashed_key == key.hashed_key) + .values(balance=ApiKey.balance + amount_msat) + ) + await session.execute(stmt) await session.commit() + + # Refresh the key object to get the updated balance + await session.refresh(key) + logger.info(f"TestmintWallet.credit_balance successfully credited {amount_msat} msat") return amount_msat diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py index b6755934..b2cf6bde 100644 --- a/tests/integration/test_error_handling_edge_cases.py +++ b/tests/integration/test_error_handling_edge_cases.py @@ -34,11 +34,10 @@ class TestNetworkFailureScenarios: AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), ): - # Try to refund when mint is down - expect the error to propagate - with pytest.raises(ConnectError, match="Mint service unavailable"): - await authenticated_client.post( - "/v1/wallet/refund", json={"amount": 1000} - ) + # Try to refund when mint is down - should return 503 status + response = await authenticated_client.post("/v1/wallet/refund") + assert response.status_code == 503 + assert "Mint service unavailable" in response.json()["detail"] @pytest.mark.asyncio async def test_upstream_llm_service_down( diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index 875d6632..a2b0aff5 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -141,40 +141,40 @@ async def test_providers_data_structure_validation( ) -> None: """Test provider data structure contains expected fields""" - # Mock comprehensive provider data + # Mock RIP-02 provider announcement event mock_events: list[dict[str, Any]] = [ { "id": "event1", - "content": "Great provider: http://comprehensive-provider.onion", + "pubkey": "test_pubkey", "created_at": 1234567890, + "content": "Comprehensive provider announcement", + "tags": [ + ["d", "provider-123"], + ["endpoint", "https://api.provider.example/v1"], + ["name", "Comprehensive Provider"], + ["description", "A comprehensive AI provider"], + ["model", "gpt-3.5-turbo"], + ["model", "gpt-4"], + ] } ] - mock_provider_data = { - "id": "provider-123", - "name": "Comprehensive Provider", - "status": "online", - "models": [ - { - "id": "gpt-3.5-turbo", - "name": "GPT-3.5 Turbo", - "pricing": {"prompt": "0.0015", "completion": "0.002"}, - }, - { - "id": "gpt-4", - "name": "GPT-4", - "pricing": {"prompt": "0.03", "completion": "0.06"}, - }, - ], - "endpoint": "https://api.provider.example/v1", - "availability": "99.9%", + mock_health_response = { + "status_code": 200, + "endpoint": "models", + "json": { + "data": [ + {"id": "gpt-3.5-turbo", "object": "model"}, + {"id": "gpt-4", "object": "model"} + ] + } } with patch( "router.discovery.query_nostr_relay_for_providers", return_value=mock_events ): with patch("router.discovery.fetch_provider_health") as mock_fetch: - mock_fetch.return_value = {"status_code": 200, "json": mock_provider_data} + mock_fetch.return_value = mock_health_response response = await integration_client.get("/v1/providers/?include_json=true") assert response.status_code == 200 @@ -184,25 +184,24 @@ async def test_providers_data_structure_validation( # Validate that provider data contains expected fields assert len(providers) > 0 - for provider_dict in providers: - url = list(provider_dict.keys())[0] - provider_info = provider_dict[url] - - # Expected fields should be present - expected_fields = ["name", "status", "models"] + for provider_data in providers: + # Should have provider and health keys based on actual implementation + assert "provider" in provider_data + assert "health" in provider_data + + provider_info = provider_data["provider"] + # Expected fields from RIP-02 parser + expected_fields = ["id", "name", "endpoint_url", "supported_models"] for field in expected_fields: - if field in mock_provider_data: - assert field in provider_info + assert field in provider_info # Validate models structure if present - if "models" in provider_info: - models = provider_info["models"] - if isinstance(models, list) and len(models) > 0: - # If models is a list of dictionaries, validate structure - for model in models: - if isinstance(model, dict): - # Model should have id at minimum - assert "id" in model or "name" in model + if "supported_models" in provider_info: + models = provider_info["supported_models"] + assert isinstance(models, list) + # Should have the models from the mocked event + assert "gpt-3.5-turbo" in models + assert "gpt-4" in models @pytest.mark.integration @@ -239,22 +238,34 @@ async def test_providers_endpoint_offline_providers( mock_events: list[dict[str, Any]] = [ { "id": "event1", - "content": "Provider: http://healthy-provider.onion", + "pubkey": "healthy_provider_pubkey", "created_at": 1234567890, + "content": "Healthy provider announcement", + "tags": [ + ["d", "healthy-provider"], + ["endpoint", "http://healthy-provider.onion"], + ["name", "Healthy Provider"], + ] }, { "id": "event2", - "content": "Provider: http://offline-provider.onion", + "pubkey": "offline_provider_pubkey", "created_at": 1234567891, + "content": "Offline provider announcement", + "tags": [ + ["d", "offline-provider"], + ["endpoint", "http://offline-provider.onion"], + ["name", "Offline Provider"], + ] }, ] # Mock one healthy and one offline provider def mock_fetch_provider_health(url: str) -> dict[str, Any]: if "healthy" in url: - return {"status_code": 200, "json": {"status": "online"}} + return {"status_code": 200, "endpoint": "root", "json": {"status": "online"}} else: - return {"status_code": 500, "json": {"error": "Service unavailable"}} + return {"status_code": 500, "endpoint": "error", "json": {"error": "Service unavailable"}} with patch( "router.discovery.query_nostr_relay_for_providers", return_value=mock_events @@ -272,16 +283,21 @@ async def test_providers_endpoint_offline_providers( assert len(data["providers"]) == 2 # Verify that offline providers are still included but marked appropriately - for provider_dict in data["providers"]: - url = list(provider_dict.keys())[0] - provider_info = provider_dict[url] - - if "offline" in url: - # Offline provider should have error information - assert "error" in provider_info + for provider_data in data["providers"]: + assert "provider" in provider_data + assert "health" in provider_data + + provider_info = provider_data["provider"] + health_info = provider_data["health"] + + if "offline" in provider_info["endpoint_url"]: + # Offline provider should have error information in health + assert health_info["status_code"] == 500 + assert "error" in health_info["json"] else: - # Healthy provider should have status info - assert "status" in provider_info or "error" not in provider_info + # Healthy provider should have successful health check + assert health_info["status_code"] == 200 + assert "status" in health_info["json"] or "error" not in health_info["json"] @pytest.mark.integration @@ -291,22 +307,29 @@ async def test_providers_endpoint_duplicate_urls( ) -> None: """Test providers endpoint handles duplicate URLs correctly""" - # Mock events with duplicate provider URLs + # Mock events with duplicate provider events (same event ID) - should be deduplicated by relay query logic mock_events: list[dict[str, Any]] = [ { "id": "event1", - "content": "Check out http://provider.onion", + "pubkey": "provider_pubkey", "created_at": 1234567890, + "content": "Provider announcement", + "tags": [ + ["d", "provider-1"], + ["endpoint", "http://provider.onion"], + ["name", "Provider"], + ] }, { "id": "event2", - "content": "Also try http://provider.onion for good service", - "created_at": 1234567891, - }, - { - "id": "event3", - "content": "Different provider: http://other-provider.onion", + "pubkey": "other_provider_pubkey", "created_at": 1234567892, + "content": "Different provider announcement", + "tags": [ + ["d", "other-provider"], + ["endpoint", "http://other-provider.onion"], + ["name", "Other Provider"], + ] }, ] @@ -314,20 +337,24 @@ async def test_providers_endpoint_duplicate_urls( "router.discovery.query_nostr_relay_for_providers", return_value=mock_events ): with patch("router.discovery.fetch_provider_health") as mock_fetch: - mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} + mock_fetch.return_value = {"status_code": 200, "endpoint": "root", "json": {"status": "online"}} response = await integration_client.get("/v1/providers/") assert response.status_code == 200 data = response.json() - # Should deduplicate URLs + # Should return 2 unique providers based on events providers = data["providers"] - assert len(providers) == 2 # Only 2 unique URLs + assert len(providers) == 2 # 2 unique events - # Verify no duplicates - unique_providers = set(providers) - assert len(unique_providers) == len(providers) + # Verify all providers are unique by endpoint_url + endpoint_urls = [] + for provider_data in providers: + endpoint_urls.append(provider_data["endpoint_url"]) + + unique_endpoints = set(endpoint_urls) + assert len(unique_endpoints) == len(endpoint_urls) @pytest.mark.integration diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 8731d5ee..17b4f965 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -14,6 +14,7 @@ from httpx import AsyncClient from sqlmodel import select from router.core.db import ApiKey +from router.wallet import CurrencyUnit @pytest.mark.integration @@ -205,13 +206,14 @@ async def test_refund_with_lightning_address( # Capture state await db_snapshot.capture() - # Mock wallet.send_to_lnurl - with patch("router.wallet.send_token") as mock_wallet_func: - mock_wallet = AsyncMock() - mock_wallet.send_to_lnurl = AsyncMock( - return_value=500 - ) # Return amount sent # type: ignore[method-assign] - mock_wallet_func.return_value = mock_wallet + # Mock send_to_lnurl function directly + with patch("router.balance.send_to_lnurl") as mock_send_to_lnurl: + mock_send_to_lnurl.return_value = { + "amount_sent": balance, + "unit": "msat", + "lnurl": refund_address, + "status": "completed" + } # Request refund integration_client.headers["Authorization"] = f"Bearer {api_key}" @@ -225,10 +227,11 @@ async def test_refund_with_lightning_address( assert data["msats"] == balance assert "token" not in data - # Verify send_to_lnurl was called - mock_wallet.send_to_lnurl.assert_called_once_with( - refund_address, - amount=500, # 500 sats + # Verify send_to_lnurl was called with correct parameters + mock_send_to_lnurl.assert_called_once_with( + balance, # amount in msats + CurrencyUnit.msat, # unit + refund_address, # lnurl ) # Verify key was deleted by trying to use it @@ -408,7 +411,7 @@ async def test_mint_unavailability_handling( # Make the send_token method raise an exception with patch( - "router.wallet.send_token", + "router.balance.send_token", side_effect=Exception("Mint unavailable: Connection refused"), ): # The exception should propagate as a 503 error (Service Unavailable) diff --git a/tests/integration/test_wallet_topup.py b/tests/integration/test_wallet_topup.py index 3964d158..814adfd2 100644 --- a/tests/integration/test_wallet_topup.py +++ b/tests/integration/test_wallet_topup.py @@ -424,19 +424,18 @@ async def test_network_failure_during_token_verification( # type: ignore[no-unt # Generate a valid token token = await testmint_wallet.mint_tokens(300) - # Mock wallet.redeem to simulate network failure - with patch("router.wallet.send_token") as mock_wallet: - mock_wallet.return_value.redeem = AsyncMock( - side_effect=Exception("Network error: Connection timeout") - ) + # Mock credit_balance to simulate network failure during token verification + with patch("router.balance.credit_balance") as mock_credit_balance: + mock_credit_balance.side_effect = Exception("Network error: Connection timeout") response = await authenticated_client.post( "/v1/wallet/topup", params={"cashu_token": token} ) - # Should return 400 error - assert response.status_code == 400 + # Should return 500 error for network issues + assert response.status_code == 500 assert "detail" in response.json() + assert response.json()["detail"] == "Internal server error" @pytest.mark.integration @@ -471,7 +470,7 @@ async def test_topup_with_zero_amount_token( # type: ignore[no-untyped-def] # Create a token with 0 amount (edge case) # The testmint wallet should handle this - with patch.object(testmint_wallet, "redeem_token", return_value=0): + with patch.object(testmint_wallet, "redeem_token", return_value=(0, "sat", testmint_wallet.mint_url)): token = await testmint_wallet.mint_tokens(0) response = await authenticated_client.post( diff --git a/tests/integration/utils.py b/tests/integration/utils.py index dc23b353..6adcd8d3 100644 --- a/tests/integration/utils.py +++ b/tests/integration/utils.py @@ -68,10 +68,6 @@ class CashuTokenGenerator: # Invalid JSON structure lambda: "cashuA" + base64.urlsafe_b64encode(b'{"invalid": "structure"}').decode(), - # Empty proofs - lambda: CashuTokenGenerator._encode_token( - {"token": [{"mint": "https://test.com", "proofs": []}], "unit": "sat"} - ), # Invalid proof structure lambda: CashuTokenGenerator._encode_token( {