diff --git a/src/livepeer_gateway/__init__.py b/src/livepeer_gateway/__init__.py index f474df5..bc23a23 100644 --- a/src/livepeer_gateway/__init__.py +++ b/src/livepeer_gateway/__init__.py @@ -61,7 +61,7 @@ ) from .discovery import discover_orchestrators, discover_runners from .orch_info import get_orch_info -from .remote_signer import LivePaymentSession, PaymentSession +from .remote_signer import LivePaymentChallenge, LivePaymentSession, PaymentSession from .scope import start_scope from .selection import ( RunnerSelectionCursor, @@ -107,6 +107,7 @@ "LiveRunnerSession", "LiveRunnerSessionCallback", "LiveRunnerSessionEvent", + "LivePaymentChallenge", "LivePaymentSession", "LiveRunnerProxy", "LivepeerGatewayError", diff --git a/src/livepeer_gateway/http.py b/src/livepeer_gateway/http.py index 01c4b5d..699d042 100644 --- a/src/livepeer_gateway/http.py +++ b/src/livepeer_gateway/http.py @@ -305,9 +305,9 @@ async def _request_body( async def request_json( url: str, *, - method: Optional[str] = None, - payload: Optional[dict[str, Any]] = None, - headers: Optional[dict[str, str]] = None, + method: str | None = None, + payload: dict[str, Any] | None = None, + headers: dict[str, str] | None = None, timeout: float = 5.0, ) -> Any: """ @@ -403,6 +403,21 @@ async def get_json( return await request_json(url, headers=headers, timeout=timeout) +async def _post_empty( + url: str, + *, + headers: dict[str, str] | None = None, + timeout: float = 5.0, +) -> None: + """POST an empty body to ``url`` and discard the response.""" + await _request_body( + url, + method="POST", + headers=headers, + timeout=timeout, + ) + + def _parse_http_url(url: str, *, context: str = "URL") -> ParseResult: """ Normalize a URL for HTTP(S) endpoints. diff --git a/src/livepeer_gateway/live_runner.py b/src/livepeer_gateway/live_runner.py index 2d06836..deb981d 100644 --- a/src/livepeer_gateway/live_runner.py +++ b/src/livepeer_gateway/live_runner.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import contextlib import inspect import json import logging @@ -26,9 +27,10 @@ from .channel_reader import ChannelReader from .errors import LivepeerGatewayError, LivepeerHTTPError, SignerRefreshRequired -from .http import _request_body, open_stream, post_json, request_json +from .http import _post_empty, _request_body, open_stream, post_json, request_json from .remote_signer import ( GetPaymentResponse, + LivePaymentChallenge, LivePaymentSession, _freeze_headers, get_signer_info, @@ -46,6 +48,9 @@ "720p-pixel-seconds": "lv2v", "fixed": "fixed", } +# Metered types are billed for as long as the work runs, so they need ongoing +# payments. Fixed pricing is settled by the upfront payment alone. +_METERED_PAYMENT_TYPES = frozenset({"live", "lv2v"}) # golang format duration, eg "10s" _DURATION_RE = re.compile(r"^\s*(?P[0-9]+(?:\.[0-9]+)?)(?Pns|us|\u00b5s|ms|s|m|h)\s*$") @@ -102,12 +107,52 @@ class LiveRunnerInstance: price_info: LiveRunnerPriceInfo | None = None -@dataclass(frozen=True) +@dataclass class LiveRunnerSession: + """A reserved live runner session.""" + session_id: str app_url: str runner_url: str + control_url: str runner: LiveRunnerInstance | None = None + # True once the orchestrator reported this session gone. + released: bool = False + _payment_task: asyncio.Task[None] | None = field( + default=None, repr=False, compare=False + ) + + def __post_init__(self) -> None: + if not isinstance(self.control_url, str) or not self.control_url.strip(): + raise LivepeerGatewayError("Live runner session requires control_url") + self.control_url = self.control_url.strip() + _ = _join_endpoint(self.control_url, "stop") + + def _start_payments(self, payment_session: LivePaymentSession) -> None: + if self._payment_task is not None: + return + self._payment_task = _start_funding( + payment_session, + lambda: setattr(self, "released", True), + ) + + async def stop_payments(self) -> None: + """Stop funding this session without releasing it. + + Useful to hand funding to something else, or to let a session lapse + deliberately. Closing the session stops payments too. + """ + await _stop_funding(self._payment_task) + self._payment_task = None + + async def aclose(self) -> None: + await stop_runner_session(self) + + async def __aenter__(self) -> LiveRunnerSession: + return self + + async def __aexit__(self, *exc: object) -> None: + await self.aclose() @dataclass(frozen=True) @@ -128,7 +173,7 @@ class LiveRunnerCallResult: compare=False, ) # Non-JSON responses (an image, say) arrive unparsed in `content`; `data` stays empty. - content: Optional[bytes] = field(default=None, repr=False) + content: bytes | None = field(default=None, repr=False) content_type: str = "" @@ -148,6 +193,11 @@ class LiveRunnerCallStream: payment_session: LivePaymentSession | None _session: aiohttp.ClientSession = field(repr=False, compare=False) _response: aiohttp.ClientResponse = field(repr=False, compare=False) + # True once the orchestrator reported the backing session gone. + released: bool = False + _payment_task: asyncio.Task[None] | None = field( + default=None, repr=False, compare=False + ) @property def content_type(self) -> str: @@ -161,7 +211,21 @@ async def aiter_lines(self) -> AsyncIterator[str]: async for line in self._response.content: yield line.decode(errors="replace").rstrip("\n") + def _start_payments( + self, + payment_session: LivePaymentSession, + ) -> None: + if self._payment_task is not None: + return + self._payment_task = _start_funding( + payment_session, + lambda: setattr(self, "released", True), + ) + async def aclose(self) -> None: + # Stop funding first: don't pay for a stream we are about to drop. + await _stop_funding(self._payment_task) + self._payment_task = None self._response.release() await self._session.close() @@ -299,8 +363,8 @@ async def close(self) -> None: try: await _post_empty( _join_endpoint(self.orchestrator_url, f"/runners/{quote(self.runner_id, safe='')}/unregister"), - {"Authorization": secret}, - self._timeout, + headers={"Authorization": secret}, + timeout=self._timeout, ) except Exception: _LOG.debug("Live runner unregister failed", exc_info=True) @@ -757,12 +821,13 @@ async def call_runner( if signer_url: signer = await get_signer_info(signer_url, _freeze_headers(signer_headers)) payer_address = cast(str, signer.address) - challenge: _RunnerPaymentChallenge | None = None + challenge: LivePaymentChallenge | None = None attempts = (max(0, int(max_payment_challenge_retries)) + 1) * 2 for attempt in range(attempts): payment_session: LivePaymentSession | None = None payment_type = "" session_id = "" + needs_ongoing_funding = False # No preferred format: the app, or the upstream it fronts, picks. Only # control-plane calls ask for JSON. request_headers: dict[str, str] = {"Accept": "*/*"} @@ -792,6 +857,9 @@ async def call_runner( request_headers["Livepeer-Segment"] = payment.seg_creds or "" session_id = challenge.manifest_id + # Metered pricing bills for as long as the work runs. + needs_ongoing_funding = payment_type in _METERED_PAYMENT_TYPES + try: request_kwargs: dict[str, Any] = {"timeout": timeout} if request_headers: @@ -806,16 +874,35 @@ async def call_runner( payload=request_payload, headers=request_headers or None, ) - return LiveRunnerCallStream( - resp.status, resp.headers, runner_url, runner, payment_session, session, resp, + call_stream = LiveRunnerCallStream( + resp.status, + resp.headers, + runner_url, + runner, + None if payment_type == "fixed" else payment_session, + session, + resp, ) - - body, content_type = await _request_body( - runner_url, - method=method, - payload=request_payload, - **request_kwargs, + # The stream outlives this call, so it owns the funding. + if needs_ongoing_funding: + call_stream._start_payments(cast(LivePaymentSession, payment_session)) + return call_stream + + # The request ends with this call, so the funding ends with it. + pay_task = ( + _start_funding(cast(LivePaymentSession, payment_session)) + if needs_ongoing_funding + else None ) + try: + body, content_type = await _request_body( + runner_url, + method=method, + payload=request_payload, + **request_kwargs, + ) + finally: + await _stop_funding(pay_task) # Non-JSON bodies (an image, ndjson) are handed back unparsed in `content`. is_json = _is_json_content_type(content_type) data: dict[str, Any] = {} @@ -854,14 +941,7 @@ async def call_runner( raise LivepeerGatewayError("Live runner call exhausted payment challenge retries") -@dataclass(frozen=True) -class _RunnerPaymentChallenge: - payment_params: str - orchestrator_url: str - manifest_id: str - - -def _parse_runner_payment_challenge(error: LivepeerHTTPError) -> _RunnerPaymentChallenge: +def _parse_runner_payment_challenge(error: LivepeerHTTPError) -> LivePaymentChallenge: try: data = json.loads(error.body) except json.JSONDecodeError as e: @@ -870,24 +950,47 @@ def _parse_runner_payment_challenge(error: LivepeerHTTPError) -> _RunnerPaymentC raise LivepeerGatewayError("Live runner payment challenge response must be a JSON object") payment_params = data.get("payment_params") - orchestrator_url = data.get("orchestrator") manifest_id = data.get("manifest_id") + payment_url = data.get("payment_url") if not isinstance(payment_params, str) or not payment_params: raise LivepeerGatewayError("Live runner payment challenge missing payment_params") - if not isinstance(orchestrator_url, str) or not orchestrator_url: - raise LivepeerGatewayError("Live runner payment challenge missing orchestrator") if not isinstance(manifest_id, str) or not manifest_id: raise LivepeerGatewayError("Live runner payment challenge missing manifest_id") + if not isinstance(payment_url, str) or not payment_url: + raise LivepeerGatewayError("Live runner payment challenge missing payment_url") - return _RunnerPaymentChallenge( + return LivePaymentChallenge( payment_params=payment_params, - orchestrator_url=orchestrator_url, manifest_id=manifest_id, + payment_url=payment_url, ) +def _start_funding( + payment_session: LivePaymentSession, + on_released: Callable[[], None] | None = None, +) -> asyncio.Task[None]: + """Run payments in the background for as long as the caller keeps the task.""" + + async def _fund() -> None: + released = await payment_session.run_payments() + if released and on_released is not None: + on_released() + + return asyncio.create_task(_fund()) + + +async def _stop_funding(task: asyncio.Task[None] | None) -> None: + if task is None: + return + if not task.done(): + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + + async def _get_runner_payment( - challenge: _RunnerPaymentChallenge, + challenge: LivePaymentChallenge, *, payment_type: str, signer_url: str, @@ -897,9 +1000,7 @@ async def _get_runner_payment( signer_url=signer_url, signer_headers=signer_headers, type=payment_type, - payment_params=challenge.payment_params, - manifest_id=challenge.manifest_id, - orchestrator_url=challenge.orchestrator_url, + challenge=challenge, ) payment = await session.get_payment() if not payment.payment: @@ -957,25 +1058,6 @@ def _live_runner_price_info_from_json(value: object) -> LiveRunnerPriceInfo | No ) -def _live_runner_session_from_json( - data: dict[str, Any], - *, - runner_url: str, - runner: LiveRunnerInstance | None, -) -> LiveRunnerSession: - session_id = data.get("session_id") - app_url = data.get("app_url") - if not isinstance(session_id, str) or not session_id.strip(): - raise LivepeerGatewayError("Live runner session reserve response missing session_id") - if not isinstance(app_url, str) or not app_url.strip(): - raise LivepeerGatewayError("Live runner session reserve response missing app_url") - return LiveRunnerSession( - session_id=session_id.strip(), - app_url=app_url.strip(), - runner_url=runner_url, - runner=runner, - ) - async def stop_runner_session( session: LiveRunnerSession | LiveRunnerSessionRequest, *, @@ -983,13 +1065,12 @@ async def stop_runner_session( ) -> None: request_headers: dict[str, str] = {} if isinstance(session, LiveRunnerSession): - runner_url = session.runner_url.strip() - session_id = session.session_id.strip() - if not runner_url: - raise LivepeerGatewayError("Live runner session stop requires runner_url") - if not session_id: - raise LivepeerGatewayError("Live runner session stop requires session_id") - url = _join_endpoint(runner_url, f"/{quote(session_id, safe='')}/stop") + # This helper is also the public cleanup path, so callers that do not + # use aclose() still stop local funding before the remote reservation. + await session.stop_payments() + if session.released: + return + url = _join_endpoint(session.control_url, "stop") else: headers = getattr(session, "headers", None) get = getattr(headers, "get", None) @@ -1002,9 +1083,11 @@ async def stop_runner_session( request_headers = {"Livepeer-Session-Token": token} await _post_empty( url, - request_headers, - timeout, + headers=request_headers, + timeout=timeout, ) + if isinstance(session, LiveRunnerSession): + session.released = True def detect_process_gpu() -> LiveRunnerGPU | None: @@ -1174,25 +1257,6 @@ def _is_trickle_channel_response(value: object) -> bool: ) and ("internal_url" not in value or isinstance(value.get("internal_url"), str)) -async def _post_empty(url: str, headers: dict[str, str], timeout: float) -> None: - try: - client_timeout = aiohttp.ClientTimeout(total=timeout) - connector = aiohttp.TCPConnector(ssl=False) - async with aiohttp.ClientSession(timeout=client_timeout, connector=connector) as session: - async with session.post(url, data=b"", headers=headers) as resp: - body = await resp.text() - if resp.status >= 400: - raise LivepeerGatewayError( - f"HTTP empty POST error: HTTP {resp.status}; body={body!r}" - ) - except LivepeerGatewayError: - raise - except getattr(aiohttp, "ClientConnectorError", ()) as e: - raise LivepeerGatewayError(f"HTTP empty POST error: {getattr(e, 'message', e)}") from e - except (TimeoutError, aiohttp.ClientError) as e: - raise LivepeerGatewayError(f"HTTP empty POST error: {getattr(e, 'message', e)}") from e - - def _detect_gpu_pynvml() -> LiveRunnerGPU | None: try: import pynvml # type: ignore[import-not-found] diff --git a/src/livepeer_gateway/remote_signer.py b/src/livepeer_gateway/remote_signer.py index 0f43c6c..6505816 100644 --- a/src/livepeer_gateway/remote_signer.py +++ b/src/livepeer_gateway/remote_signer.py @@ -1,29 +1,46 @@ from __future__ import annotations +import asyncio import base64 import json import logging import re import ssl -from dataclasses import dataclass +from dataclasses import dataclass, replace from functools import lru_cache from typing import Any, Optional from urllib.error import HTTPError, URLError from urllib.request import Request, urlopen -import aiohttp - from . import lp_rpc_pb2 from .async_cache import async_lru_cache -from .errors import LivepeerGatewayError, PaymentError, SignerRefreshRequired +from .errors import ( + LivepeerGatewayError, + LivepeerHTTPError, + PaymentError, + SignerRefreshRequired, + SkipPaymentCycle, +) _LOG = logging.getLogger(__name__) +# Must stay under the signer's opening payment: 10s per-second, 60s pixel. +PAYMENT_INTERVAL_S = 3.0 + @dataclass(frozen=True) class GetPaymentResponse: payment: str seg_creds: Optional[str] = None +@dataclass(frozen=True) +class LivePaymentChallenge: + """The complete payment contract returned by a live-runner 402.""" + + payment_params: str + manifest_id: str + payment_url: str + + @dataclass(frozen=True) class SignerMaterial: """ @@ -202,19 +219,15 @@ def __init__( *, signer_headers: dict[str, str] | None = None, type: str, - payment_params: str, - manifest_id: str, - orchestrator_url: str | None = None, + challenge: LivePaymentChallenge, max_refresh_retries: int = 3, ) -> None: self._signer_url = signer_url self._signer_headers = _freeze_headers(signer_headers) self._type = type - self._payment_params = payment_params - self._manifest_id = manifest_id + self._challenge = challenge self._max_refresh_retries = max(0, int(max_refresh_retries)) self._state: dict[str, Any] | None = None - self._orchestrator_url = orchestrator_url async def get_payment(self) -> GetPaymentResponse: if not self._signer_url: @@ -231,61 +244,63 @@ async def get_payment(self) -> GetPaymentResponse: ) from e if self._state is None: raise - orchestrator_url = e.orchestrator_url - if not orchestrator_url: - raise PaymentError( - "Signer refresh response missing Livepeer-Orchestrator-URL header" - ) from e - await self._refresh_payment_params(orchestrator_url) + await self._refresh_payment_params() attempts += 1 - async def send_payment(self, orchestrator_url: str | None = None) -> None: + async def send_payment(self) -> None: + """Generate a payment and POST it to the challenge's endpoint. + + Raises LivepeerHTTPError on error responses so callers can branch on + the status code, and SkipPaymentCycle when the signer gates the cycle. + """ if not self._signer_url: return - target = orchestrator_url or self._orchestrator_url - if not target: - raise PaymentError("orchestrator_url is required before sending payment") - - from .http import _extract_error_message_from_body, _http_origin + from .http import _post_empty payment = await self.get_payment() - url = f"{_http_origin(target)}/payment" + if not payment.seg_creds: + # An empty segment header fails the orchestrator's sig check and + # comes back 403, which reads as a dead session, not a bad signer. + raise PaymentError("Signer returned a payment with no segCreds") headers = { "Livepeer-Payment": payment.payment, "Livepeer-Segment": payment.seg_creds, } - try: - timeout = aiohttp.ClientTimeout(total=5.0) - async with aiohttp.ClientSession(timeout=timeout) as session: - async with session.post(url, data=b"", headers=headers) as resp: - if resp.status >= 400: - body = await resp.text() - message = _extract_error_message_from_body(body) - body_part = f"; body={message!r}" if message else "" - raise PaymentError( - f"HTTP payment error: HTTP {resp.status} from endpoint (url={url}){body_part}" - ) - await resp.read() - except PaymentError: - raise - except getattr(aiohttp, "ClientConnectorError", ()) as e: - raise PaymentError( - f"HTTP payment error: failed to reach endpoint: {getattr(e, 'message', e)} (url={url})" - ) from e - except (aiohttp.ClientError, TimeoutError) as e: - raise PaymentError( - f"HTTP payment error: failed to reach endpoint: {getattr(e, 'message', e)} (url={url})" - ) from e + await _post_empty(self._challenge.payment_url, headers=headers, timeout=5.0) + + async def run_payments(self) -> bool: + """Keep a metered session funded until cancelled or the session ends. + + Cancel the task to stop; the first payment waits one interval, since + the caller pays upfront. Returns True if the orchestrator reports that + the challenge's session-scoped endpoint is gone. + """ + while True: + await asyncio.sleep(PAYMENT_INTERVAL_S) + try: + await self.send_payment() + except SkipPaymentCycle as e: + _LOG.debug("Payment loop skipped cycle: %s", e) + except LivepeerHTTPError as e: + # A 4xx will not change on a retry (404 gone, 409 fixed price, + # 403 mismatch), so stop rather than mint tickets nobody will + # honour. 408 and 429 are the two that do ask to be retried. + if 400 <= e.status_code < 500 and e.status_code not in (408, 429): + _LOG.info("Payment loop stopping (HTTP %d): %s", e.status_code, e) + return e.status_code == 404 + _LOG.warning("Payment failed; retrying next cycle: %s", e) + except Exception as e: + _LOG.warning("Payment failed; retrying next cycle: %s", e) async def _payment_request(self) -> GetPaymentResponse: from .http import _http_origin, post_json url = f"{_http_origin(self._signer_url)}/generate-live-payment" payload: dict[str, Any] = { - "orchestrator": self._payment_params, + "orchestrator": self._challenge.payment_params, "type": self._type, - "ManifestID": self._manifest_id, + "ManifestID": self._challenge.manifest_id, } if self._state is not None: payload["state"] = self._state @@ -313,19 +328,19 @@ async def _payment_request(self) -> GetPaymentResponse: self._state = state return GetPaymentResponse(payment=payment, seg_creds=seg_creds) - async def _refresh_payment_params(self, orchestrator_url: str) -> None: + async def _refresh_payment_params(self) -> None: from .http import _http_origin, post_json signer = await get_signer_info(self._signer_url or "", self._signer_headers) if not signer.address: raise PaymentError("Cannot refresh payment without signer address") - url = f"{_http_origin(orchestrator_url)}/refresh-payment" + url = f"{_http_origin(self._challenge.payment_url)}/refresh-payment" data = await post_json( url, { "sender": signer.address, - "manifest_id": self._manifest_id, + "manifest_id": self._challenge.manifest_id, }, ) payment_params = data.get("payment_params") @@ -333,12 +348,11 @@ async def _refresh_payment_params(self, orchestrator_url: str) -> None: raise PaymentError( f"RefreshPayment error: missing/invalid 'payment_params' in response (url={url})" ) - self._payment_params = payment_params - refreshed_orchestrator_url = data.get("orchestrator") - self._orchestrator_url = ( - refreshed_orchestrator_url - if isinstance(refreshed_orchestrator_url, str) and refreshed_orchestrator_url.strip() - else orchestrator_url + # Refresh rotates the embedded payment material. The initial scoped + # endpoint remains authoritative for the lifetime of this session. + self._challenge = replace( + self._challenge, + payment_params=payment_params, ) diff --git a/src/livepeer_gateway/selection.py b/src/livepeer_gateway/selection.py index 968dedf..9a5d47e 100644 --- a/src/livepeer_gateway/selection.py +++ b/src/livepeer_gateway/selection.py @@ -296,17 +296,26 @@ async def reserve_session( result = await cursor.next() session_id = result.data.get("session_id") app_url = result.data.get("app_url") + control_url = _string_value(result.data.get("control_url")) if not isinstance(session_id, str) or not session_id.strip(): raise LivepeerGatewayError("runner session response missing session_id") if not isinstance(app_url, str) or not app_url.strip(): raise LivepeerGatewayError("runner session response missing app_url") - return LiveRunnerSession( + if not control_url: + raise LivepeerGatewayError("runner session response missing control_url") + session = LiveRunnerSession( session_id=session_id.strip(), app_url=app_url.strip(), runner_url=result.runner_url, + control_url=control_url, runner=result.runner, ) + # No payment session means fixed price or offchain: nothing to fund. + if result.payment_session is not None: + session._start_payments(result.payment_session) + return session + def _runner_candidates_from_discovery(entries: Sequence[dict[str, Any]]) -> list[LiveRunnerInstance]: candidates: list[LiveRunnerInstance] = [] diff --git a/tests/test_live_payment_session.py b/tests/test_live_payment_session.py index 094571f..65dda6a 100644 --- a/tests/test_live_payment_session.py +++ b/tests/test_live_payment_session.py @@ -7,16 +7,28 @@ from livepeer_gateway import lp_rpc_pb2 from livepeer_gateway.errors import ( - PaymentError, + LivepeerHTTPError, SignerRefreshRequired, ) from livepeer_gateway.remote_signer import ( + LivePaymentChallenge, LivePaymentSession, PaymentSession, get_signer_info, ) +_PAYMENT_URL = "https://orch.example.com/apps/runner/session/manifest-1/payment" + + +def _challenge(*, payment_params: str = "opaque") -> LivePaymentChallenge: + return LivePaymentChallenge( + payment_params=payment_params, + manifest_id="manifest-1", + payment_url=_PAYMENT_URL, + ) + + class TestPaymentSession: def test_get_payment_round_trips_state_without_cross_session_leak(self) -> None: calls: list[tuple[str, dict[str, object], dict[str, str] | None]] = [] @@ -87,55 +99,23 @@ async def test_none_signer_exits_early(self) -> None: session = LivePaymentSession( None, type="lv2v", - payment_params="opaque", - manifest_id="manifest-1", + challenge=_challenge(), ) payment = await session.get_payment() - await session.send_payment("https://orchestrator.example.com") + await session.send_payment() assert payment.payment == "" assert payment.seg_creds is None - async def test_send_payment_uses_default_tls_verification(self) -> None: - class _Response: - status = 204 - headers: dict[str, str] = {} - - async def __aenter__(self) -> _Response: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - async def read(self) -> bytes: - return b"" - - async def text(self) -> str: - raise AssertionError( - "successful payment responses must not be decoded as text" - ) - - class _Session: - def __init__(self, **kwargs: object) -> None: - self.kwargs = kwargs - - async def __aenter__(self) -> _Session: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - def post(self, *args: object, **kwargs: object) -> _Response: - return _Response() - + async def test_send_payment_reuses_empty_post_helper(self) -> None: session = LivePaymentSession( "https://signer.example.com", type="lv2v", - payment_params="opaque", - manifest_id="manifest-1", + challenge=_challenge(), ) + post_empty = mock.AsyncMock() with ( mock.patch.object( session, @@ -144,171 +124,29 @@ def post(self, *args: object, **kwargs: object) -> _Response: return_value=types.SimpleNamespace(payment="p", seg_creds="s") ), ), - mock.patch( - "livepeer_gateway.remote_signer.aiohttp.TCPConnector" - ) as connector_mock, - mock.patch( - "livepeer_gateway.remote_signer.aiohttp.ClientSession", - side_effect=_Session, - ) as client_session_mock, - ): - await session.send_payment("https://orchestrator.example.com") - - connector_mock.assert_not_called() - assert "connector" not in client_session_mock.call_args.kwargs - - async def test_send_payment_uses_constructor_orchestrator_url(self) -> None: - posts: list[tuple[object, dict[str, object]]] = [] - - class _Response: - status = 204 - headers: dict[str, str] = {} - - async def __aenter__(self) -> _Response: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - async def read(self) -> bytes: - return b"" - - class _Session: - def __init__(self, **kwargs: object) -> None: - del kwargs - - async def __aenter__(self) -> _Session: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - def post(self, url: object, **kwargs: object) -> _Response: - posts.append((url, kwargs)) - return _Response() - - session = LivePaymentSession( - "https://signer.example.com", - type="lv2v", - payment_params="opaque", - manifest_id="manifest-1", - orchestrator_url="https://orchestrator.example.com/base", - ) - - with ( - mock.patch.object( - session, - "get_payment", - new=mock.AsyncMock( - return_value=types.SimpleNamespace(payment="p", seg_creds="s") - ), - ), - mock.patch( - "livepeer_gateway.remote_signer.aiohttp.ClientSession", - side_effect=_Session, - ), + mock.patch("livepeer_gateway.http._post_empty", post_empty), ): await session.send_payment() - assert posts[0][0] == "https://orchestrator.example.com/payment" - assert posts[0][1]["headers"] == { - "Livepeer-Payment": "p", - "Livepeer-Segment": "s", - } - - async def test_send_payment_accepts_binary_payment_result_response(self) -> None: - class _Response: - status = 200 - headers: dict[str, str] = {} - - async def __aenter__(self) -> _Response: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - async def read(self) -> bytes: - return b"\x82\x01protobuf-payment-result" - - async def text(self) -> str: - raise AssertionError( - "successful binary payment response decoded as text" - ) - - class _Session: - def __init__(self, **kwargs: object) -> None: - del kwargs - - async def __aenter__(self) -> _Session: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - def post(self, *args: object, **kwargs: object) -> _Response: - del args, kwargs - return _Response() - - session = LivePaymentSession( - "https://signer.example.com", - type="lv2v", - payment_params="opaque", - manifest_id="manifest-1", + post_empty.assert_awaited_once_with( + _PAYMENT_URL, + headers={"Livepeer-Payment": "p", "Livepeer-Segment": "s"}, + timeout=5.0, ) - with ( - mock.patch.object( - session, - "get_payment", - new=mock.AsyncMock( - return_value=types.SimpleNamespace(payment="p", seg_creds="s") - ), - ), - mock.patch( - "livepeer_gateway.remote_signer.aiohttp.ClientSession", - side_effect=_Session, - ), - ): - await session.send_payment("https://orchestrator.example.com") - - async def test_send_payment_error_decodes_body_for_message(self) -> None: - class _Response: - status = 400 - headers: dict[str, str] = {} - - async def __aenter__(self) -> _Response: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - async def read(self) -> bytes: - raise AssertionError("error payment responses should use text decoding") - - async def text(self) -> str: - return '{"error":{"message":"payment rejected"}}' - - class _Session: - def __init__(self, **kwargs: object) -> None: - del kwargs - - async def __aenter__(self) -> _Session: - return self - - async def __aexit__(self, *args: object) -> None: - return None - - def post(self, *args: object, **kwargs: object) -> _Response: - del args, kwargs - return _Response() - + async def test_send_payment_preserves_typed_http_error(self) -> None: session = LivePaymentSession( "https://signer.example.com", type="lv2v", - payment_params="opaque", - manifest_id="manifest-1", + challenge=_challenge(), ) + error = LivepeerHTTPError( + 400, + "https://orchestrator.example.com/payment", + body='{"error":{"message":"payment rejected"}}', + message="payment rejected", + ) with ( mock.patch.object( session, @@ -318,14 +156,14 @@ def post(self, *args: object, **kwargs: object) -> _Response: ), ), mock.patch( - "livepeer_gateway.remote_signer.aiohttp.ClientSession", - side_effect=_Session, + "livepeer_gateway.http._post_empty", + new=mock.AsyncMock(side_effect=error), ), ): - with pytest.raises(PaymentError) as raised: - await session.send_payment("https://orchestrator.example.com") + with pytest.raises(LivepeerHTTPError) as raised: + await session.send_payment() - assert "payment rejected" in str(raised.value) + assert raised.value is error async def test_get_payment_sends_opaque_payment_params_and_state(self) -> None: calls: list[tuple[str, dict[str, object], dict[str, str] | None]] = [] @@ -350,8 +188,7 @@ async def _post_json( "https://signer.example.com", signer_headers={"Authorization": "token"}, type="lv2v", - payment_params="opaque-payment-params", - manifest_id="manifest-1", + challenge=_challenge(payment_params="opaque-payment-params"), ) first = await session.get_payment() second = await session.get_payment() @@ -388,8 +225,7 @@ async def _post_json( session = LivePaymentSession( "https://signer.example.com", type="lv2v", - payment_params="old-payment-params", - manifest_id="manifest-1", + challenge=_challenge(payment_params="old-payment-params"), ) with pytest.raises(SignerRefreshRequired): await session.get_payment() @@ -405,7 +241,7 @@ async def _post_json( ) ] - async def test_stateful_480_refreshes_payment_params_from_orchestrator_header( + async def test_stateful_480_refreshes_params_from_payment_url_origin( self, ) -> None: calls: list[tuple[str, dict[str, object]]] = [] @@ -430,10 +266,7 @@ async def _post_json( "state": {"state": "one"}, } if payment_requests == 2: - raise SignerRefreshRequired( - "refresh", - orchestrator_url="https://orch.example.com", - ) + raise SignerRefreshRequired("refresh") return { "payment": "payment-2", "segCreds": "segment-2", @@ -445,6 +278,8 @@ async def _post_json( return { "payment_params": "new-payment-params", "orchestrator": "https://orch.example.com", + "manifest_id": "manifest-1", + "payment_url": "https://orch.example.com/payment", } raise AssertionError(f"unexpected POST {url}") @@ -452,8 +287,7 @@ async def _post_json( session = LivePaymentSession( "https://signer.example.com", type="lv2v", - payment_params="old-payment-params", - manifest_id="manifest-1", + challenge=_challenge(payment_params="old-payment-params"), ) first_payment = await session.get_payment() payment = await session.get_payment() @@ -467,37 +301,7 @@ async def _post_json( ) assert calls[4][1]["orchestrator"] == "new-payment-params" - async def test_480_without_orchestrator_header_fails(self) -> None: - payment_requests = 0 - - async def _post_json( - url: str, - payload: dict[str, object], - *, - headers: dict[str, str] | None = None, - timeout: float = 5.0, - ) -> dict[str, object]: - nonlocal payment_requests - del url, payload, headers, timeout - payment_requests += 1 - if payment_requests == 1: - return { - "payment": "payment-1", - "segCreds": "segment-1", - "state": {"state": "one"}, - } - raise SignerRefreshRequired("refresh") - - with mock.patch("livepeer_gateway.http.post_json", side_effect=_post_json): - session = LivePaymentSession( - "https://signer.example.com", - type="lv2v", - payment_params="old-payment-params", - manifest_id="manifest-1", - ) - await session.get_payment() - with pytest.raises(PaymentError, match="missing Livepeer-Orchestrator-URL"): - await session.get_payment() + assert session._challenge.payment_url == _PAYMENT_URL async def test_get_signer_info_caches_result(self) -> None: calls: list[tuple[str, dict[str, object]]] = [] diff --git a/tests/test_live_runner.py b/tests/test_live_runner.py index e85d1f2..e2ca113 100644 --- a/tests/test_live_runner.py +++ b/tests/test_live_runner.py @@ -26,6 +26,7 @@ stop_runner_session, create_proxy, ) +from livepeer_gateway.remote_signer import LivePaymentChallenge class TestLiveRunnerHelpers: @@ -43,6 +44,38 @@ def test_join_endpoint_preserves_base_path(self) -> None: == "https://orch.example.com:8935/base/runners/heartbeat" ) + def test_payment_challenge_uses_server_supplied_url(self) -> None: + body = json.dumps( + { + "payment_params": "opaque-payment-params", + "manifest_id": "manifest-1", + "payment_url": _payment_url("manifest-1"), + } + ) + + challenge = live_runner._parse_runner_payment_challenge( + LivepeerHTTPError(402, "https://runner.example.com", body) + ) + + assert challenge == _payment_challenge("manifest-1") + + @pytest.mark.parametrize("payment_url", [None, ""]) + def test_payment_challenge_requires_payment_url( + self, payment_url: str | None + ) -> None: + body = json.dumps( + { + "payment_params": "opaque-payment-params", + "manifest_id": "manifest-1", + "payment_url": payment_url, + } + ) + + with pytest.raises(LivepeerGatewayError, match="missing payment_url"): + live_runner._parse_runner_payment_challenge( + LivepeerHTTPError(402, "https://runner.example.com", body) + ) + def test_parse_go_duration(self) -> None: assert live_runner._parse_go_duration_s("500ms", default=5.0) == 0.5 assert live_runner._parse_go_duration_s("5s", default=1.0) == 5.0 @@ -132,6 +165,7 @@ def _post_empty(url: str, headers: dict[str, str], timeout: float) -> None: session_id="session-1", app_url="https://service.example.com/app", runner_url="https://service.example.com/apps/runner-1/session", + control_url="https://service.example.com/apps/runner-1/session/session-1", ) with mock.patch.object(live_runner, "_post_empty", side_effect=_post_empty): @@ -288,9 +322,7 @@ def _request_body( "signer_url": "https://signer.example.com", "signer_headers": {"Authorization": "token"}, "type": "live", - "payment_params": "opaque-payment-params", - "manifest_id": "manifest-1", - "orchestrator_url": "https://orchestrator.example.com", + "challenge": _payment_challenge("manifest-1"), } ] assert result.payment_session is payment_sessions[0] @@ -359,9 +391,7 @@ def _request_body( "signer_url": "https://signer.example.com", "signer_headers": None, "type": "lv2v", - "payment_params": "opaque-payment-params", - "manifest_id": "manifest-scope", - "orchestrator_url": "https://orchestrator.example.com", + "challenge": _payment_challenge("manifest-scope"), } ] @@ -438,8 +468,8 @@ def _request_body( assert len(sessions) == 2 assert sessions[0]["type"] == "fixed" assert sessions[1]["type"] == "fixed" - assert sessions[0]["manifest_id"] == "fixed-manifest" - assert sessions[1]["manifest_id"] == "fixed-manifest" + assert sessions[0]["challenge"] == _payment_challenge("fixed-manifest") + assert sessions[1]["challenge"] == _payment_challenge("fixed-manifest") assert result.payment_session is None async def test_paid_call_restarts_challenge_when_signer_requests_refresh( @@ -532,17 +562,13 @@ def _request_body( "signer_url": "https://signer.example.com", "signer_headers": None, "type": "live", - "payment_params": "opaque-payment-params", - "manifest_id": "manifest-1", - "orchestrator_url": "https://orchestrator.example.com", + "challenge": _payment_challenge("manifest-1"), }, { "signer_url": "https://signer.example.com", "signer_headers": None, "type": "live", - "payment_params": "opaque-payment-params", - "manifest_id": "manifest-2", - "orchestrator_url": "https://orchestrator.example.com", + "challenge": _payment_challenge("manifest-2"), }, ] assert sig_mock.call_count == 1 @@ -614,10 +640,26 @@ def _payment_challenge_body(manifest_id: str) -> str: "payment_params": "opaque-payment-params", "orchestrator": "https://orchestrator.example.com", "manifest_id": manifest_id, + "payment_url": _payment_url(manifest_id), } ) +def _payment_challenge(manifest_id: str) -> LivePaymentChallenge: + return LivePaymentChallenge( + payment_params="opaque-payment-params", + manifest_id=manifest_id, + payment_url=_payment_url(manifest_id), + ) + + +def _payment_url(manifest_id: str) -> str: + return ( + "https://orchestrator.example.com/apps/runner-1/session/" + f"{manifest_id}/payment" + ) + + def _json_data(data: dict[str, object]) -> tuple[bytes, str]: return json.dumps(data).encode("utf-8"), "application/json" diff --git a/tests/test_live_runner_payments.py b/tests/test_live_runner_payments.py new file mode 100644 index 0000000..c6fab96 --- /dev/null +++ b/tests/test_live_runner_payments.py @@ -0,0 +1,302 @@ +from __future__ import annotations + +import asyncio +from unittest import mock + +import pytest + +from livepeer_gateway import live_runner, remote_signer, selection +from livepeer_gateway.errors import ( + LivepeerGatewayError, + LivepeerHTTPError, + SkipPaymentCycle, +) +from livepeer_gateway.live_runner import LiveRunnerCallResult, LiveRunnerSession +from livepeer_gateway.remote_signer import LivePaymentChallenge, LivePaymentSession + + +_CONTROL_URL = "https://orch.example.com/apps/runner-1/session/session-1" +_PAYMENT_URL = f"{_CONTROL_URL}/payment" + + +def _http_error(status: int) -> LivepeerHTTPError: + return LivepeerHTTPError(status, _PAYMENT_URL) + + +def _live_payment_session() -> LivePaymentSession: + return LivePaymentSession( + "https://signer.example.com", + type="live", + challenge=LivePaymentChallenge( + payment_params="opaque", + manifest_id="session-1", + payment_url=_PAYMENT_URL, + ), + ) + + +class _FundingSession: + def __init__(self, *, released: bool = False) -> None: + self.released = released + self.started = asyncio.Event() + self.cancelled = asyncio.Event() + + async def run_payments(self) -> bool: + self.started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + self.cancelled.set() + raise + return self.released + + +def _session(*, control_url: str = _CONTROL_URL) -> LiveRunnerSession: + return LiveRunnerSession( + session_id="session-1", + app_url=f"{_CONTROL_URL}/app", + runner_url="https://orch.example.com/apps/runner-1/session", + control_url=control_url, + ) + + +class TestPaymentLoop: + @pytest.mark.parametrize( + "status, released", [(403, False), (404, True), (409, False)] + ) + async def test_terminal_status_stops_loop( + self, status: int, released: bool + ) -> None: + payment_session = _live_payment_session() + with ( + mock.patch.object( + payment_session, + "send_payment", + new=mock.AsyncMock(side_effect=_http_error(status)), + ) as send_payment, + mock.patch.object(remote_signer, "PAYMENT_INTERVAL_S", 0), + ): + result = await asyncio.wait_for( + payment_session.run_payments(), + timeout=1.0, + ) + + assert result is released + send_payment.assert_awaited_once_with() + + @pytest.mark.parametrize( + "first_error", + [RuntimeError("network"), _http_error(408), SkipPaymentCycle("paid up")], + ) + async def test_retryable_error_reaches_next_cycle( + self, first_error: Exception + ) -> None: + payment_session = _live_payment_session() + send_payment = mock.AsyncMock(side_effect=[first_error, _http_error(404)]) + with ( + mock.patch.object(payment_session, "send_payment", new=send_payment), + mock.patch.object(remote_signer, "PAYMENT_INTERVAL_S", 0), + ): + released = await asyncio.wait_for( + payment_session.run_payments(), + timeout=1.0, + ) + + assert released + assert send_payment.await_count == 2 + + +class TestSessionPaymentLifecycle: + @pytest.mark.parametrize("control_url", ["", "ftp://orch/session/session-1"]) + def test_session_rejects_missing_or_invalid_control_url( + self, control_url: str + ) -> None: + with pytest.raises(LivepeerGatewayError): + _session(control_url=control_url) + + async def test_start_payments_starts_challenge_owned_session(self) -> None: + payment_session = _FundingSession() + session = _session() + + session._start_payments(payment_session) # type: ignore[arg-type] + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + + await session.stop_payments() + + async def test_start_payments_is_idempotent(self) -> None: + payment_session = _FundingSession() + session = _session() + + session._start_payments(payment_session) # type: ignore[arg-type] + task = session._payment_task + session._start_payments(payment_session) # type: ignore[arg-type] + + assert session._payment_task is task + await session.stop_payments() + + async def test_stop_payments_cancels_and_clears_task(self) -> None: + payment_session = _FundingSession() + session = _session() + session._start_payments(payment_session) # type: ignore[arg-type] + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + + await session.stop_payments() + + assert session._payment_task is None + assert payment_session.cancelled.is_set() + + async def test_aclose_delegates_to_stop_runner_session(self) -> None: + session = _session() + + stop = mock.AsyncMock() + with mock.patch.object(live_runner, "stop_runner_session", stop): + await session.aclose() + + stop.assert_awaited_once_with(session) + + async def test_async_context_manager_closes_session(self) -> None: + session = _session() + stop = mock.AsyncMock() + + with mock.patch.object(live_runner, "stop_runner_session", stop): + async with session as entered: + assert entered is session + + stop.assert_awaited_once_with(session) + + +class TestStopRunnerSession: + async def test_stops_payment_task_and_uses_control_url(self) -> None: + payment_session = _FundingSession() + session = _session() + session._start_payments(payment_session) # type: ignore[arg-type] + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + post_empty = mock.AsyncMock() + + with mock.patch.object(live_runner, "_post_empty", post_empty): + await live_runner.stop_runner_session(session, timeout=12.0) + + assert payment_session.cancelled.is_set() + assert session._payment_task is None + assert session.released + post_empty.assert_awaited_once_with( + f"{_CONTROL_URL}/stop", + headers={}, + timeout=12.0, + ) + + async def test_remote_stop_failure_still_stops_payment_task(self) -> None: + payment_session = _FundingSession() + session = _session() + session._start_payments(payment_session) # type: ignore[arg-type] + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + + with ( + mock.patch.object( + live_runner, + "_post_empty", + new=mock.AsyncMock(side_effect=LivepeerGatewayError("stop failed")), + ), + pytest.raises(LivepeerGatewayError, match="stop failed"), + ): + await live_runner.stop_runner_session(session) + + assert payment_session.cancelled.is_set() + assert session._payment_task is None + assert not session.released + + async def test_already_released_session_only_stops_local_funding(self) -> None: + payment_session = _FundingSession() + session = _session() + session._start_payments(payment_session) # type: ignore[arg-type] + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + session.released = True + post_empty = mock.AsyncMock() + + with mock.patch.object(live_runner, "_post_empty", post_empty): + await live_runner.stop_runner_session(session) + + assert payment_session.cancelled.is_set() + post_empty.assert_not_awaited() + + +class _Cursor: + def __init__(self, *results: LiveRunnerCallResult) -> None: + self.results = list(results) + self.rejections = [] + + async def next(self) -> LiveRunnerCallResult: + if self.results: + return self.results.pop(0) + raise AssertionError("unexpected extra runner selection") + + +def _reservation( + name: str, + *, + control_url: str | None, + payment_session: object | None = None, +) -> LiveRunnerCallResult: + data = { + "session_id": f"session-{name}", + "app_url": f"https://orch.example.com/session-{name}/app", + } + if control_url is not None: + data["control_url"] = control_url + return LiveRunnerCallResult( + data, + runner_url=f"https://orch.example.com/runner-{name}/session", + payment_session=payment_session, # type: ignore[arg-type] + ) + + +class TestReservationSelection: + async def test_paid_reservation_starts_scoped_funding(self) -> None: + payment_session = _FundingSession() + cursor = _Cursor( + _reservation( + "1", + control_url=_CONTROL_URL, + payment_session=payment_session, + ) + ) + + with mock.patch.object( + selection, "runner_selector", new=mock.AsyncMock(return_value=cursor) + ): + session = await selection.reserve_session() + + await asyncio.wait_for(payment_session.started.wait(), timeout=1.0) + await session.stop_payments() + + @pytest.mark.parametrize( + "bad_control_url", + [None, "ftp://orch.example.com/session/bad"], + ) + async def test_invalid_control_url_fails_immediately( + self, bad_control_url: str | None + ) -> None: + cursor = _Cursor( + _reservation("bad", control_url=bad_control_url), + _reservation("good", control_url=_CONTROL_URL), + ) + + with mock.patch.object( + selection, + "runner_selector", + new=mock.AsyncMock(return_value=cursor), + ): + with pytest.raises(LivepeerGatewayError): + await selection.reserve_session() + + assert len(cursor.results) == 1 + + async def test_missing_control_url_is_contract_error(self) -> None: + cursor = _Cursor(_reservation("bad", control_url=None)) + with mock.patch.object( + selection, + "runner_selector", + new=mock.AsyncMock(return_value=cursor), + ): + with pytest.raises(LivepeerGatewayError, match="missing control_url"): + await selection.reserve_session() diff --git a/tests/test_selection.py b/tests/test_selection.py index b7d3140..e56d371 100644 --- a/tests/test_selection.py +++ b/tests/test_selection.py @@ -86,7 +86,11 @@ async def _call_runner( ) -> LiveRunnerCallResult: calls.append((runner.url, payload, method, timeout)) return LiveRunnerCallResult( - {"session_id": "session-1", "app_url": "https://orch-a/apps/a/app"}, + { + "session_id": "session-1", + "app_url": "https://orch-a/apps/a/app", + "control_url": "https://orch-a/apps/a/session/session-1", + }, runner_url=runner.url, runner=runner, session_id="session-1", @@ -336,7 +340,11 @@ async def _call_runner( ) -> LiveRunnerCallResult: del payload, method, timeout return LiveRunnerCallResult( - {"session_id": "session-1", "app_url": "https://orch-a/apps/a/app"}, + { + "session_id": "session-1", + "app_url": "https://orch-a/apps/a/app", + "control_url": "https://orch-a/apps/a/session/session-1", + }, runner_url=runner.url, runner=runner, session_id="session-1",