# Conflicts: # packages/cloud-platform/cloud/schema.py
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -10,7 +10,7 @@ from time import monotonic
|
||||
from typing import TYPE_CHECKING
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request, status
|
||||
from fastapi import APIRouter, HTTPException, Request, Response, status
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from cloud.auth import (
|
||||
@@ -29,6 +29,7 @@ from cloud.internal_api.models import (
|
||||
HostGovernancePolicyModel,
|
||||
HostEnrollmentRequest,
|
||||
HostEnrollmentResponse,
|
||||
HostTaskCancellationResponse,
|
||||
HostTaskSubmissionRequest,
|
||||
HostTaskSubmissionResponse,
|
||||
LeaseRenewalRequest,
|
||||
@@ -250,6 +251,50 @@ def create_internal_router(
|
||||
)
|
||||
return HostTaskSubmissionResponse(task_id=task_id)
|
||||
|
||||
@router.post(
|
||||
"/hosts/{host_id}/tasks/{task_id}/cancel",
|
||||
response_model=HostTaskCancellationResponse,
|
||||
responses={
|
||||
status.HTTP_202_ACCEPTED: {"model": HostTaskCancellationResponse},
|
||||
},
|
||||
)
|
||||
def cancel_host_task(
|
||||
host_id: str,
|
||||
task_id: str,
|
||||
request: Request,
|
||||
response: Response,
|
||||
) -> HostTaskCancellationResponse:
|
||||
authorize_host(request, host_id)
|
||||
if scheduler is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="task cancellation is unavailable",
|
||||
)
|
||||
task = scheduler.store.get_task(task_id)
|
||||
if task is None or task.constraints.target_host_id != host_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"task {task_id!r} not found",
|
||||
)
|
||||
result = scheduler.store.request_task_cancellation(
|
||||
task_id, requested_at=utc_now()
|
||||
)
|
||||
if result == "not_found":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"task {task_id!r} not found",
|
||||
)
|
||||
if result == "already_terminal":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"task {task_id!r} has already reached a terminal state",
|
||||
)
|
||||
task = scheduler.store.get_task(task_id)
|
||||
assert task is not None
|
||||
if result == "requested" and task.status != "cancelled":
|
||||
response.status_code = status.HTTP_202_ACCEPTED
|
||||
return HostTaskCancellationResponse(task_id=task_id, status=task.status)
|
||||
|
||||
@router.post(
|
||||
"/hosts/{host_id}/assignments/claim",
|
||||
response_model=ClaimResponse,
|
||||
@@ -316,7 +361,7 @@ def create_internal_router(
|
||||
summary=payload.progress.summary[:500],
|
||||
updated_at=now,
|
||||
)
|
||||
renewal_status = pool.store.renew_lease(
|
||||
renewal = pool.store.renew_lease(
|
||||
task_id=task_id,
|
||||
attempt=payload.attempt,
|
||||
lease_id=payload.lease_id,
|
||||
@@ -325,16 +370,17 @@ def create_internal_router(
|
||||
now=now,
|
||||
progress=progress_snapshot,
|
||||
)
|
||||
if renewal_status == "not_found":
|
||||
if renewal.status == "not_found":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="assignment not found",
|
||||
)
|
||||
if renewal_status != "renewed":
|
||||
if renewal.status != "renewed":
|
||||
return _stale_lease_conflict("assignment lease is stale or expired")
|
||||
return LeaseRenewalResponse(
|
||||
status="renewed",
|
||||
lease_expires_at=lease_expires_at,
|
||||
cancel_requested=renewal.cancel_requested,
|
||||
)
|
||||
|
||||
@router.post(
|
||||
|
||||
@@ -95,6 +95,7 @@ class LeaseRenewalRequest(BaseModel):
|
||||
class LeaseRenewalResponse(BaseModel):
|
||||
status: Literal["renewed"]
|
||||
lease_expires_at: datetime
|
||||
cancel_requested: bool = False
|
||||
|
||||
|
||||
class TerminalResultRequest(BaseModel):
|
||||
@@ -102,7 +103,7 @@ class TerminalResultRequest(BaseModel):
|
||||
task_id: str = Field(min_length=1)
|
||||
attempt: int = Field(ge=1)
|
||||
lease_id: str = Field(min_length=1)
|
||||
status: Literal["done", "failed"]
|
||||
status: Literal["done", "failed", "cancelled"]
|
||||
failure_reason: str | None = None
|
||||
result: dict[str, Any] | None = None
|
||||
|
||||
@@ -121,6 +122,11 @@ class HostTaskSubmissionResponse(BaseModel):
|
||||
task_id: str
|
||||
|
||||
|
||||
class HostTaskCancellationResponse(BaseModel):
|
||||
task_id: str
|
||||
status: str
|
||||
|
||||
|
||||
class StaleLeaseConflict(BaseModel):
|
||||
code: Literal["stale_lease"] = "stale_lease"
|
||||
detail: str
|
||||
|
||||
@@ -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 = "0013_task_cancellation"
|
||||
down_revision = "0012_planner_decision_action_metadata"
|
||||
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")
|
||||
@@ -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
|
||||
@@ -487,7 +503,7 @@ class CloudRepository(Protocol):
|
||||
lease_expires_at: datetime,
|
||||
now: datetime,
|
||||
progress: AssignmentProgressSnapshot | None = None,
|
||||
) -> LeaseRenewalStatus: ...
|
||||
) -> LeaseRenewalResult: ...
|
||||
|
||||
def record_task_result(
|
||||
self,
|
||||
@@ -502,6 +518,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,
|
||||
*,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,7 +9,7 @@ from alembic.runtime.migration import MigrationContext
|
||||
from cloud.database import create_database_engine, normalize_database_url
|
||||
|
||||
|
||||
HEAD_REVISION = "0012_planner_decision_action_metadata"
|
||||
HEAD_REVISION = "0013_task_cancellation"
|
||||
|
||||
|
||||
class SchemaVersionError(RuntimeError):
|
||||
|
||||
@@ -11,6 +11,7 @@ authentication can be added later without changing route signatures.
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Callable, Literal
|
||||
|
||||
from cloud.auth import (
|
||||
@@ -31,6 +32,7 @@ from cloud.sdk.models import (
|
||||
PluginRegistrationRequest,
|
||||
PluginResponse,
|
||||
TaskAttemptResponse,
|
||||
TaskCancellationResponse,
|
||||
TaskListItem,
|
||||
TaskListResponse,
|
||||
TaskPlannerDecisionItem,
|
||||
@@ -39,7 +41,7 @@ from cloud.sdk.models import (
|
||||
TaskSubmissionRequest,
|
||||
TaskSubmissionResponse,
|
||||
)
|
||||
from fastapi import APIRouter, HTTPException, Query, Request, status
|
||||
from fastapi import APIRouter, HTTPException, Query, Request, Response, status
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cloud.plugins import PluginRegistry
|
||||
@@ -156,7 +158,9 @@ def create_cloud_router(
|
||||
@router.get("/tasks", response_model=TaskListResponse)
|
||||
def list_tasks(
|
||||
request: Request,
|
||||
status_filter: Literal["queued", "assigned", "dispatched", "done", "failed"]
|
||||
status_filter: Literal[
|
||||
"queued", "assigned", "dispatched", "done", "failed", "cancelled"
|
||||
]
|
||||
| None = Query(default=None, alias="status"),
|
||||
limit: int = Query(default=50, ge=1, le=100),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
@@ -272,6 +276,36 @@ def create_cloud_router(
|
||||
)
|
||||
return TaskPlannerDecisionListResponse(items=items)
|
||||
|
||||
@router.post(
|
||||
"/tasks/{task_id}/cancel",
|
||||
response_model=TaskCancellationResponse,
|
||||
responses={
|
||||
status.HTTP_202_ACCEPTED: {"model": TaskCancellationResponse},
|
||||
},
|
||||
)
|
||||
def cancel_task(
|
||||
task_id: str, request: Request, response: Response
|
||||
) -> TaskCancellationResponse:
|
||||
_authorize(request, TASKS_SUBMIT_SCOPE)
|
||||
result = scheduler.store.request_task_cancellation(
|
||||
task_id, requested_at=datetime.now(UTC)
|
||||
)
|
||||
if result == "not_found":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"task {task_id!r} not found",
|
||||
)
|
||||
if result == "already_terminal":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"task {task_id!r} has already reached a terminal state",
|
||||
)
|
||||
task = scheduler.store.get_task(task_id)
|
||||
assert task is not None
|
||||
if result == "requested" and task.status != "cancelled":
|
||||
response.status_code = status.HTTP_202_ACCEPTED
|
||||
return TaskCancellationResponse(task_id=task_id, status=task.status)
|
||||
|
||||
@router.get("/devices", response_model=list[DeviceResponse])
|
||||
def list_devices(request: Request) -> list[DeviceResponse]:
|
||||
_authorize(request, POOL_READ_SCOPE)
|
||||
|
||||
@@ -115,6 +115,10 @@ class CloudClient:
|
||||
resp = self._request("GET", f"/tasks/{task_id}/attempts")
|
||||
return resp.json()
|
||||
|
||||
def cancel_task(self, task_id: str) -> dict[str, Any]:
|
||||
resp = self._request("POST", f"/tasks/{task_id}/cancel")
|
||||
return resp.json()
|
||||
|
||||
# ----------------------------------------------------------------- devices
|
||||
|
||||
def list_devices(self) -> list[dict[str, Any]]:
|
||||
|
||||
@@ -25,6 +25,11 @@ class TaskSubmissionResponse(BaseModel):
|
||||
task_id: str
|
||||
|
||||
|
||||
class TaskCancellationResponse(BaseModel):
|
||||
task_id: str
|
||||
status: str
|
||||
|
||||
|
||||
class TaskStatusResponse(BaseModel):
|
||||
id: str
|
||||
status: str
|
||||
|
||||
@@ -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(
|
||||
@@ -2217,6 +2267,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