diff --git a/routstr/balance.py b/routstr/balance.py index 7128916d..37106c4f 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -211,8 +211,20 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: _refund_cache[key] = (expiry, value) +async def _lookup_key_no_create( + bearer_value: str, session: AsyncSession +) -> ApiKey | None: + """Look up an existing API key without creating one Used by the refund endpoint""" + if bearer_value.startswith("sk-"): + return await session.get(ApiKey, bearer_value[3:]) + if bearer_value.startswith("cashu"): + hashed = hashlib.sha256(bearer_value.encode()).hexdigest() + return await session.get(ApiKey, hashed) + return None + + async def _restore_balance( - session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int + session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -227,7 +239,7 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={"hashed_key": hashed_key, "restored_balance": balance}, + extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, ) @@ -282,7 +294,12 @@ async def refund_wallet_endpoint( ) bearer_value: str = authorization[7:] - key: ApiKey = await validate_bearer_key(bearer_value, session) + key: ApiKey | None = await _lookup_key_no_create(bearer_value, session) + if key is None: + raise HTTPException( + status_code=401, + detail="Key not found. Deposit first via /v1/wallet/create before requesting a refund.", + ) if key.total_balance <= 0: if cached := await _refund_cache_get(bearer_value): @@ -373,11 +390,11 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") error_msg = str(e) if ( "mint" in error_msg.lower() diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 4c79675c..1c62385d 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -372,17 +372,17 @@ async def test_refund_rejects_concurrent_topup_on_same_key( topup_amount_sat = 500 topup_token = await testmint_wallet.mint_tokens(topup_amount_sat) - validate_called = asyncio.Event() + key_looked_up = asyncio.Event() allow_refund_to_continue = asyncio.Event() - original_validate_bearer_key = balance_module.validate_bearer_key + original_lookup = balance_module._lookup_key_no_create delayed_once = False - async def delayed_validate_bearer_key(*args: Any, **kwargs: Any) -> ApiKey: + async def delayed_lookup_key_no_create(*args: Any, **kwargs: Any) -> ApiKey | None: nonlocal delayed_once - key = await original_validate_bearer_key(*args, **kwargs) - if not delayed_once: + key = await original_lookup(*args, **kwargs) + if not delayed_once and key is not None: delayed_once = True - validate_called.set() + key_looked_up.set() await allow_refund_to_continue.wait() return key @@ -390,7 +390,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key( return await authenticated_client.post("/v1/wallet/refund") async def issue_topup() -> Any: - await validate_called.wait() + await key_looked_up.wait() try: return await authenticated_client.post( "/v1/wallet/topup", params={"cashu_token": topup_token} @@ -399,7 +399,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key( allow_refund_to_continue.set() with patch( - "routstr.balance.validate_bearer_key", new=delayed_validate_bearer_key + "routstr.balance._lookup_key_no_create", new=delayed_lookup_key_no_create ): refund_response, topup_response = await asyncio.gather( issue_refund(), issue_topup() diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 11456367..ba9d54da 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -168,12 +168,12 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No refund_token = "cashuArefund_apikey_token" session = MagicMock() + session.get = AsyncMock(return_value=key) session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() with ( - patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store, @@ -203,12 +203,12 @@ async def test_apikey_refund_logs_token() -> None: refund_token = "cashuAlogged_token" session = MagicMock() + session.get = AsyncMock(return_value=key) session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() with ( - patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()), @@ -232,12 +232,12 @@ async def test_apikey_refund_log_includes_path() -> None: refund_token = "cashuApath_token" session = MagicMock() + session.get = AsyncMock(return_value=key) session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() with ( - patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.balance.store_cashu_transaction", AsyncMock()), @@ -269,6 +269,7 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None: key = _make_api_key(balance=5000, refund_currency="sat") session = MagicMock() + session.get = AsyncMock(return_value=key) # Debit returns rowcount=0 → balance changed concurrently session.exec = AsyncMock(return_value=_update_result(0)) session.commit = AsyncMock() @@ -276,7 +277,6 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None: mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted") with ( - patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", mock_send_token), patch("routstr.balance.store_cashu_transaction", AsyncMock()), @@ -333,11 +333,11 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: # First exec call = debit (succeeds), second = restore session = MagicMock() + session.get = AsyncMock(return_value=key) session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)]) session.commit = AsyncMock() with ( - patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)), patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))), patch("routstr.balance.store_cashu_transaction", AsyncMock()), @@ -355,3 +355,49 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: assert exc_info.value.status_code == 503 # Verify two exec calls: debit + restore assert session.exec.await_count == 2 + + +# --------------------------------------------------------------------------- +# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_refund_fresh_cashu_bearer_returns_401() -> None: + """Fresh Cashu token not in DB must get 401, never create a new ApiKey.""" + from fastapi import HTTPException + + session = MagicMock() + session.get = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer cashuAfresh_never_deposited_token", + x_cashu=None, + session=session, + ) + + assert exc_info.value.status_code == 401 + session.get.assert_awaited_once() + # No add/commit → no key was persisted + session.add.assert_not_called() + session.commit.assert_not_called() + + +@pytest.mark.asyncio +async def test_refund_unknown_sk_bearer_returns_401() -> None: + """Unknown sk- key not in DB must get 401.""" + from fastapi import HTTPException + + session = MagicMock() + session.get = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-unknownhash", + x_cashu=None, + session=session, + ) + + assert exc_info.value.status_code == 401 + session.get.assert_awaited_once()