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
|
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
@@ -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
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user