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
+19
View File
@@ -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"]
+69
View File
@@ -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)