feat(guard): read the Zerto task, and ask before guarding unknown tools
Two changes that both come from the same mistake: assuming an answer
instead of reading one.
1. Read the task after inserting a checkpoint.
POST /v1/vpgs/{id}/checkpoints returns a TASK ID, not a result. A 200
only means queued. The outcome lives in GET /v1/tasks/{id} under
Status.State: 1 InProgress, 4 Failed, 5 Stopped, 6 Completed, with
4/5/6 terminal.
tag_vpgs now waits for that task and refuses unless it Completed, and
reports task_id and task_state. Measured: two inserts fired back to
back at one VPG give Completed for the first and Failed for the
second. That is exactly the case an earlier comment in this file
called "silently dropped" -- it was never silent, we just never read
the task. Comment corrected.
Previously a failed insert surfaced only as wait_for_tag timing out
45s later with a misleading hint about Azure/AWS. Now it says the
task failed and the operation did not happen.
2. Unknown tools ask the human instead of passing through.
The catalog is opt-in, so an unlisted tool ran unguarded. But the set
of mutating tools is unbounded and grows with every MCP installed,
while the set of read-only ones is small, so a mutating-only list is
permanently behind and being behind fails open.
Adds read_only_tools and zerto_check_tool(server, tool, vm) returning
read_only / mutating / unknown. unknown does not mean safe: it means
nobody classified it, so the tool hands the agent a question to put
to the human, and the human decides whether to checkpoint. On yes the
agent guards; on no it runs and says plainly that Zerto cannot rewind
it; if they want it remembered, zerto_add_mutating_tool.
This stays advisory. An MCP server cannot see or block another
server's tool calls, so real enforcement belongs in a host PreToolUse
hook. The skill carries the flow.
Known issue, not addressed here: two concurrent tool calls race on
Keycloak token acquisition in the shared client and one gets HTTP 401.
pytest 41 passed (12 new).
Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>
Claude-Session: https://claude.ai/code/session_016yVfC5nvZowoLFnEGWhLGn
This commit is contained in:
@@ -50,3 +50,22 @@ def test_example_config_covers_windows_and_linux():
|
||||
# every entry must name the arg holding the VM, or the guard cannot resolve one
|
||||
for entry in cat.list():
|
||||
assert entry.vm_arg, f"{entry.server}/{entry.tool} has no vm_arg"
|
||||
|
||||
|
||||
def test_classify_three_way():
|
||||
data = {
|
||||
"mutating_tools": [{"server": "winrm", "tool": "run_ps", "vm_arg": "host"}],
|
||||
"read_only_tools": [{"server": "ssh", "tool": "read_file"}],
|
||||
}
|
||||
cat = MutatingCatalog.from_config(data)
|
||||
assert cat.classify("winrm", "run_ps") == "mutating"
|
||||
assert cat.classify("SSH", "Read_File") == "read_only"
|
||||
# unknown is not safe; it means nobody classified it
|
||||
assert cat.classify("anything", "else") == "unknown"
|
||||
|
||||
|
||||
def test_read_only_entries_need_no_vm_arg():
|
||||
data = {"read_only_tools": [{"server": "ssh", "tool": "stat"}]}
|
||||
cat = MutatingCatalog.from_config(data)
|
||||
assert cat.classify("ssh", "stat") == "read_only"
|
||||
assert [e.tool for e in cat.read_only()] == ["stat"]
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from zerto_rewind_mcp.client import ZertoError
|
||||
from zerto_rewind_mcp.tasks import (
|
||||
COMPLETED,
|
||||
FAILED,
|
||||
task_id_from,
|
||||
task_label,
|
||||
task_state,
|
||||
wait_for_task,
|
||||
)
|
||||
|
||||
|
||||
def test_task_id_from_shapes():
|
||||
# a write endpoint returns the id as a bare, often quoted, string
|
||||
assert task_id_from('"abc.def"') == "abc.def"
|
||||
assert task_id_from("abc.def") == "abc.def"
|
||||
assert task_id_from({"TaskIdentifier": "t1"}) == "t1"
|
||||
with pytest.raises(ZertoError):
|
||||
task_id_from(None)
|
||||
|
||||
|
||||
def test_task_state_handles_nested_and_bare():
|
||||
assert task_state({"Status": {"State": 6, "Progress": 100}}) == COMPLETED
|
||||
assert task_state({"Status": 4}) == FAILED
|
||||
assert task_state({"Status": None}) is None
|
||||
assert task_state("nope") is None
|
||||
|
||||
|
||||
def test_task_labels():
|
||||
assert task_label(COMPLETED) == "Completed"
|
||||
assert task_label(FAILED) == "Failed"
|
||||
|
||||
|
||||
class _TaskClient:
|
||||
def __init__(self, states):
|
||||
self.states = list(states)
|
||||
self.calls = 0
|
||||
|
||||
async def get_task(self, task_id):
|
||||
self.calls += 1
|
||||
state = self.states.pop(0) if self.states else self.states
|
||||
return {"Type": "InsertTaggedCP", "Status": {"State": state, "Progress": 100}}
|
||||
|
||||
|
||||
def test_wait_for_task_returns_on_completed():
|
||||
client = _TaskClient([1, 1, COMPLETED])
|
||||
out = asyncio.run(wait_for_task(client, "t1", interval_s=0))
|
||||
assert out["state"] == COMPLETED
|
||||
assert out["label"] == "Completed"
|
||||
assert client.calls == 3
|
||||
|
||||
|
||||
def test_wait_for_task_raises_on_failed():
|
||||
# this is the case that used to look like success: 200 queued, task Failed
|
||||
client = _TaskClient([1, FAILED])
|
||||
with pytest.raises(ZertoError) as err:
|
||||
asyncio.run(wait_for_task(client, "t1", interval_s=0))
|
||||
assert "Failed" in str(err.value)
|
||||
assert "did not happen" in str(err.value)
|
||||
|
||||
|
||||
def test_wait_for_task_times_out_while_in_progress():
|
||||
client = _TaskClient([1] * 50)
|
||||
with pytest.raises(ZertoError) as err:
|
||||
asyncio.run(wait_for_task(client, "t1", timeout_s=0.05, interval_s=0.01))
|
||||
assert "InProgress" in str(err.value)
|
||||
Reference in New Issue
Block a user