mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: harden LNURL refund destination checks and surface in-progress refunds
This commit is contained in:
@@ -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
|
||||
|
||||
+57
-15
@@ -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
|
||||
|
||||
+30
-10
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user