feat(cloud): add durable cancellation support to task repository
- Add nullable cancel_requested_at column (migration 0012) - Widen ScheduledTaskStatus/TerminalTaskStatus to include cancelled - Add CancellationRequestStatus + request_task_cancellation() to CloudRepository protocol and SQLAlchemy implementation - renew_lease() now returns LeaseRenewalResult, surfacing whether cancellation is pending, instead of a bare status string - reap_expired_leases() resolves pending-cancellation tasks to cancelled instead of requeuing/failing them - record_task_result() accepts cancelled and clears cancel_requested_at on any terminal write Note: internal_api/api.py's renew_assignment route still compares renew_lease()'s return value against a bare string; it will be updated in the next task (Internal Host<->Cloud protocol) to consume LeaseRenewalResult and populate the new cancel_requested wire field.
This commit is contained in:
@@ -40,7 +40,7 @@ from cloud.observability import current_correlation_id
|
||||
from core.models import utc_now
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cloud.repository import AssignmentProgressSnapshot
|
||||
from cloud.repository import AssignmentProgressSnapshot, LeaseRenewalResult
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -1497,7 +1497,9 @@ class SQLAlchemyCloudRepository:
|
||||
lease_expires_at: datetime,
|
||||
now: datetime,
|
||||
progress: AssignmentProgressSnapshot | None = None,
|
||||
) -> str:
|
||||
) -> "LeaseRenewalResult":
|
||||
from cloud.repository import LeaseRenewalResult
|
||||
|
||||
with self._sessions.begin() as session:
|
||||
task = session.get(
|
||||
ScheduledTaskRow,
|
||||
@@ -1505,19 +1507,19 @@ class SQLAlchemyCloudRepository:
|
||||
with_for_update=self.engine.dialect.name == "postgresql",
|
||||
)
|
||||
if task is None:
|
||||
return "not_found"
|
||||
return LeaseRenewalResult(status="not_found")
|
||||
if (
|
||||
task.status not in {"assigned", "dispatched"}
|
||||
or task.attempt_count != attempt
|
||||
or task.lease_id != lease_id
|
||||
or task.assigned_host_id != host_id
|
||||
):
|
||||
return "conflict"
|
||||
return LeaseRenewalResult(status="conflict")
|
||||
current_expiry = _parse_dt(task.lease_expires_at)
|
||||
if current_expiry is None or current_expiry <= now:
|
||||
return "expired"
|
||||
return LeaseRenewalResult(status="expired")
|
||||
if lease_expires_at <= now:
|
||||
return "conflict"
|
||||
return LeaseRenewalResult(status="conflict")
|
||||
|
||||
attempt_row = session.get(
|
||||
TaskAttemptRow,
|
||||
@@ -1530,7 +1532,7 @@ class SQLAlchemyCloudRepository:
|
||||
or attempt_row.lease_id != lease_id
|
||||
or attempt_row.host_id != host_id
|
||||
):
|
||||
return "conflict"
|
||||
return LeaseRenewalResult(status="conflict")
|
||||
|
||||
renewed_until = _iso(lease_expires_at)
|
||||
task.lease_expires_at = renewed_until
|
||||
@@ -1542,7 +1544,10 @@ class SQLAlchemyCloudRepository:
|
||||
task.progress_summary = progress.summary
|
||||
task.progress_updated_at = _iso(progress.updated_at)
|
||||
_log_task_lifecycle("renewed", task)
|
||||
return "renewed"
|
||||
return LeaseRenewalResult(
|
||||
status="renewed",
|
||||
cancel_requested=task.cancel_requested_at is not None,
|
||||
)
|
||||
|
||||
def record_task_result(
|
||||
self,
|
||||
@@ -1579,7 +1584,7 @@ class SQLAlchemyCloudRepository:
|
||||
):
|
||||
return "conflict"
|
||||
|
||||
if task.status in {"done", "failed"}:
|
||||
if task.status in {"done", "failed", "cancelled"}:
|
||||
if (
|
||||
task.status == status
|
||||
and task.failure_reason == failure_reason
|
||||
@@ -1593,7 +1598,7 @@ class SQLAlchemyCloudRepository:
|
||||
current_expiry = _parse_dt(task.lease_expires_at)
|
||||
if current_expiry is None or current_expiry <= completed_at:
|
||||
return "conflict"
|
||||
if status not in {"done", "failed"}:
|
||||
if status not in {"done", "failed", "cancelled"}:
|
||||
return "conflict"
|
||||
|
||||
result_json = (
|
||||
@@ -1610,11 +1615,18 @@ class SQLAlchemyCloudRepository:
|
||||
task.progress_step_status = None
|
||||
task.progress_summary = None
|
||||
task.progress_updated_at = None
|
||||
task.cancel_requested_at = None
|
||||
attempt_row.status = status
|
||||
attempt_row.completed_at = completed_at_iso
|
||||
attempt_row.failure_reason = failure_reason
|
||||
attempt_row.result_json = result_json
|
||||
_log_task_lifecycle("completed" if status == "done" else "failed", task)
|
||||
if status == "done":
|
||||
lifecycle_event = "completed"
|
||||
elif status == "cancelled":
|
||||
lifecycle_event = "cancelled"
|
||||
else:
|
||||
lifecycle_event = "failed"
|
||||
_log_task_lifecycle(lifecycle_event, task)
|
||||
return "recorded"
|
||||
|
||||
def reap_expired_leases(
|
||||
@@ -1657,7 +1669,12 @@ class SQLAlchemyCloudRepository:
|
||||
task.lease_id = None
|
||||
task.lease_expires_at = None
|
||||
task.result_json = None
|
||||
if task.attempt_count < max_attempts:
|
||||
if task.cancel_requested_at is not None:
|
||||
task.status = "cancelled"
|
||||
task.assigned_host_id = None
|
||||
task.assigned_device_id = None
|
||||
task.cancel_requested_at = None
|
||||
elif task.attempt_count < max_attempts:
|
||||
task.status = "queued"
|
||||
task.assigned_host_id = None
|
||||
task.assigned_device_id = None
|
||||
@@ -1667,13 +1684,46 @@ class SQLAlchemyCloudRepository:
|
||||
task.failure_reason = (
|
||||
f"lease expired after {task.attempt_count} attempts"
|
||||
)
|
||||
_log_task_lifecycle(
|
||||
"retried" if task.status == "queued" else "failed",
|
||||
task,
|
||||
)
|
||||
if task.status == "queued":
|
||||
lifecycle_event = "retried"
|
||||
elif task.status == "cancelled":
|
||||
lifecycle_event = "cancelled"
|
||||
else:
|
||||
lifecycle_event = "failed"
|
||||
_log_task_lifecycle(lifecycle_event, task)
|
||||
reaped_task_ids.append(task.id)
|
||||
return reaped_task_ids
|
||||
|
||||
def request_task_cancellation(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
requested_at: datetime,
|
||||
) -> str:
|
||||
with self._sessions.begin() as session:
|
||||
task = session.get(
|
||||
ScheduledTaskRow,
|
||||
task_id,
|
||||
with_for_update=self.engine.dialect.name == "postgresql",
|
||||
)
|
||||
if task is None:
|
||||
return "not_found"
|
||||
if task.status == "queued":
|
||||
task.status = "cancelled"
|
||||
task.cancel_requested_at = None
|
||||
task.updated_at = _iso(requested_at)
|
||||
_log_task_lifecycle("cancelled", task)
|
||||
return "requested"
|
||||
if task.status in {"assigned", "dispatched"}:
|
||||
if task.cancel_requested_at is not None:
|
||||
return "already_requested"
|
||||
task.cancel_requested_at = _iso(requested_at)
|
||||
task.updated_at = _iso(requested_at)
|
||||
return "requested"
|
||||
if task.status == "cancelled":
|
||||
return "requested"
|
||||
return "already_terminal"
|
||||
|
||||
def list_task_attempts(self, task_id: str) -> list[Any]:
|
||||
with self._sessions() as session:
|
||||
rows = session.scalars(
|
||||
@@ -2223,6 +2273,7 @@ def _task_from_row(row: ScheduledTaskRow) -> Any:
|
||||
progress_step_status=row.progress_step_status,
|
||||
progress_summary=row.progress_summary,
|
||||
progress_updated_at=_parse_dt(row.progress_updated_at),
|
||||
cancel_requested_at=_parse_dt(row.cancel_requested_at),
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user