diff --git a/routstr/mint.py b/routstr/mint.py index b09b497c..e1b60ca2 100644 --- a/routstr/mint.py +++ b/routstr/mint.py @@ -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 diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 35f29f2a..def5bca7 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -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: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 8606f201..8a8b06ea 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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) diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 0eec7bfe..65ccbb57 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -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" diff --git a/tests/unit/test_lnurl_amount_and_destination.py b/tests/unit/test_lnurl_amount_and_destination.py index 97737cf7..e722141a 100644 --- a/tests/unit/test_lnurl_amount_and_destination.py +++ b/tests/unit/test_lnurl_amount_and_destination.py @@ -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 diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py index dc58d10c..f6873fa4 100644 --- a/tests/unit/test_mint.py +++ b/tests/unit/test_mint.py @@ -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