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:
@@ -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."
|
||||
)
|
||||
Reference in New Issue
Block a user