"""Platform-neutral ProcessingRunner with bounded retries and strict result acceptance.""" from __future__ import annotations import hashlib import time from dataclasses import dataclass from typing import Callable, Optional, Protocol from arr_ingestion.contracts import OPAQUE_ID_RE from arr_processing.contracts import ProcessingRequest, ProcessingResult from arr_processing.errors import ProcessingError, ProcessingTransportError from arr_processing.ledger import ProcessingLedger, ProcessingRecord from arr_processing.signatures import SignedResultCodec ACTIVE_REMOTE_STATUSES = frozenset({"pending", "queued", "running", "cancelling"}) SUCCESS_REMOTE_STATUSES = frozenset({"success", "completed", "succeeded"}) FAILED_REMOTE_STATUSES = frozenset({"failed", "error"}) CANCELLED_REMOTE_STATUSES = frozenset({"cancelled", "canceled", "interrupted"}) ALL_REMOTE_STATUSES = ( ACTIVE_REMOTE_STATUSES | SUCCESS_REMOTE_STATUSES | FAILED_REMOTE_STATUSES | CANCELLED_REMOTE_STATUSES ) def processing_dispatch_idempotency_key(request: ProcessingRequest) -> str: """Return the one persisted dispatch identity for a processing attempt.""" request_sha = request.sha256() return hashlib.sha256( (request_sha + "\x1fdispatch").encode("ascii") ).hexdigest() @dataclass(frozen=True) class RemoteRunSnapshot: remote_run_id: str status: str final_content: Optional[str] = None class ProcessingTransport(Protocol): def submit(self, request: ProcessingRequest, idempotency_key: str) -> RemoteRunSnapshot: ... def get(self, request: ProcessingRequest, remote_run_id: str) -> RemoteRunSnapshot: ... def cancel(self, request: ProcessingRequest, remote_run_id: str) -> RemoteRunSnapshot: ... @dataclass(frozen=True) class ProcessingCompletion: result: ProcessingResult callback_replayed: bool class ProcessingRunner: def __init__( self, transport: ProcessingTransport, ledger: ProcessingLedger, result_codec: SignedResultCodec, *, max_transport_attempts: int = 3, retry_delay_seconds: float = 0.25, sleep: Callable[[float], None] = time.sleep, monotonic: Callable[[], float] = time.monotonic, ) -> None: if max_transport_attempts < 1 or retry_delay_seconds < 0: raise ValueError("runner retry policy is invalid") self._transport = transport self._ledger = ledger self._codec = result_codec self._max_transport_attempts = max_transport_attempts self._retry_delay_seconds = retry_delay_seconds self._sleep = sleep self._monotonic = monotonic def start(self, request: ProcessingRequest) -> ProcessingRecord: request_sha = request.sha256() idempotency_key = processing_dispatch_idempotency_key(request) record = self._ledger.reserve(request, request_sha, idempotency_key) if record.remote_run_id is not None: return record if record.remote_status == "failed": raise ProcessingError( record.failure_code or "PROCESSING_ATTEMPT_FAILED", "processing attempt already failed before remote dispatch", ) try: snapshot = self._with_transport_retries( lambda: self._transport.submit(request, idempotency_key) ) except ProcessingTransportError as error: if error.active_run_conflict: concurrent = self._ledger.get(request.job_id, request.attempt_no) if concurrent is not None and concurrent.remote_run_id is not None: return concurrent raise ProcessingError( "PROCESSING_ACTIVE_RUN_CONFLICT", "remote session already has an uncorrelated active run", retryable=True, ) from None concurrent = self._ledger.get(request.job_id, request.attempt_no) if concurrent is not None and concurrent.remote_run_id is not None: return concurrent self._ledger.fail_start( request.job_id, request.attempt_no, error.code, ) self._raise_transport(error) try: self._validate_snapshot(snapshot) except ProcessingError as error: self._ledger.fail_start( request.job_id, request.attempt_no, error.code, ) raise return self._ledger.bind_run( request.job_id, request.attempt_no, snapshot.remote_run_id, snapshot.status, ) def poll( self, job_id: str, attempt_no: int, *, poll_interval_seconds: float = 1.0, timeout_seconds: float = 120.0, ) -> ProcessingCompletion: if poll_interval_seconds <= 0 or timeout_seconds <= 0: raise ValueError("poll timing must be greater than zero") record = self._required_record(job_id, attempt_no) if record.result is not None: return ProcessingCompletion(record.result, callback_replayed=True) if record.remote_run_id is None: raise ProcessingError( "PROCESSING_RUN_NOT_STARTED", "processing attempt has no remote run" ) deadline = self._monotonic() + timeout_seconds while True: try: snapshot = self._with_transport_retries( lambda: self._transport.get(record.request, str(record.remote_run_id)) ) except ProcessingTransportError as error: self._raise_transport(error) self._validate_snapshot(snapshot, expected_run_id=record.remote_run_id) self._ledger.update_status(job_id, attempt_no, snapshot.status) if snapshot.status in SUCCESS_REMOTE_STATUSES: if not isinstance(snapshot.final_content, str) or not snapshot.final_content.strip(): self._ledger.update_status(job_id, attempt_no, "delivery_missing") raise ProcessingError( "PROCESSING_RESULT_MISSING", "remote run completed without an authenticated result", retryable=True, ) return self.accept_callback(snapshot.final_content.encode("utf-8")) if snapshot.status in FAILED_REMOTE_STATUSES: raise ProcessingError( "PROCESSING_REMOTE_FAILED", "remote processing run failed" ) if snapshot.status in CANCELLED_REMOTE_STATUSES: raise ProcessingError( "PROCESSING_REMOTE_CANCELLED", "remote processing run was cancelled" ) now = self._monotonic() if now >= deadline: raise ProcessingError( "PROCESSING_POLL_TIMEOUT", "processing run did not finish before the polling deadline", retryable=True, ) self._sleep(min(poll_interval_seconds, max(0.0, deadline - now))) def accept_callback(self, raw_signed_result: bytes) -> ProcessingCompletion: verified = self._codec.verify(raw_signed_result) result = verified.result record = self._required_record(result.job_id, result.attempt_no) request = record.request if ( record.remote_run_id != result.remote_run_id or request.processor_version != result.processor_version or request.rule_set_sha256 != result.rule_set_sha256 ): raise ProcessingError( "PROCESSING_RESULT_MISMATCH", "processing result does not match its request" ) callback_replayed = self._ledger.accept_callback( result.delivery_id, verified.signed_sha256, ) _updated, result_replayed = self._ledger.attach_result(result) return ProcessingCompletion( result=result, callback_replayed=callback_replayed or result_replayed, ) def cancel(self, job_id: str, attempt_no: int) -> ProcessingRecord: record = self._required_record(job_id, attempt_no) if record.remote_run_id is None: raise ProcessingError( "PROCESSING_RUN_NOT_STARTED", "processing attempt has no remote run" ) try: snapshot = self._with_transport_retries( lambda: self._transport.cancel(record.request, str(record.remote_run_id)) ) except ProcessingTransportError as error: self._raise_transport(error) self._validate_snapshot(snapshot, expected_run_id=record.remote_run_id) return self._ledger.update_status(job_id, attempt_no, snapshot.status) def _with_transport_retries(self, operation: Callable[[], RemoteRunSnapshot]) -> RemoteRunSnapshot: last_error: Optional[ProcessingTransportError] = None for attempt in range(1, self._max_transport_attempts + 1): try: return operation() except ProcessingTransportError as error: last_error = error if error.active_run_conflict or not error.retryable or attempt >= self._max_transport_attempts: raise self._sleep(self._retry_delay_seconds * attempt) assert last_error is not None raise last_error @staticmethod def _validate_snapshot( snapshot: RemoteRunSnapshot, expected_run_id: Optional[str] = None, ) -> None: if ( not isinstance(snapshot, RemoteRunSnapshot) or not OPAQUE_ID_RE.fullmatch(snapshot.remote_run_id) or snapshot.status not in ALL_REMOTE_STATUSES or (expected_run_id is not None and snapshot.remote_run_id != expected_run_id) or ( snapshot.final_content is not None and not isinstance(snapshot.final_content, str) ) ): raise ProcessingError( "PROCESSING_REMOTE_PROTOCOL_INVALID", "remote run response is invalid" ) def _required_record(self, job_id: str, attempt_no: int) -> ProcessingRecord: record = self._ledger.get(job_id, attempt_no) if record is None: raise ProcessingError( "PROCESSING_JOB_NOT_FOUND", "processing attempt is not registered" ) return record @staticmethod def _raise_transport(error: ProcessingTransportError) -> None: raise ProcessingError( error.code, "remote processing transport did not complete", retryable=error.retryable, ) from None