91 lines
2.9 KiB
Python
91 lines
2.9 KiB
Python
"""Zerto async task polling.
|
|
|
|
Most write operations return a task id, not a result. The POST returning 200
|
|
only means the task was queued. Whether it actually happened is in
|
|
GET /v1/tasks/{id}. Skipping that read is how a failed checkpoint insert looks
|
|
like a success.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from typing import Any
|
|
|
|
from zerto_rewind_mcp.client import ZertoClient, ZertoError
|
|
|
|
# Zerto task states (9.0 Tasks API; mirrored by Zerto's published quickstarts).
|
|
FIRST_UNUSED_VALUE = 0
|
|
IN_PROGRESS = 1
|
|
WAITING_FOR_USER_INPUT = 2
|
|
PAUSED = 3
|
|
FAILED = 4
|
|
STOPPED = 5
|
|
COMPLETED = 6
|
|
CANCELLING = 7
|
|
|
|
TASK_STATE_LABELS: dict[int, str] = {
|
|
FIRST_UNUSED_VALUE: "FirstUnusedValue",
|
|
IN_PROGRESS: "InProgress",
|
|
WAITING_FOR_USER_INPUT: "WaitingForUserInput",
|
|
PAUSED: "Paused",
|
|
FAILED: "Failed",
|
|
STOPPED: "Stopped",
|
|
COMPLETED: "Completed",
|
|
CANCELLING: "Cancelling",
|
|
}
|
|
TERMINAL_STATES = {FAILED, STOPPED, COMPLETED}
|
|
|
|
|
|
def task_id_from(payload: Any) -> str:
|
|
"""A write endpoint returns the task id as a bare (often quoted) string."""
|
|
if isinstance(payload, str):
|
|
value = payload.strip().strip('"')
|
|
if value:
|
|
return value
|
|
if isinstance(payload, dict):
|
|
for key in ("TaskIdentifier", "taskIdentifier", "Identifier", "identifier", "id"):
|
|
if payload.get(key):
|
|
return str(payload[key]).strip().strip('"')
|
|
raise ZertoError(f"Could not read a task id from {payload!r}")
|
|
|
|
|
|
def task_state(payload: Any) -> int | None:
|
|
"""Status is {"State": n, "Progress": n}, but can be a bare int."""
|
|
if not isinstance(payload, dict):
|
|
return None
|
|
status = payload.get("Status")
|
|
if isinstance(status, dict):
|
|
status = status.get("State")
|
|
return status if isinstance(status, bool) is False and isinstance(status, int) else None
|
|
|
|
|
|
def task_label(state: int | None) -> str:
|
|
return TASK_STATE_LABELS.get(state, f"Unknown({state})") if state is not None else "Unknown"
|
|
|
|
|
|
async def wait_for_task(
|
|
client: ZertoClient,
|
|
task_id: str,
|
|
*,
|
|
timeout_s: float = 120.0,
|
|
interval_s: float = 3.0,
|
|
) -> dict[str, Any]:
|
|
"""Poll until the task reaches a terminal state. Raise unless it Completed."""
|
|
deadline = asyncio.get_event_loop().time() + timeout_s
|
|
last: dict[str, Any] = {}
|
|
while asyncio.get_event_loop().time() < deadline:
|
|
last = await client.get_task(task_id)
|
|
state = task_state(last)
|
|
if state in TERMINAL_STATES:
|
|
if state == COMPLETED:
|
|
return {"task_id": task_id, "state": state, "label": task_label(state)}
|
|
raise ZertoError(
|
|
f"Zerto task {task_id} ({last.get('Type')}) finished as "
|
|
f"{task_label(state)}. The operation did not happen."
|
|
)
|
|
await asyncio.sleep(interval_s)
|
|
raise ZertoError(
|
|
f"Zerto task {task_id} ({last.get('Type')}) still "
|
|
f"{task_label(task_state(last))} after {timeout_s:.0f}s."
|
|
)
|