feat(guard): read the Zerto task, and ask before guarding unknown tools (#6)

This commit was merged in pull request #6.
This commit is contained in:
2026-09-21 15:07:43 -04:00
parent 5039f7378d
commit 1d53029038
10 changed files with 359 additions and 22 deletions
+90
View File
@@ -0,0 +1,90 @@
"""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."
)