mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +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:
|
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
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user