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:
@@ -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