mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
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:
+5
-4
@@ -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
|
||||
|
||||
@@ -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:
|
||||
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
@@ -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)
|
||||
|
||||
@@ -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,11 +208,8 @@ 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
|
||||
)
|
||||
response = await _safe_read_request(client, "GET", url, headers=headers)
|
||||
data = response.json()
|
||||
|
||||
models_data = data.get("data", [])
|
||||
@@ -244,9 +240,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
if or_model:
|
||||
input_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M"
|
||||
)
|
||||
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
|
||||
|
||||
@@ -255,9 +249,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
output_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M"
|
||||
)
|
||||
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
|
||||
|
||||
@@ -320,9 +312,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
return models
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
async def on_upstream_error_redirect(
|
||||
self, status_code: int, error_message: str
|
||||
) -> None:
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user