diff --git a/routstr/balance.py b/routstr/balance.py index 2f4140cf..4e163a53 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -367,6 +367,10 @@ async def refund_wallet_endpoint( return persisted if paid: return refund.describe(paid) + # Balance reads zero because a prior refund already debited it and is + # still settling; surface that as 409 rather than "no balance". + if await refund.latest_open(session, key): + raise refund.refund_in_progress_error() if key.reserved_balance > 0: # Release only durable reservations old enough to be stale. A newer diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 2f5c59a6..9201bca7 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,6 +1,9 @@ from __future__ import annotations +import asyncio import ipaddress +import json +import socket from collections.abc import Awaitable, Callable from typing import Any, TypedDict @@ -48,15 +51,18 @@ class MeltOutcomeAmbiguousError(LNURLError): _MAX_LNURL_REDIRECTS = 3 +_MAX_LNURL_RESPONSE_BYTES = 64 * 1024 _NON_PUBLIC_HOST_SUFFIXES = (".localhost", ".local", ".internal") -def _require_public_https_destination(url: httpx.URL) -> None: +async def _require_public_https_destination(url: httpx.URL) -> None: """Reject anything that is not a public HTTPS endpoint. LNURL destinations and their redirect targets are attacker-influenced, so every hop has to be re-checked: a single ``https://`` origin says nothing - about where a 302 points. + about where a 302 points. A bare hostname check is not enough either: a + public-looking name can resolve to a loopback/link-local/private address + (SSRF), so DNS is resolved here and every resulting address must be global. """ if url.scheme != "https": raise LNURLError("LNURL destination must be an HTTPS URL") @@ -66,15 +72,37 @@ def _require_public_https_destination(url: httpx.URL) -> None: raise LNURLError("LNURL destination has no host") try: - address = ipaddress.ip_address(host) + literal = ipaddress.ip_address(host) except ValueError: - if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES): - raise LNURLError("LNURL destination is not a public host") from None + literal = None + + if literal is not None: + if not literal.is_global: + raise LNURLError("LNURL destination is not a public host") return - if not address.is_global: + if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES): raise LNURLError("LNURL destination is not a public host") + port = url.port or 443 + try: + infos = await asyncio.get_running_loop().getaddrinfo( + host, port, proto=socket.IPPROTO_TCP + ) + except socket.gaierror as e: + raise LNURLError("LNURL destination could not be resolved") from e + if not infos: + raise LNURLError("LNURL destination could not be resolved") + for info in infos: + try: + resolved = ipaddress.ip_address(info[4][0]) + except ValueError as e: + raise LNURLError( + "LNURL destination resolved to an invalid address" + ) from e + if not resolved.is_global: + raise LNURLError("LNURL destination is not a public host") + async def _fetch_lnurl_json( url: str, params: dict[str, int] | None = None @@ -88,21 +116,34 @@ async def _fetch_lnurl_json( target = httpx.URL(url, params=params) if params else httpx.URL(url) except httpx.InvalidURL as e: raise LNURLError("LNURL destination is not a usable URL") from e - _require_public_https_destination(target) + await _require_public_https_destination(target) + raw: bytes | None = None async with httpx.AsyncClient() as client: for _ in range(_MAX_LNURL_REDIRECTS + 1): - response = await client.get(target, follow_redirects=False, timeout=10) - if not response.is_redirect: - break - target = target.join(response.headers.get("location", "")) - _require_public_https_destination(target) + async with client.stream( + "GET", target, follow_redirects=False, timeout=10 + ) as response: + if response.is_redirect: + target = target.join(response.headers.get("location", "")) + await _require_public_https_destination(target) + continue + response.raise_for_status() + chunks = bytearray() + async for chunk in response.aiter_bytes(): + chunks.extend(chunk) + if len(chunks) > _MAX_LNURL_RESPONSE_BYTES: + raise LNURLError("LNURL response exceeded the size limit") + raw = bytes(chunks) + break else: raise LNURLError("LNURL destination exceeded the redirect limit") - response.raise_for_status() + + if raw is None: + raise LNURLError("LNURL destination exceeded the redirect limit") try: - data = response.json() + data = json.loads(raw) except ValueError as e: raise LNURLError("LNURL response was not valid JSON") from e @@ -190,9 +231,10 @@ async def get_lnurl_data(lnurl: str) -> LNURLData: if not isinstance(callback_url, str): raise LNURLError("Invalid LNURL payRequest: missing callback URL") try: - _require_public_https_destination(httpx.URL(callback_url)) + callback_target = httpx.URL(callback_url) except httpx.InvalidURL as e: raise LNURLError("Invalid LNURL callback URL") from e + await _require_public_https_destination(callback_target) min_sendable = lnurl_data.get("minSendable", 1000) # Default 1 sat max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC diff --git a/routstr/refund.py b/routstr/refund.py index 5c328a57..028b3323 100644 --- a/routstr/refund.py +++ b/routstr/refund.py @@ -94,16 +94,7 @@ async def open_claim( await session.commit() except IntegrityError: await session.rollback() - raise HTTPException( - status_code=409, - detail={ - "error": { - "message": "A refund for this key is already in progress.", - "type": "invalid_request_error", - "code": "refund_in_progress", - } - }, - ) + raise refund_in_progress_error() return refund @@ -205,6 +196,35 @@ async def hold(session: AsyncSession, refund: Refund, quote_id: str | None) -> N await session.commit() +def refund_in_progress_error() -> HTTPException: + """The 409 raised when a key already has an in-flight refund claim.""" + return HTTPException( + status_code=409, + detail={ + "error": { + "message": "A refund for this key is already in progress.", + "type": "invalid_request_error", + "code": "refund_in_progress", + } + }, + ) + + +async def latest_open(session: AsyncSession, key: ApiKey) -> Refund | None: + """Latest non-terminal (in-flight) claim for the key, if any. + + An open claim means a prior refund already debited the balance and is still + settling, so the balance reads as zero even though a refund is under way. + """ + result = await session.exec( + select(Refund) + .where(Refund.api_key_hashed_key == key.hashed_key) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .order_by(col(Refund.created_at).desc(), col(Refund.updated_at).desc()) + ) + return result.first() + + async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None: """Latest paid refund of either method. diff --git a/tests/integration/test_refund_claims.py b/tests/integration/test_refund_claims.py index 57e66795..aa8a4a32 100644 --- a/tests/integration/test_refund_claims.py +++ b/tests/integration/test_refund_claims.py @@ -583,6 +583,32 @@ async def test_endpoint_refund_while_ambiguous_returns_409( assert (await _load_key(integration_session)).balance == 2_000_000 +@pytest.mark.asyncio +async def test_endpoint_zero_balance_with_open_claim_returns_409( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A prior refund debited the balance and is still settling: the retry must + report refund_in_progress (409), not "no balance to refund" (400).""" + await _open_ambiguous(integration_session, "quote-123") + assert (await _load_key(integration_session)).balance == 0 + + send = AsyncMock() + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", send), + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 409 + assert isinstance(exc_info.value.detail, dict) + assert exc_info.value.detail["error"]["code"] == "refund_in_progress" + send.assert_not_awaited() + + @pytest.mark.asyncio async def test_cashu_token_survives_failed_ledger_write( integration_session: AsyncSession, patched_db_engine: None diff --git a/tests/unit/test_lnurl_amount_and_destination.py b/tests/unit/test_lnurl_amount_and_destination.py index ab6a996e..a8e08dfb 100644 --- a/tests/unit/test_lnurl_amount_and_destination.py +++ b/tests/unit/test_lnurl_amount_and_destination.py @@ -6,6 +6,7 @@ is what stops a pre-dispatch failure from stranding proofs. """ import math +import socket from collections.abc import Callable from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -347,6 +348,61 @@ async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None: assert raw_send.await_args.kwargs["amount"] == 1000 +def _patch_getaddrinfo(ip: str) -> Any: + """Force DNS resolution of any hostname to a single fixed IP.""" + + async def fake_getaddrinfo(host: str, port: int, **_kw: object) -> list[Any]: + return [ + (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", (ip, port)) + ] + + loop = MagicMock() + loop.getaddrinfo = fake_getaddrinfo + return patch.object( + lnurl_module.asyncio, "get_running_loop", return_value=loop + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "private_ip", ["127.0.0.1", "10.0.0.5", "169.254.169.254", "192.168.1.1"] +) +async def test_guard_rejects_public_hostname_resolving_to_private( + private_ip: str, +) -> None: + """A public-looking name must be rejected when DNS points it inward (SSRF).""" + with ( + _patch_getaddrinfo(private_ip), + pytest.raises(LNURLError, match="public host"), + ): + await lnurl_module._require_public_https_destination( + httpx.URL("https://totally-public.example.com/cb") + ) + + +@pytest.mark.asyncio +async def test_guard_allows_public_hostname_resolving_to_public() -> None: + with _patch_getaddrinfo("93.184.216.34"): + await lnurl_module._require_public_https_destination( + httpx.URL("https://totally-public.example.com/cb") + ) + + +@pytest.mark.asyncio +async def test_get_lnurl_data_rejects_oversized_response() -> None: + big = b"x" * (lnurl_module._MAX_LNURL_RESPONSE_BYTES + 1) + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, content=big) + + with ( + _patch_getaddrinfo("93.184.216.34"), + _mock_client(handler), + pytest.raises(LNURLError, match="size limit"), + ): + await get_lnurl_data("owner@ln.tld") + + def test_select_melt_proofs_ignores_fees_for_unneeded_wallet_proofs() -> None: from routstr.payment.lnurl import _select_melt_proofs