fix: harden LNURL refund destination checks and surface in-progress refunds

This commit is contained in:
9qeklajc
2026-09-16 12:08:03 +02:00
parent eb9587b290
commit c3970de510
5 changed files with 173 additions and 25 deletions
+4
View File
@@ -367,6 +367,10 @@ async def refund_wallet_endpoint(
return persisted return persisted
if paid: if paid:
return refund.describe(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: if key.reserved_balance > 0:
# Release only durable reservations old enough to be stale. A newer # Release only durable reservations old enough to be stale. A newer
+57 -15
View File
@@ -1,6 +1,9 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import ipaddress import ipaddress
import json
import socket
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Any, TypedDict from typing import Any, TypedDict
@@ -48,15 +51,18 @@ class MeltOutcomeAmbiguousError(LNURLError):
_MAX_LNURL_REDIRECTS = 3 _MAX_LNURL_REDIRECTS = 3
_MAX_LNURL_RESPONSE_BYTES = 64 * 1024
_NON_PUBLIC_HOST_SUFFIXES = (".localhost", ".local", ".internal") _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. """Reject anything that is not a public HTTPS endpoint.
LNURL destinations and their redirect targets are attacker-influenced, so LNURL destinations and their redirect targets are attacker-influenced, so
every hop has to be re-checked: a single ``https://`` origin says nothing 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": if url.scheme != "https":
raise LNURLError("LNURL destination must be an HTTPS URL") 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") raise LNURLError("LNURL destination has no host")
try: try:
address = ipaddress.ip_address(host) literal = ipaddress.ip_address(host)
except ValueError: except ValueError:
if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES): literal = None
raise LNURLError("LNURL destination is not a public host") from None
if literal is not None:
if not literal.is_global:
raise LNURLError("LNURL destination is not a public host")
return 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") 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( async def _fetch_lnurl_json(
url: str, params: dict[str, int] | None = None 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) target = httpx.URL(url, params=params) if params else httpx.URL(url)
except httpx.InvalidURL as e: except httpx.InvalidURL as e:
raise LNURLError("LNURL destination is not a usable URL") from 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: async with httpx.AsyncClient() as client:
for _ in range(_MAX_LNURL_REDIRECTS + 1): for _ in range(_MAX_LNURL_REDIRECTS + 1):
response = await client.get(target, follow_redirects=False, timeout=10) async with client.stream(
if not response.is_redirect: "GET", target, follow_redirects=False, timeout=10
break ) as response:
target = target.join(response.headers.get("location", "")) if response.is_redirect:
_require_public_https_destination(target) 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: else:
raise LNURLError("LNURL destination exceeded the redirect limit") 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: try:
data = response.json() data = json.loads(raw)
except ValueError as e: except ValueError as e:
raise LNURLError("LNURL response was not valid JSON") from 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): if not isinstance(callback_url, str):
raise LNURLError("Invalid LNURL payRequest: missing callback URL") raise LNURLError("Invalid LNURL payRequest: missing callback URL")
try: try:
_require_public_https_destination(httpx.URL(callback_url)) callback_target = httpx.URL(callback_url)
except httpx.InvalidURL as e: except httpx.InvalidURL as e:
raise LNURLError("Invalid LNURL callback URL") from 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 min_sendable = lnurl_data.get("minSendable", 1000) # Default 1 sat
max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC
+30 -10
View File
@@ -94,16 +94,7 @@ async def open_claim(
await session.commit() await session.commit()
except IntegrityError: except IntegrityError:
await session.rollback() await session.rollback()
raise HTTPException( raise refund_in_progress_error()
status_code=409,
detail={
"error": {
"message": "A refund for this key is already in progress.",
"type": "invalid_request_error",
"code": "refund_in_progress",
}
},
)
return refund return refund
@@ -205,6 +196,35 @@ async def hold(session: AsyncSession, refund: Refund, quote_id: str | None) -> N
await session.commit() 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: async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None:
"""Latest paid refund of either method. """Latest paid refund of either method.
+26
View File
@@ -583,6 +583,32 @@ async def test_endpoint_refund_while_ambiguous_returns_409(
assert (await _load_key(integration_session)).balance == 2_000_000 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 @pytest.mark.asyncio
async def test_cashu_token_survives_failed_ledger_write( async def test_cashu_token_survives_failed_ledger_write(
integration_session: AsyncSession, patched_db_engine: None integration_session: AsyncSession, patched_db_engine: None
@@ -6,6 +6,7 @@ is what stops a pre-dispatch failure from stranding proofs.
""" """
import math import math
import socket
from collections.abc import Callable from collections.abc import Callable
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch 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 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: def test_select_melt_proofs_ignores_fees_for_unneeded_wallet_proofs() -> None:
from routstr.payment.lnurl import _select_melt_proofs from routstr.payment.lnurl import _select_melt_proofs