Files
wyndham-ARR/arr_processing/postgres.py
2026-07-29 16:38:05 +08:00

637 lines
23 KiB
Python

"""PostgreSQL processing state for restart-safe Agent dispatch and callbacks."""
from __future__ import annotations
import hashlib
import os
from dataclasses import dataclass
from datetime import date
from typing import Any, Callable, Optional, Tuple
from arr_ingestion.contracts import ArtifactRef
from arr_ingestion.repository import IngestionOutcome
from arr_processing.contracts import ProcessingRequest, ProcessingResult
from arr_processing.errors import ProcessingError
from arr_processing.ledger import ProcessingRecord
TARGET_DATABASE = "booking_test"
DATABASE_ENV = "ARR_DATABASE_URL"
_RUN_TERMINAL = frozenset({"accepted", "rejected", "failed", "cancelled"})
_SAFE_START_FAILURE_CODES = frozenset(
{
"PROCESSING_REMOTE_PROTOCOL_INVALID",
"PROCESSING_REMOTE_RATE_LIMITED",
"PROCESSING_REMOTE_REJECTED",
"PROCESSING_REMOTE_UNAVAILABLE",
}
)
def _default_connect(dsn: str) -> Any:
try:
import psycopg # type: ignore[import-not-found]
except ImportError:
raise ProcessingError(
"DATABASE_DRIVER_UNAVAILABLE",
"PostgreSQL driver is unavailable",
retryable=True,
) from None
try:
return psycopg.connect(dsn, autocommit=False)
except Exception:
raise ProcessingError(
"DATABASE_UNAVAILABLE",
"processing database is unavailable",
retryable=True,
) from None
@dataclass(frozen=True)
class ProcessingDatabaseConfig:
dsn: str
@classmethod
def from_environment(cls) -> "ProcessingDatabaseConfig":
dsn = os.environ.get(DATABASE_ENV, "").strip()
if not dsn:
raise ProcessingError(
"DATABASE_CONFIG_MISSING",
f"{DATABASE_ENV} is required",
)
return cls(dsn)
class PostgresProcessingState:
"""Implements ProcessingLedger plus source and committed-outcome lookup."""
def __init__(
self,
config: ProcessingDatabaseConfig,
connect: Optional[Callable[[str], Any]] = None,
) -> None:
self._config = config
self._connect = connect or _default_connect
def _run(self, operation: Callable[[Any], Any], *, read_only: bool = False) -> Any:
try:
connection = self._connect(self._config.dsn)
except ProcessingError:
raise
except Exception:
raise ProcessingError(
"DATABASE_UNAVAILABLE",
"processing database is unavailable",
retryable=True,
) from None
try:
with connection.transaction():
with connection.cursor() as cursor:
cursor.execute(
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ"
+ (" READ ONLY" if read_only else "")
)
cursor.execute("SET LOCAL statement_timeout = '15s'")
cursor.execute("SET LOCAL lock_timeout = '5s'")
cursor.execute("SELECT current_database()")
row = cursor.fetchone()
if not row or row[0] != TARGET_DATABASE:
raise ProcessingError(
"DATABASE_TARGET_INVALID",
"processing database target is invalid",
)
return operation(cursor)
except ProcessingError:
raise
except Exception:
raise ProcessingError(
"DATABASE_WRITE_FAILED" if not read_only else "DATABASE_QUERY_FAILED",
"processing state could not be persisted"
if not read_only
else "processing state could not be read",
retryable=True,
) from None
finally:
connection.close()
def reserve(
self,
request: ProcessingRequest,
request_sha256: str,
idempotency_key: str,
) -> ProcessingRecord:
def operation(cursor: Any) -> ProcessingRecord:
row = self._attempt_row(
cursor,
request.job_id,
request.attempt_no,
for_update=True,
)
if (
row[1] != "opera_daily"
or row[2] != request.processor_version
or str(row[3]) != request.rule_set_sha256
or str(row[5]) != idempotency_key
):
raise ProcessingError(
"PROCESSING_JOB_CONFLICT",
"processing attempt identity conflicts",
)
return ProcessingRecord(
request=request,
request_sha256=request_sha256,
idempotency_key=idempotency_key,
remote_run_id=str(row[6]) if row[6] is not None else None,
remote_status=self._remote_status(str(row[7])),
failure_code=str(row[9]) if row[9] is not None else None,
)
return self._run(operation)
def get(self, job_id: str, attempt_no: int) -> Optional[ProcessingRecord]:
def operation(cursor: Any) -> Optional[ProcessingRecord]:
try:
row = self._attempt_row(cursor, job_id, attempt_no)
except ProcessingError as error:
if error.code == "PROCESSING_JOB_NOT_FOUND":
return None
raise
placeholder = "persisted_" + hashlib.sha256(
f"{job_id}\x1f{attempt_no}".encode("utf-8")
).hexdigest()[:48]
request = ProcessingRequest(
job_id=job_id,
attempt_no=attempt_no,
source_file_id=placeholder,
processor_version=str(row[2]),
rule_set_sha256=str(row[3]),
)
return ProcessingRecord(
request=request,
request_sha256=request.sha256(),
idempotency_key=str(row[5]),
remote_run_id=str(row[6]) if row[6] is not None else None,
remote_status=self._remote_status(str(row[7])),
failure_code=str(row[9]) if row[9] is not None else None,
)
return self._run(operation, read_only=True)
def bind_run(
self,
job_id: str,
attempt_no: int,
remote_run_id: str,
remote_status: str,
) -> ProcessingRecord:
return self._set_remote_status(
job_id,
attempt_no,
remote_status,
remote_run_id=remote_run_id,
)
def update_status(
self,
job_id: str,
attempt_no: int,
remote_status: str,
) -> ProcessingRecord:
return self._set_remote_status(job_id, attempt_no, remote_status)
def fail_start(
self,
job_id: str,
attempt_no: int,
failure_code: str,
) -> ProcessingRecord:
if failure_code not in _SAFE_START_FAILURE_CODES:
failure_code = "PROCESSING_REMOTE_PROTOCOL_INVALID"
def operation(cursor: Any) -> ProcessingRecord:
row = self._attempt_row(cursor, job_id, attempt_no, for_update=True)
if row[6] is not None:
raise ProcessingError(
"PROCESSING_RUN_CONFLICT",
"processing attempt already has a remote run",
)
if str(row[4]) not in _RUN_TERMINAL:
cursor.execute(
"""
UPDATE ingestion.processing_attempts
SET attempt_status = 'failed',
failure_code = %s,
failure_message = 'remote processing submission did not complete',
finished_at = now()
WHERE id = %s
""",
(failure_code, row[8]),
)
cursor.execute(
"""
UPDATE ingestion.processing_runs
SET run_status = 'failed',
failure_code = %s,
failure_message = 'remote processing submission did not complete',
updated_at = now(),
finished_at = now()
WHERE id = %s
""",
(failure_code, row[0]),
)
persisted_failure_code = failure_code
else:
persisted_failure_code = (
str(row[9]) if row[9] is not None else failure_code
)
return self._record_from_row(
row,
job_id=job_id,
attempt_no=attempt_no,
remote_status="failed",
failure_code=persisted_failure_code,
)
return self._run(operation)
def attach_result(self, result: ProcessingResult) -> Tuple[ProcessingRecord, bool]:
def operation(cursor: Any) -> Tuple[ProcessingRecord, bool]:
row = self._attempt_row(
cursor,
result.job_id,
result.attempt_no,
for_update=True,
)
if (
row[6] != result.remote_run_id
or row[2] != result.processor_version
or str(row[3]) != result.rule_set_sha256
):
raise ProcessingError(
"PROCESSING_RESULT_MISMATCH",
"processing result does not match its request",
)
replayed = self._delivery_matches(cursor, result)
if str(row[4]) not in _RUN_TERMINAL:
cursor.execute(
"""
UPDATE ingestion.processing_attempts
SET attempt_status = 'delivered',
failure_code = NULL,
failure_message = NULL
WHERE id = %s
""",
(row[8],),
)
cursor.execute(
"""
UPDATE ingestion.processing_runs
SET run_status = 'validating',
updated_at = now()
WHERE id = %s
""",
(row[0],),
)
record = self._record_from_row(row, remote_status="delivered", result=result)
return record, replayed
return self._run(operation)
def accept_callback(self, delivery_id: str, signed_sha256: str) -> bool:
del signed_sha256
def operation(cursor: Any) -> bool:
cursor.execute(
"""
SELECT 1
FROM ingestion.processing_deliveries
WHERE delivery_key = %s
""",
(delivery_id,),
)
return cursor.fetchone() is not None
return self._run(operation, read_only=True)
def source_for_attempt(self, job_id: str, attempt_no: int) -> ArtifactRef:
def operation(cursor: Any) -> ArtifactRef:
cursor.execute(
"""
SELECT
artifact.object_key,
artifact.original_filename,
artifact.sha256,
artifact.byte_size,
artifact.mime_type
FROM ingestion.processing_runs AS run
JOIN ingestion.processing_attempts AS attempt
ON attempt.processing_run_id = run.id
AND attempt.attempt_no = %s
JOIN ingestion.artifacts AS artifact
ON artifact.id = run.source_artifact_id
WHERE run.run_key = %s
AND run.pipeline_type = 'opera_daily'
AND artifact.artifact_kind = 'opera_xml'
""",
(attempt_no, job_id),
)
row = cursor.fetchone()
if row is None:
raise ProcessingError(
"PROCESSING_JOB_NOT_FOUND",
"processing attempt is not registered",
)
return ArtifactRef.from_dict(
"source_xml",
{
"object_key": str(row[0]),
"original_filename": str(row[1]),
"sha256": str(row[2]),
"byte_size": int(row[3]),
"mime_type": str(row[4]),
},
)
return self._run(operation, read_only=True)
def outcome_for_delivery(self, delivery_id: str) -> Optional[IngestionOutcome]:
def operation(cursor: Any) -> Optional[IngestionOutcome]:
cursor.execute(
"""
SELECT
delivery.delivery_status,
run.run_key,
delivery.business_date,
delivery.daily_version_id,
version.version_no
FROM ingestion.processing_deliveries AS delivery
JOIN ingestion.processing_runs AS run
ON run.id = delivery.processing_run_id
LEFT JOIN finance.daily_versions AS version
ON version.id = delivery.daily_version_id
WHERE delivery.delivery_key = %s
""",
(delivery_id,),
)
row = cursor.fetchone()
if row is None or row[0] not in {"committed", "recorded_failure"}:
return None
return IngestionOutcome(
status=(
"already_committed"
if row[0] == "committed"
else "recorded_failure"
),
job_id=str(row[1]),
business_date=row[2] if isinstance(row[2], date) else None,
daily_version_id=int(row[3]) if row[3] is not None else None,
version_no=int(row[4]) if row[4] is not None else None,
)
return self._run(operation, read_only=True)
def _set_remote_status(
self,
job_id: str,
attempt_no: int,
remote_status: str,
*,
remote_run_id: Optional[str] = None,
) -> ProcessingRecord:
attempt_status, run_status, failure_code, terminal = self._database_status(
remote_status
)
def operation(cursor: Any) -> ProcessingRecord:
row = self._attempt_row(cursor, job_id, attempt_no, for_update=True)
persisted_run_id = str(row[6]) if row[6] is not None else None
if remote_run_id is not None:
cursor.execute(
"""
SELECT processing_run_id
FROM ingestion.processing_attempts
WHERE remote_run_id = %s
""",
(remote_run_id,),
)
owner = cursor.fetchone()
if owner is not None and int(owner[0]) != int(row[0]):
raise ProcessingError(
"PROCESSING_RUN_CONFLICT",
"remote run is already correlated",
)
if persisted_run_id is not None and persisted_run_id != remote_run_id:
raise ProcessingError(
"PROCESSING_RUN_CONFLICT",
"processing attempt has a different remote run",
)
persisted_run_id = remote_run_id
if persisted_run_id is None:
raise ProcessingError(
"PROCESSING_RUN_NOT_STARTED",
"processing attempt has no remote run",
)
if str(row[4]) not in _RUN_TERMINAL:
cursor.execute(
"""
UPDATE ingestion.processing_attempts
SET remote_run_id = %s,
attempt_status = %s,
failure_code = %s,
failure_message = CASE
WHEN %s::text IS NULL THEN NULL
ELSE 'remote processing did not complete'
END,
started_at = CASE
WHEN %s IN ('dispatched', 'running')
THEN COALESCE(started_at, now())
ELSE started_at
END,
finished_at = CASE WHEN %s THEN now() ELSE NULL END
WHERE id = %s
""",
(
persisted_run_id,
attempt_status,
failure_code,
failure_code,
attempt_status,
terminal,
row[8],
),
)
cursor.execute(
"""
UPDATE ingestion.processing_runs
SET run_status = %s,
failure_code = %s,
failure_message = CASE
WHEN %s::text IS NULL THEN NULL
ELSE 'remote processing did not complete'
END,
updated_at = now(),
finished_at = CASE WHEN %s THEN now() ELSE NULL END
WHERE id = %s
""",
(run_status, failure_code, failure_code, terminal, row[0]),
)
return self._record_from_row(
row,
job_id=job_id,
attempt_no=attempt_no,
remote_run_id=persisted_run_id,
remote_status=remote_status,
)
return self._run(operation)
@staticmethod
def _database_status(remote_status: str) -> Tuple[str, str, Optional[str], bool]:
if remote_status in {"pending", "queued"}:
return "dispatched", "queued", None, False
if remote_status in {"running", "cancelling"}:
return "running", "running", None, False
if remote_status in {"success", "completed", "succeeded"}:
return "delivered", "validating", None, False
if remote_status in {"failed", "error"}:
return "failed", "failed", "PROCESSING_REMOTE_FAILED", True
if remote_status in {"cancelled", "canceled", "interrupted"}:
return "cancelled", "cancelled", None, True
raise ProcessingError(
"PROCESSING_REMOTE_PROTOCOL_INVALID",
"remote processing status is invalid",
)
@staticmethod
def _remote_status(attempt_status: str) -> str:
return {
"queued": "queued",
"dispatched": "pending",
"running": "running",
"delivered": "success",
"succeeded": "success",
"failed": "failed",
"cancelled": "cancelled",
}.get(attempt_status, "failed")
@staticmethod
def _attempt_row(
cursor: Any,
job_id: str,
attempt_no: int,
*,
for_update: bool = False,
) -> Tuple[Any, ...]:
cursor.execute(
"""
SELECT
run.id,
run.pipeline_type,
run.requested_processor_version,
run.requested_rule_set_sha256,
run.run_status,
attempt.idempotency_key,
attempt.remote_run_id,
attempt.attempt_status,
attempt.id,
attempt.failure_code
FROM ingestion.processing_runs AS run
JOIN ingestion.processing_attempts AS attempt
ON attempt.processing_run_id = run.id
AND attempt.attempt_no = %s
WHERE run.run_key = %s
"""
+ (" FOR UPDATE OF run, attempt" if for_update else ""),
(attempt_no, job_id),
)
row = cursor.fetchone()
if row is None:
raise ProcessingError(
"PROCESSING_JOB_NOT_FOUND",
"processing attempt is not registered",
)
return tuple(row)
@staticmethod
def _record_from_row(
row: Tuple[Any, ...],
*,
remote_status: str,
job_id: Optional[str] = None,
attempt_no: Optional[int] = None,
remote_run_id: Optional[str] = None,
result: Optional[ProcessingResult] = None,
failure_code: Optional[str] = None,
) -> ProcessingRecord:
resolved_job_id = result.job_id if result is not None else job_id
resolved_attempt_no = result.attempt_no if result is not None else attempt_no
if resolved_job_id is None or resolved_attempt_no is None:
raise ProcessingError(
"PROCESSING_JOB_NOT_FOUND",
"processing attempt identity is unavailable",
)
request = ProcessingRequest(
job_id=resolved_job_id,
attempt_no=resolved_attempt_no,
source_file_id="persisted-source",
processor_version=str(row[2]),
rule_set_sha256=str(row[3]),
)
return ProcessingRecord(
request=request,
request_sha256=request.sha256(),
idempotency_key=str(row[5]),
remote_run_id=remote_run_id
if remote_run_id is not None
else (str(row[6]) if row[6] is not None else None),
remote_status=remote_status,
failure_code=(
failure_code
if failure_code is not None
else (str(row[9]) if row[9] is not None else None)
),
result=result,
)
@staticmethod
def _delivery_matches(cursor: Any, result: ProcessingResult) -> bool:
cursor.execute(
"""
SELECT
run.run_key,
attempt.attempt_no,
attempt.remote_run_id,
delivery.result_status,
delivery.business_date,
delivery.processor_version,
delivery.rule_set_sha256,
delivery.result_schema_version
FROM ingestion.processing_deliveries AS delivery
JOIN ingestion.processing_runs AS run
ON run.id = delivery.processing_run_id
LEFT JOIN ingestion.processing_attempts AS attempt
ON attempt.id = delivery.attempt_id
WHERE delivery.delivery_key = %s
""",
(result.delivery_id,),
)
row = cursor.fetchone()
if row is None:
return False
expected_date = result.business_date
if (
row[0] != result.job_id
or int(row[1]) != result.attempt_no
or row[2] != result.remote_run_id
or row[3] != result.status
or row[4] != expected_date
or row[5] != result.processor_version
or str(row[6]) != result.rule_set_sha256
or row[7] != result.result_schema_version
):
raise ProcessingError(
"PROCESSING_CALLBACK_CONFLICT",
"delivery identifier conflicts",
)
return True