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:
raise
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:
# Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait.
backoff = (2**attempt) + (time.monotonic() % 1.0)
logger.warning(
"Mint operation timed out, retrying",
@@ -323,6 +320,10 @@ async def run_mint_operation(
)
await asyncio.sleep(backoff)
continue
if guard is not None:
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise httpx.TimeoutException(
f"{op_name} timed out (attempts: {attempt + 1})"
) from exc
+8 -15
View File
@@ -9,8 +9,8 @@ from cashu.core.base import MeltQuoteState
from cashu.wallet.wallet import Proof, Wallet
from ..mint import (
MINT_TRANSPORT_EXCEPTIONS,
is_mint_rate_limited,
is_mint_transport_error,
run_mint_operation,
)
@@ -106,17 +106,6 @@ async def _fetch_lnurl_json(
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:
"""Decode LNURL to get the actual URL.
@@ -260,8 +249,11 @@ def _select_melt_proofs(
selected_amount += proof.amount
input_fees = int(wallet.get_fees_for_proofs(selected))
required = quote_amount + fee_reserve + input_fees
if required <= gross_budget and selected_amount >= required:
return selected, 0
if selected_amount >= required:
if required <= gross_budget:
return selected, 0
# Covered but over budget; more proofs only raise input fees.
break
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:
await on_melt_quote(melt_quote_resp.quote)
assert selected_proofs is not None
proofs = selected_proofs
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.
await wallet.set_reserved_for_send(proofs, reserved=False)
raise
if not _contains_mint_transport_error(error):
if not is_mint_transport_error(error):
raise
# Cashu clears reservations on transport errors despite an unknown outcome.
try:
+36 -23
View File
@@ -7,7 +7,7 @@ import math
import traceback
import typing
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
import httpx
@@ -82,6 +82,21 @@ async def _aclose_if_needed(resource: object | None) -> None:
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]
@@ -1049,9 +1064,7 @@ class BaseUpstreamProvider:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
return
try:
@@ -1318,14 +1331,15 @@ class BaseUpstreamProvider:
)
raise
finally:
try:
if not usage_finalized:
await finalize_db_only()
finally:
try:
await _aclose_if_needed(response)
finally:
await _aclose_if_needed(client)
# Shielded so a client disconnect cannot cancel billing
# finalization or leak the upstream connection.
await asyncio.shield(
_finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
client,
)
)
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
@@ -1521,9 +1535,7 @@ class BaseUpstreamProvider:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
return
try:
@@ -1750,14 +1762,15 @@ class BaseUpstreamProvider:
)
raise
finally:
try:
if not usage_finalized:
await finalize_db_only()
finally:
try:
await _aclose_if_needed(response)
finally:
await _aclose_if_needed(client)
# Shielded so a client disconnect cannot cancel billing
# finalization or leak the upstream connection.
await asyncio.shield(
_finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
client,
)
)
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
+95 -108
View File
@@ -40,7 +40,8 @@ _ppq_circuits: dict[str, _PPQCircuitState] = {}
def _ppq_origin(url: str) -> str:
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(
@@ -66,9 +67,7 @@ async def _safe_read_request(
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
try:
response = await client.request(
method, url, headers=headers, json=json
)
response = await client.request(method, url, headers=headers, json=json)
response.raise_for_status()
state.consecutive_failures = 0
state.cooldown_until = 0.0
@@ -209,119 +208,109 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
url = f"{self.base_url}/models"
headers = {"Authorization": f"Bearer {self.api_key}"}
try:
async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request(
client, "GET", url, headers=headers
)
data = response.json()
async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request(client, "GET", url, headers=headers)
data = response.json()
models_data = data.get("data", [])
models_data = data.get("data", [])
or_models = [
Model(**model) # type: ignore
for model in await async_fetch_openrouter_models()
]
or_models = [
Model(**model) # type: ignore
for model in await async_fetch_openrouter_models()
]
models = []
for model_data in models_data:
try:
ppqai_model = PPQAIModel.parse_obj(model_data)
if ppqai_model.id in self.IGNORED_MODEL_IDS:
continue
models = []
for model_data in models_data:
try:
ppqai_model = PPQAIModel.parse_obj(model_data)
if ppqai_model.id in self.IGNORED_MODEL_IDS:
continue
or_model = next(
(
model
for model in or_models
if (model.id == ppqai_model.id)
or (model.id.split("/")[-1] == ppqai_model.id)
or (model.id == ppqai_model.id.split("/")[-1])
),
None,
)
or_model = next(
(
model
for model in or_models
if (model.id == ppqai_model.id)
or (model.id.split("/")[-1] == ppqai_model.id)
or (model.id == ppqai_model.id.split("/")[-1])
),
None,
)
if or_model:
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M"
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
if or_model:
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get("input_per_1M")
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
if input_price is not None:
or_model.pricing.prompt = input_price / 1_000_000
if input_price is not None:
or_model.pricing.prompt = input_price / 1_000_000
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M"
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get("output_per_1M")
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else:
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append(
Model(
id=ppqai_model.id,
name=ppqai_model.name,
created=ppqai_model.created_at // 1000,
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
context_length=ppqai_model.context_length,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
)
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else:
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append(
Model(
id=ppqai_model.id,
name=ppqai_model.name,
created=ppqai_model.created_at // 1000,
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
context_length=ppqai_model.context_length,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
)
except Exception as e:
logger.warning(
"Failed to parse PPQ.AI model",
extra={
"model_id": model_data.get("id", "unknown"),
"error": str(e),
"error_type": type(e).__name__,
},
)
except Exception as e:
logger.warning(
"Failed to parse PPQ.AI model",
extra={
"model_id": model_data.get("id", "unknown"),
"error": str(e),
"error_type": type(e).__name__,
},
)
return models
except Exception:
raise
return models
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
@@ -443,9 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request(
client, "GET", url, headers=headers
)
response = await _safe_read_request(client, "GET", url, headers=headers)
status_data = response.json()
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.args[1] is proofs
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)
@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:
from routstr.core.settings import settings