From a833cf429e042efe09cf0991a5b8b1daca24d4df Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 29 Mar 2026 21:22:09 +0200 Subject: [PATCH 1/3] add refund x-cashu to default refund endpoint --- routstr/balance.py | 19 +++++++ tests/unit/test_balance.py | 100 +++++++++++++++++++++++++++++++++++++ 2 files changed, 119 insertions(+) create mode 100644 tests/unit/test_balance.py diff --git a/routstr/balance.py b/routstr/balance.py index 9adbd9b7..6eed34aa 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -5,6 +5,7 @@ from time import monotonic from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException +from fastapi.responses import JSONResponse from pydantic import BaseModel from sqlmodel import select @@ -207,6 +208,7 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: @router.post("/refund") async def refund_wallet_endpoint( authorization: Annotated[str, Header(...)], + x_cashu: Annotated[str | None, Header()] = None, session: AsyncSession = Depends(get_session), ) -> dict[str, str]: if not authorization.startswith("Bearer "): @@ -217,6 +219,23 @@ async def refund_wallet_endpoint( bearer_value: str = authorization[7:] + if x_cashu: + payment_token_hash = hashlib.sha256(x_cashu.strip().encode()).hexdigest() + result = await session.get(CashuTransaction, payment_token_hash) + if result is None: + raise HTTPException(status_code=404, detail="Refund not found") + if result.swept: + raise HTTPException(status_code=410, detail="Refund has been swept") + result.collected = True + session.add(result) + await session.commit() + body: dict[str, str] = {"token": result.token} + if result.unit == "sat": + body["sats"] = str(result.amount) + else: + body["msats"] = str(result.amount) + return JSONResponse(content=body, headers={"X-Cashu": result.token}) + key: ApiKey = await validate_bearer_key(bearer_value, session) if key.total_balance <= 0: diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py new file mode 100644 index 00000000..fdb4d0fc --- /dev/null +++ b/tests/unit/test_balance.py @@ -0,0 +1,100 @@ +import hashlib +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from routstr.balance import refund_wallet_endpoint +from routstr.core.db import CashuTransaction + + +def _make_cashu_tx(token: str, amount: int, unit: str, swept: bool = False) -> CashuTransaction: + tx = CashuTransaction(token=token, amount=amount, unit=unit) + tx.swept = swept + tx.collected = False + return tx + + +@pytest.mark.asyncio +async def test_refund_x_cashu_returns_token() -> None: + x_cashu_token = "cashuAtest_token_value" + expected_hash = hashlib.sha256(x_cashu_token.strip().encode()).hexdigest() + tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat") + + session = MagicMock() + session.get = AsyncMock(return_value=tx) + session.add = MagicMock() + session.commit = AsyncMock() + + result = await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu=x_cashu_token, + session=session, + ) + + session.get.assert_awaited_once_with(CashuTransaction, expected_hash) + import json + body = json.loads(result.body) + assert body["token"] == "cashuArefund_token" + assert body["msats"] == "1000" + assert result.headers["X-Cashu"] == "cashuArefund_token" + assert tx.collected is True + + +@pytest.mark.asyncio +async def test_refund_x_cashu_sat_unit() -> None: + x_cashu_token = "cashuAsat_token" + tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat") + + session = MagicMock() + session.get = AsyncMock(return_value=tx) + session.add = MagicMock() + session.commit = AsyncMock() + + result = await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu=x_cashu_token, + session=session, + ) + + import json + body = json.loads(result.body) + assert body["token"] == "cashuArefund_sat" + assert body["sats"] == "500" + assert "msats" not in body + assert result.headers["X-Cashu"] == "cashuArefund_sat" + + +@pytest.mark.asyncio +async def test_refund_x_cashu_not_found_raises_404() -> None: + 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-somekey", + x_cashu="cashuAmissing_token", + session=session, + ) + + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_refund_x_cashu_swept_raises_410() -> None: + from fastapi import HTTPException + + tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", swept=True) + + session = MagicMock() + session.get = AsyncMock(return_value=tx) + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-somekey", + x_cashu="cashuAswept_token", + session=session, + ) + + assert exc_info.value.status_code == 410 From 6fde846be277ba754e74e4834ba68cb618bc07fd Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Apr 2026 15:04:38 +0200 Subject: [PATCH 2/3] refund x cashu --- routstr/balance.py | 62 +++++++++++++++++++++++++------------- tests/unit/test_balance.py | 6 ++++ 2 files changed, 47 insertions(+), 21 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index 6eed34aa..d8ec98f5 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -207,35 +207,55 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: @router.post("/refund") async def refund_wallet_endpoint( - authorization: Annotated[str, Header(...)], + authorization: Annotated[str | None, Header()] = None, x_cashu: Annotated[str | None, Header()] = None, session: AsyncSession = Depends(get_session), -) -> dict[str, str]: - if not authorization.startswith("Bearer "): +) -> JSONResponse | dict[str, str]: + if x_cashu: + # Find the "in" transaction by the original payment token + in_tx_result = await session.exec( + select(CashuTransaction).where( + CashuTransaction.token == x_cashu, + CashuTransaction.type == "in", + ) + ) + in_tx = in_tx_result.first() + if in_tx is None: + raise HTTPException(status_code=404, detail="Refund not found") + + # Use the request_id to find the associated "out" (refund) transaction + if in_tx.request_id is None: + raise HTTPException(status_code=404, detail="Refund not found") + + out_tx_result = await session.exec( + select(CashuTransaction).where( + CashuTransaction.request_id == in_tx.request_id, + CashuTransaction.type == "out", + ) + ) + out_tx = out_tx_result.first() + if out_tx is None: + raise HTTPException(status_code=404, detail="Refund not found") + if out_tx.swept: + raise HTTPException(status_code=410, detail="Refund has been swept") + + out_tx.collected = True + session.add(out_tx) + await session.commit() + body: dict[str, str] = {"token": out_tx.token} + if out_tx.unit == "sat": + body["sats"] = str(out_tx.amount) + else: + body["msats"] = str(out_tx.amount) + return JSONResponse(content=body, headers={"X-Cashu": out_tx.token}) + + if authorization is None or not authorization.startswith("Bearer "): raise HTTPException( status_code=401, detail="Invalid authorization. Use 'Bearer ' or 'Bearer '", ) bearer_value: str = authorization[7:] - - if x_cashu: - payment_token_hash = hashlib.sha256(x_cashu.strip().encode()).hexdigest() - result = await session.get(CashuTransaction, payment_token_hash) - if result is None: - raise HTTPException(status_code=404, detail="Refund not found") - if result.swept: - raise HTTPException(status_code=410, detail="Refund has been swept") - result.collected = True - session.add(result) - await session.commit() - body: dict[str, str] = {"token": result.token} - if result.unit == "sat": - body["sats"] = str(result.amount) - else: - body["msats"] = str(result.amount) - return JSONResponse(content=body, headers={"X-Cashu": result.token}) - key: ApiKey = await validate_bearer_key(bearer_value, session) if key.total_balance <= 0: diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index fdb4d0fc..f36ef158 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -33,6 +33,9 @@ async def test_refund_x_cashu_returns_token() -> None: session.get.assert_awaited_once_with(CashuTransaction, expected_hash) import json + + from fastapi.responses import JSONResponse + assert isinstance(result, JSONResponse) body = json.loads(result.body) assert body["token"] == "cashuArefund_token" assert body["msats"] == "1000" @@ -57,6 +60,9 @@ async def test_refund_x_cashu_sat_unit() -> None: ) import json + + from fastapi.responses import JSONResponse + assert isinstance(result, JSONResponse) body = json.loads(result.body) assert body["token"] == "cashuArefund_sat" assert body["sats"] == "500" From 905ae670151c6f3fb8affbda4212e0cd65e2e46c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Apr 2026 15:11:43 +0200 Subject: [PATCH 3/3] update tests --- routstr/balance.py | 2 +- tests/unit/test_balance.py | 50 +++++++++++++++++++++++--------------- 2 files changed, 31 insertions(+), 21 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index d8ec98f5..e63f2cc3 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -205,7 +205,7 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: _refund_cache[key] = (expiry, value) -@router.post("/refund") +@router.post("/refund", response_model=None) async def refund_wallet_endpoint( authorization: Annotated[str | None, Header()] = None, x_cashu: Annotated[str | None, Header()] = None, diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index f36ef158..b9ff980d 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -1,27 +1,42 @@ -import hashlib +import json from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi.responses import JSONResponse from routstr.balance import refund_wallet_endpoint from routstr.core.db import CashuTransaction -def _make_cashu_tx(token: str, amount: int, unit: str, swept: bool = False) -> CashuTransaction: - tx = CashuTransaction(token=token, amount=amount, unit=unit) +def _make_cashu_tx( + token: str, + amount: int, + unit: str, + type: str = "out", + request_id: str | None = "req-abc", + swept: bool = False, + collected: bool = False, +) -> CashuTransaction: + tx = CashuTransaction(token=token, amount=amount, unit=unit, type=type, request_id=request_id) tx.swept = swept - tx.collected = False + tx.collected = collected return tx +def _exec_result(tx: CashuTransaction | None) -> MagicMock: + result = MagicMock() + result.first.return_value = tx + return result + + @pytest.mark.asyncio async def test_refund_x_cashu_returns_token() -> None: x_cashu_token = "cashuAtest_token_value" - expected_hash = hashlib.sha256(x_cashu_token.strip().encode()).hexdigest() - tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat") + in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc") + out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc") session = MagicMock() - session.get = AsyncMock(return_value=tx) + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() @@ -31,25 +46,22 @@ async def test_refund_x_cashu_returns_token() -> None: session=session, ) - session.get.assert_awaited_once_with(CashuTransaction, expected_hash) - import json - - from fastapi.responses import JSONResponse assert isinstance(result, JSONResponse) body = json.loads(result.body) assert body["token"] == "cashuArefund_token" assert body["msats"] == "1000" assert result.headers["X-Cashu"] == "cashuArefund_token" - assert tx.collected is True + assert out_tx.collected is True @pytest.mark.asyncio async def test_refund_x_cashu_sat_unit() -> None: x_cashu_token = "cashuAsat_token" - tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat") + in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat") + out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat") session = MagicMock() - session.get = AsyncMock(return_value=tx) + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() @@ -59,9 +71,6 @@ async def test_refund_x_cashu_sat_unit() -> None: session=session, ) - import json - - from fastapi.responses import JSONResponse assert isinstance(result, JSONResponse) body = json.loads(result.body) assert body["token"] == "cashuArefund_sat" @@ -75,7 +84,7 @@ async def test_refund_x_cashu_not_found_raises_404() -> None: from fastapi import HTTPException session = MagicMock() - session.get = AsyncMock(return_value=None) + session.exec = AsyncMock(return_value=_exec_result(None)) with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -91,10 +100,11 @@ async def test_refund_x_cashu_not_found_raises_404() -> None: async def test_refund_x_cashu_swept_raises_410() -> None: from fastapi import HTTPException - tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", swept=True) + in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept") + out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True) session = MagicMock() - session.get = AsyncMock(return_value=tx) + session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint(