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:
2026-07-15 18:08:48 +08:00
parent 88189770ff
commit 19c6669800
9 changed files with 385 additions and 35 deletions
@@ -110,6 +110,7 @@ class ScheduledTaskRow(Base):
progress_step_status: Mapped[str | None] = mapped_column(String, nullable=True)
progress_summary: Mapped[str | None] = mapped_column(String, nullable=True)
progress_updated_at: Mapped[str | None] = mapped_column(String, nullable=True)
cancel_requested_at: Mapped[str | None] = mapped_column(String, nullable=True)
failure_reason: Mapped[str | None] = mapped_column(Text, nullable=True)
result_json: Mapped[str | None] = mapped_column(Text, nullable=True)
updated_at: Mapped[str | None] = mapped_column(String, nullable=True)
@@ -0,0 +1,22 @@
"""Add nullable cancel_requested_at column to scheduled_tasks."""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "0012_task_cancellation"
down_revision = "0011_planner_decision_log_reflection"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"scheduled_tasks",
sa.Column("cancel_requested_at", sa.String(), nullable=True),
)
def downgrade() -> None:
op.drop_column("scheduled_tasks", "cancel_requested_at")
+25 -2
View File
@@ -30,9 +30,12 @@ if TYPE_CHECKING:
AttemptStatus = Literal["assigned", "dispatched", "done", "failed", "expired"]
TerminalTaskStatus = Literal["done", "failed"]
TerminalTaskStatus = Literal["done", "failed", "cancelled"]
ResultRecordStatus = Literal["recorded", "already_recorded", "conflict"]
LeaseRenewalStatus = Literal["renewed", "not_found", "conflict", "expired"]
CancellationRequestStatus = Literal[
"requested", "already_terminal", "already_requested", "not_found"
]
class HostEnrollmentConflictError(RuntimeError):
@@ -91,6 +94,19 @@ class TaskAttemptRecord:
terminal_result: dict[str, Any] | None = None
@dataclass(frozen=True)
class LeaseRenewalResult:
"""Outcome of a lease renewal, including whether cancellation is pending.
``cancel_requested`` reflects the task's durable ``cancel_requested_at``
column at renewal time regardless of ``status`` — callers only act on it
when ``status == "renewed"``.
"""
status: LeaseRenewalStatus
cancel_requested: bool = False
@dataclass(frozen=True)
class LeasedAssignment:
task_id: str
@@ -485,7 +501,7 @@ class CloudRepository(Protocol):
lease_expires_at: datetime,
now: datetime,
progress: AssignmentProgressSnapshot | None = None,
) -> LeaseRenewalStatus: ...
) -> LeaseRenewalResult: ...
def record_task_result(
self,
@@ -500,6 +516,13 @@ class CloudRepository(Protocol):
completed_at: datetime,
) -> ResultRecordStatus: ...
def request_task_cancellation(
self,
task_id: str,
*,
requested_at: datetime,
) -> CancellationRequestStatus: ...
def reap_expired_leases(
self,
*,
+4 -1
View File
@@ -22,7 +22,9 @@ if TYPE_CHECKING:
from cloud.store import CloudStore
ScheduledTaskStatus = Literal["queued", "assigned", "dispatched", "done", "failed"]
ScheduledTaskStatus = Literal[
"queued", "assigned", "dispatched", "done", "failed", "cancelled"
]
@dataclass(frozen=True)
@@ -57,6 +59,7 @@ class ScheduledTask:
progress_step_status: str | None = None
progress_summary: str | None = None
progress_updated_at: datetime | None = None
cancel_requested_at: datetime | None = None
@runtime_checkable
+1 -1
View File
@@ -9,7 +9,7 @@ from alembic.runtime.migration import MigrationContext
from cloud.database import create_database_engine, normalize_database_url
HEAD_REVISION = "0011_planner_decision_log_reflection"
HEAD_REVISION = "0012_task_cancellation"
class SchemaVersionError(RuntimeError):
+67 -16
View File
@@ -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),
)