"""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." )