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:
@@ -36,12 +36,42 @@ def entry_from_dict(raw: dict[str, Any]) -> CatalogEntry:
|
||||
)
|
||||
|
||||
|
||||
MUTATING = "mutating"
|
||||
READ_ONLY = "read_only"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
class MutatingCatalog:
|
||||
def __init__(self, entries: list[CatalogEntry] | None = None, path: Path | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
entries: list[CatalogEntry] | None = None,
|
||||
path: Path | None = None,
|
||||
read_only: list[CatalogEntry] | None = None,
|
||||
):
|
||||
self._entries: dict[tuple[str, str], CatalogEntry] = {}
|
||||
self._read_only: dict[tuple[str, str], CatalogEntry] = {}
|
||||
self.path = path
|
||||
for entry in entries or []:
|
||||
self._entries[entry.key()] = entry
|
||||
for entry in read_only or []:
|
||||
self._read_only[entry.key()] = entry
|
||||
|
||||
def read_only(self) -> list[CatalogEntry]:
|
||||
return sorted(self._read_only.values(), key=lambda e: e.key())
|
||||
|
||||
def classify(self, server: str, tool: str) -> str:
|
||||
"""mutating / read_only / unknown.
|
||||
|
||||
unknown is the important one. It does not mean safe: it means nobody has
|
||||
said this tool is read-only, so the agent must ask a human before
|
||||
changing a protected VM with it.
|
||||
"""
|
||||
key = (server.lower(), tool.lower())
|
||||
if key in self._entries:
|
||||
return MUTATING
|
||||
if key in self._read_only:
|
||||
return READ_ONLY
|
||||
return UNKNOWN
|
||||
|
||||
def list(self) -> list[CatalogEntry]:
|
||||
return sorted(self._entries.values(), key=lambda e: e.key())
|
||||
@@ -68,6 +98,9 @@ class MutatingCatalog:
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, data: dict[str, Any], path: Path | None = None) -> MutatingCatalog:
|
||||
raw = data.get("mutating_tools") or []
|
||||
entries = [entry_from_dict(item) for item in raw]
|
||||
return cls(entries, path=path)
|
||||
entries = [entry_from_dict(item) for item in data.get("mutating_tools") or []]
|
||||
read_only = [
|
||||
entry_from_dict({**item, "vm_arg": item.get("vm_arg") or "-"})
|
||||
for item in data.get("read_only_tools") or []
|
||||
]
|
||||
return cls(entries, path=path, read_only=read_only)
|
||||
|
||||
@@ -9,6 +9,7 @@ from typing import Any
|
||||
|
||||
from zerto_rewind_mcp.client import ZertoClient, ZertoError
|
||||
from zerto_rewind_mcp.protection import FindResult
|
||||
from zerto_rewind_mcp.tasks import task_id_from, wait_for_task
|
||||
from zerto_rewind_mcp.util import pick
|
||||
|
||||
_CTRL = re.compile(r"[\x00-\x1f\x7f]+")
|
||||
@@ -108,16 +109,20 @@ async def tag_vpgs(
|
||||
"message": result.message,
|
||||
"find": result.as_dict(),
|
||||
}
|
||||
# Insert then wait, one VPG at a time. Measured on 10.x: tagged checkpoint
|
||||
# inserts fired back to back at the same VPG are silently dropped -- the POST
|
||||
# returns 200 and queues a task, but only the first checkpoint ever appears.
|
||||
# Do not turn this loop into an asyncio.gather().
|
||||
# Insert then wait for the TASK, one VPG at a time. The POST returns a task
|
||||
# id, not a result: 200 only means queued. Measured on 10.x, a second insert
|
||||
# fired at the same VPG while the first is running comes back Failed (state
|
||||
# 4) from GET /v1/tasks/{id} while the first is Completed (6). That failure
|
||||
# is reported, not silent -- but only if you read the task. Do not turn this
|
||||
# loop into an asyncio.gather().
|
||||
tagged: list[dict[str, Any]] = []
|
||||
skipped = [v.as_dict() for v in result.vm.vpgs if not v.can_tag]
|
||||
errors: list[str] = []
|
||||
for vpg in result.taggable_vpgs:
|
||||
try:
|
||||
insert = await client.insert_checkpoint(vpg.vpg_identifier, tag)
|
||||
task_id = task_id_from(insert)
|
||||
task = await wait_for_task(client, task_id)
|
||||
row = await wait_for_tag(client, vpg.vpg_identifier, tag)
|
||||
tagged.append(
|
||||
{
|
||||
@@ -125,7 +130,8 @@ async def tag_vpgs(
|
||||
"vpg_name": vpg.vpg_name,
|
||||
"checkpoint_id": checkpoint_id(row),
|
||||
"tag": checkpoint_tag(row) or tag,
|
||||
"insert_result": insert if not isinstance(insert, dict) else "ok",
|
||||
"task_id": task_id,
|
||||
"task_state": task["label"],
|
||||
}
|
||||
)
|
||||
except ZertoError as exc:
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any
|
||||
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
|
||||
from zerto_rewind_mcp.catalog import CatalogEntry, MutatingCatalog
|
||||
from zerto_rewind_mcp.catalog import MUTATING, READ_ONLY, CatalogEntry, MutatingCatalog
|
||||
from zerto_rewind_mcp.checkpoints import checkpoint_id, checkpoint_tag, make_tag, tag_vpgs
|
||||
from zerto_rewind_mcp.client import ZertoClient, ZertoError
|
||||
from zerto_rewind_mcp.config import load_catalog, load_config
|
||||
@@ -219,6 +219,69 @@ async def zerto_guard_before_mutate(
|
||||
)
|
||||
|
||||
|
||||
@mcp.tool(
|
||||
name="zerto_check_tool",
|
||||
annotations={
|
||||
"title": "Is this tool safe to run on a protected VM?",
|
||||
"readOnlyHint": True,
|
||||
"destructiveHint": False,
|
||||
"idempotentHint": True,
|
||||
"openWorldHint": False,
|
||||
},
|
||||
)
|
||||
async def zerto_check_tool(server: str, tool: str, vm: str = "") -> str:
|
||||
"""Classify a tool before you run it against a VM. Call this first.
|
||||
|
||||
Returns one of:
|
||||
|
||||
read_only - known not to change the guest. Run it. No checkpoint.
|
||||
mutating - known to change the guest. Guard first, do not ask.
|
||||
unknown - NOT known to be read-only. Do not assume it is safe and do not
|
||||
silently guard either. ASK THE HUMAN whether to insert a
|
||||
tagged checkpoint before running it, then do what they say.
|
||||
|
||||
unknown is the normal case for any tool nobody has classified yet, which is
|
||||
most of them. The catalogs are short; the world of tools is not.
|
||||
"""
|
||||
verdict = get_catalog().classify(server, tool)
|
||||
payload: dict[str, Any] = {
|
||||
"ok": True,
|
||||
"server": server,
|
||||
"tool": tool,
|
||||
"vm": vm or None,
|
||||
"verdict": verdict,
|
||||
}
|
||||
target = vm or "<the VM this tool targets>"
|
||||
if verdict == READ_ONLY:
|
||||
payload["checkpoint_required"] = False
|
||||
payload["next"] = "Run the tool. Reads do not need a checkpoint."
|
||||
elif verdict == MUTATING:
|
||||
payload["checkpoint_required"] = True
|
||||
payload["next"] = (
|
||||
f"Call zerto_guard_before_mutate(query={target!r}, change_id=..., "
|
||||
"action=<what you are about to do>) and only run the tool if it "
|
||||
"returns ok true."
|
||||
)
|
||||
else:
|
||||
payload["checkpoint_required"] = None
|
||||
payload["ask_user"] = True
|
||||
payload["question"] = (
|
||||
f"{server}/{tool} is not a known read-only command. It may change "
|
||||
f"{target}. Insert a Zerto tagged checkpoint first so this is "
|
||||
"reversible?"
|
||||
)
|
||||
payload["next"] = (
|
||||
"Ask the human the question above and wait for an answer. Do not "
|
||||
"guess. If they say yes, call zerto_guard_before_mutate("
|
||||
f"query={target!r}, change_id=..., action=<what you are about to "
|
||||
"do>) and run the tool only if it returns ok true. If they say no, "
|
||||
"run the tool without a checkpoint and tell them it is not "
|
||||
"reversible through Zerto. If they want this remembered, call "
|
||||
"zerto_add_mutating_tool so it is guarded automatically next time."
|
||||
)
|
||||
return _dump(payload)
|
||||
|
||||
|
||||
@mcp.tool(
|
||||
name="zerto_list_mutating_tools",
|
||||
annotations={
|
||||
@@ -230,9 +293,22 @@ async def zerto_guard_before_mutate(
|
||||
},
|
||||
)
|
||||
async def zerto_list_mutating_tools() -> str:
|
||||
"""Tools that must be guarded. Unlisted tools pass through."""
|
||||
entries = [e.as_dict() for e in get_catalog().list()]
|
||||
return _dump({"ok": True, "count": len(entries), "mutating_tools": entries})
|
||||
"""Both catalogs. A tool in neither is unknown: ask the human before mutating."""
|
||||
cat = get_catalog()
|
||||
entries = [e.as_dict() for e in cat.list()]
|
||||
read_only = [e.as_dict() for e in cat.read_only()]
|
||||
return _dump(
|
||||
{
|
||||
"ok": True,
|
||||
"count": len(entries),
|
||||
"mutating_tools": entries,
|
||||
"read_only_tools": read_only,
|
||||
"unlisted": (
|
||||
"A tool in neither list is unknown, not safe. Call "
|
||||
"zerto_check_tool and ask the human before changing a protected VM."
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@mcp.tool(
|
||||
|
||||
@@ -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