feat: prepare ARR for controlled public deployment
This commit is contained in:
272
arr_processing/runner.py
Normal file
272
arr_processing/runner.py
Normal file
@@ -0,0 +1,272 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user