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
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
+56 -14
View File
@@ -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,13 +72,35 @@ 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")
@@ -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
async with client.stream(
"GET", target, follow_redirects=False, timeout=10
) as response:
if response.is_redirect:
target = target.join(response.headers.get("location", ""))
_require_public_https_destination(target)
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
+30 -10
View File
@@ -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.
+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
@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
@@ -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