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)