fix: address review findings on cooldowns, melt sizing, and stream finalization

- open transport cooldown only after timeout retries are exhausted so
  bounded retries keep their short backoff instead of waiting out a
  full 30s cooldown in the rate guard
- shield streaming billing finalization and upstream connection close
  from client-disconnect cancellation
- stop proof selection at the minimal covering set when over budget so
  the re-quote shortfall is not inflated by extra input fees
- reuse is_mint_transport_error instead of a duplicate helper in lnurl
- ppqai: normalize circuit origin default ports, drop dead try/except
  in fetch_models
This commit is contained in:
9qeklajc
2026-08-26 23:15:40 +02:00
parent 6086fa92b0
commit f4e02d4ee2
6 changed files with 212 additions and 150 deletions
+5 -4
View File
@@ -306,11 +306,8 @@ async def run_mint_operation(
except MintCooldownError: except MintCooldownError:
raise raise
except (asyncio.TimeoutError, httpx.TimeoutException) as exc: except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
if guard is not None:
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
if retry_timeouts and attempt < max_attempts - 1: if retry_timeouts and attempt < max_attempts - 1:
# Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait.
backoff = (2**attempt) + (time.monotonic() % 1.0) backoff = (2**attempt) + (time.monotonic() % 1.0)
logger.warning( logger.warning(
"Mint operation timed out, retrying", "Mint operation timed out, retrying",
@@ -323,6 +320,10 @@ async def run_mint_operation(
) )
await asyncio.sleep(backoff) await asyncio.sleep(backoff)
continue continue
if guard is not None:
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise httpx.TimeoutException( raise httpx.TimeoutException(
f"{op_name} timed out (attempts: {attempt + 1})" f"{op_name} timed out (attempts: {attempt + 1})"
) from exc ) from exc
+7 -14
View File
@@ -9,8 +9,8 @@ from cashu.core.base import MeltQuoteState
from cashu.wallet.wallet import Proof, Wallet from cashu.wallet.wallet import Proof, Wallet
from ..mint import ( from ..mint import (
MINT_TRANSPORT_EXCEPTIONS,
is_mint_rate_limited, is_mint_rate_limited,
is_mint_transport_error,
run_mint_operation, run_mint_operation,
) )
@@ -106,17 +106,6 @@ async def _fetch_lnurl_json(
return data return data
def _contains_mint_transport_error(error: BaseException) -> bool:
seen: set[int] = set()
current: BaseException | None = error
while current is not None and id(current) not in seen:
seen.add(id(current))
if isinstance(current, MINT_TRANSPORT_EXCEPTIONS):
return True
current = current.__cause__ or current.__context__
return False
async def decode_lnurl(lnurl: str) -> str: async def decode_lnurl(lnurl: str) -> str:
"""Decode LNURL to get the actual URL. """Decode LNURL to get the actual URL.
@@ -260,8 +249,11 @@ def _select_melt_proofs(
selected_amount += proof.amount selected_amount += proof.amount
input_fees = int(wallet.get_fees_for_proofs(selected)) input_fees = int(wallet.get_fees_for_proofs(selected))
required = quote_amount + fee_reserve + input_fees required = quote_amount + fee_reserve + input_fees
if required <= gross_budget and selected_amount >= required: if selected_amount >= required:
if required <= gross_budget:
return selected, 0 return selected, 0
# Covered but over budget; more proofs only raise input fees.
break
return None, max(1, required - min(selected_amount, gross_budget)) return None, max(1, required - min(selected_amount, gross_budget))
@@ -362,6 +354,7 @@ async def raw_send_to_lnurl(
if on_melt_quote is not None: if on_melt_quote is not None:
await on_melt_quote(melt_quote_resp.quote) await on_melt_quote(melt_quote_resp.quote)
assert selected_proofs is not None
proofs = selected_proofs proofs = selected_proofs
await wallet.set_reserved_for_send(proofs, reserved=True) await wallet.set_reserved_for_send(proofs, reserved=True)
@@ -384,7 +377,7 @@ async def raw_send_to_lnurl(
# reserved as though a Lightning payment could still settle. # reserved as though a Lightning payment could still settle.
await wallet.set_reserved_for_send(proofs, reserved=False) await wallet.set_reserved_for_send(proofs, reserved=False)
raise raise
if not _contains_mint_transport_error(error): if not is_mint_transport_error(error):
raise raise
# Cashu clears reservations on transport errors despite an unknown outcome. # Cashu clears reservations on transport errors despite an unknown outcome.
try: try:
+36 -23
View File
@@ -7,7 +7,7 @@ import math
import traceback import traceback
import typing import typing
import uuid import uuid
from collections.abc import AsyncGenerator, AsyncIterator, Iterator from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator
from typing import Any, Mapping, Self, cast from typing import Any, Mapping, Self, cast
import httpx import httpx
@@ -82,6 +82,21 @@ async def _aclose_if_needed(resource: object | None) -> None:
await result await result
async def _finalize_and_close_stream(
finalize: Callable[[], Awaitable[None]] | None,
response: object | None,
client: httpx.AsyncClient | None,
) -> None:
try:
if finalize is not None:
await finalize()
finally:
try:
await _aclose_if_needed(response)
finally:
await _aclose_if_needed(client)
CostMetadata = CostData | MaxCostData | dict[str, Any] CostMetadata = CostData | MaxCostData | dict[str, Any]
@@ -1049,9 +1064,7 @@ class BaseUpstreamProvider:
return return
try: try:
async with create_session() as new_session: async with create_session() as new_session:
fresh_key = await new_session.get( fresh_key = await new_session.get(key.__class__, key.hashed_key)
key.__class__, key.hashed_key
)
if not fresh_key: if not fresh_key:
return return
try: try:
@@ -1318,14 +1331,15 @@ class BaseUpstreamProvider:
) )
raise raise
finally: finally:
try: # Shielded so a client disconnect cannot cancel billing
if not usage_finalized: # finalization or leak the upstream connection.
await finalize_db_only() await asyncio.shield(
finally: _finalize_and_close_stream(
try: None if usage_finalized else finalize_db_only,
await _aclose_if_needed(response) response,
finally: client,
await _aclose_if_needed(client) )
)
# Remove inaccurate encoding headers from upstream response # Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers) response_headers = dict(response.headers)
@@ -1521,9 +1535,7 @@ class BaseUpstreamProvider:
return return
try: try:
async with create_session() as new_session: async with create_session() as new_session:
fresh_key = await new_session.get( fresh_key = await new_session.get(key.__class__, key.hashed_key)
key.__class__, key.hashed_key
)
if not fresh_key: if not fresh_key:
return return
try: try:
@@ -1750,14 +1762,15 @@ class BaseUpstreamProvider:
) )
raise raise
finally: finally:
try: # Shielded so a client disconnect cannot cancel billing
if not usage_finalized: # finalization or leak the upstream connection.
await finalize_db_only() await asyncio.shield(
finally: _finalize_and_close_stream(
try: None if usage_finalized else finalize_db_only,
await _aclose_if_needed(response) response,
finally: client,
await _aclose_if_needed(client) )
)
# Remove inaccurate encoding headers from upstream response # Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers) response_headers = dict(response.headers)
+7 -20
View File
@@ -40,7 +40,8 @@ _ppq_circuits: dict[str, _PPQCircuitState] = {}
def _ppq_origin(url: str) -> str: def _ppq_origin(url: str) -> str:
parsed = httpx.URL(url) parsed = httpx.URL(url)
return f"{parsed.scheme}://{parsed.host}:{parsed.port}" port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0)
return f"{parsed.scheme}://{parsed.host}:{port}"
async def _safe_read_request( async def _safe_read_request(
@@ -66,9 +67,7 @@ async def _safe_read_request(
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1): for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
try: try:
response = await client.request( response = await client.request(method, url, headers=headers, json=json)
method, url, headers=headers, json=json
)
response.raise_for_status() response.raise_for_status()
state.consecutive_failures = 0 state.consecutive_failures = 0
state.cooldown_until = 0.0 state.cooldown_until = 0.0
@@ -209,11 +208,8 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
url = f"{self.base_url}/models" url = f"{self.base_url}/models"
headers = {"Authorization": f"Bearer {self.api_key}"} headers = {"Authorization": f"Bearer {self.api_key}"}
try:
async with httpx.AsyncClient(timeout=30.0) as client: async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request( response = await _safe_read_request(client, "GET", url, headers=headers)
client, "GET", url, headers=headers
)
data = response.json() data = response.json()
models_data = data.get("data", []) models_data = data.get("data", [])
@@ -244,9 +240,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
if or_model: if or_model:
input_price = None input_price = None
if ppqai_model.pricing.api: if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get( input_price = ppqai_model.pricing.api.get("input_per_1M")
"input_per_1M"
)
elif ppqai_model.pricing.input_per_1M_tokens: elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens input_price = ppqai_model.pricing.input_per_1M_tokens
@@ -255,9 +249,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
output_price = None output_price = None
if ppqai_model.pricing.api: if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get( output_price = ppqai_model.pricing.api.get("output_per_1M")
"output_per_1M"
)
elif ppqai_model.pricing.output_per_1M_tokens: elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens output_price = ppqai_model.pricing.output_per_1M_tokens
@@ -320,9 +312,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
return models return models
except Exception:
raise
async def on_upstream_error_redirect( async def on_upstream_error_redirect(
self, status_code: int, error_message: str self, status_code: int, error_message: str
) -> None: ) -> None:
@@ -443,9 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
) )
async with httpx.AsyncClient(timeout=30.0) as client: async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request( response = await _safe_read_request(client, "GET", url, headers=headers)
client, "GET", url, headers=headers
)
status_data = response.json() status_data = response.json()
is_paid = status_data.get("status") == "Settled" is_paid = status_data.get("status") == "Settled"
@@ -345,3 +345,23 @@ async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None:
assert raw_send.await_args is not None assert raw_send.await_args is not None
assert raw_send.await_args.args[1] is proofs assert raw_send.await_args.args[1] is proofs
assert raw_send.await_args.kwargs["amount"] == 1000 assert raw_send.await_args.kwargs["amount"] == 1000
def test_select_melt_proofs_stops_at_minimal_cover_when_over_budget() -> None:
from routstr.payment.lnurl import _select_melt_proofs
wallet = MagicMock()
wallet.get_fees_for_proofs = MagicMock(side_effect=lambda selected: len(selected))
proofs = [MagicMock(amount=600, reserved=False) for _ in range(3)]
selected, shortfall = _select_melt_proofs(
wallet,
proofs,
quote_amount=1000,
fee_reserve=0,
gross_budget=1000,
)
assert selected is None
assert shortfall == 2
assert wallet.get_fees_for_proofs.call_count == 2
+48
View File
@@ -132,6 +132,54 @@ async def test_wrapped_transport_failure_opens_central_cooldown() -> None:
MintRateGuard._guards.pop(mint_url, None) MintRateGuard._guards.pop(mint_url, None)
@pytest.mark.asyncio
async def test_timeout_retry_succeeds_without_opening_cooldown() -> None:
from routstr.core.settings import settings
mint_url = "https://retryable-timeout.test"
MintRateGuard._guards.pop(mint_url, None)
calls = 0
async def flaky() -> str:
nonlocal calls
calls += 1
if calls == 1:
raise httpx.ReadTimeout("first attempt stalled")
return "ok"
with (
patch.object(settings, "mint_retry_max_attempts", 2),
patch("routstr.mint.asyncio.sleep", AsyncMock()),
):
result = await run_mint_operation(flaky, mint_url=mint_url)
assert result == "ok"
assert calls == 2
assert MintRateGuard.get(mint_url).cooldown_remaining() == 0.0
MintRateGuard._guards.pop(mint_url, None)
@pytest.mark.asyncio
async def test_exhausted_timeout_retries_open_transport_cooldown() -> None:
from routstr.core.settings import settings
mint_url = "https://exhausted-timeout.test"
MintRateGuard._guards.pop(mint_url, None)
async def always_timeout() -> None:
raise httpx.ReadTimeout("stalled")
with (
patch.object(settings, "mint_retry_max_attempts", 1),
patch("routstr.mint.asyncio.sleep", AsyncMock()),
pytest.raises(httpx.TimeoutException),
):
await run_mint_operation(always_timeout, mint_url=mint_url)
assert MintRateGuard.get(mint_url).cooldown_remaining() > 29
MintRateGuard._guards.pop(mint_url, None)
async def test_guard_concurrency_change_preserves_cooldown_state() -> None: async def test_guard_concurrency_change_preserves_cooldown_state() -> None:
from routstr.core.settings import settings from routstr.core.settings import settings